from torch import nn from torch.nn import init def weights_init_kaiming(m): classname = m.__class__.__name__ # print(classname) if classname.find('Conv') != -1: init.kaiming_normal_(m.weight.data, a=0, mode='fan_in') elif classname.find('Linear') != -1: init.kaiming_normal_(m.weight.data, a=0, mode='fan_out') init.constant_(m.bias.data, 0.0) elif classname.find('BatchNorm1d') != -1: init.normal_(m.weight.data, 1.0, 0.02) init.constant_(m.bias.data, 0.0) def weights_init_classifier(m): classname = m.__class__.__name__ if classname.find('Linear') != -1: init.normal_(m.weight.data, std=0.001) init.constant_(m.bias.data, 0.0) # Defines the new fc layer and classification layer # |--Linear--|--bn--|--relu--|--Linear--| class ClassBlock(nn.Module): def __init__(self, input_dim, class_num=1, activ='sigmoid', num_bottleneck=512): super(ClassBlock, self).__init__() add_block = [] add_block += [nn.Linear(input_dim, num_bottleneck)] add_block += [nn.BatchNorm1d(num_bottleneck)] add_block += [nn.LeakyReLU(0.1)] add_block += [nn.Dropout(p=0.5)] add_block = nn.Sequential(*add_block) add_block.apply(weights_init_kaiming) classifier = [] classifier += [nn.Linear(num_bottleneck, class_num)] if activ == 'sigmoid': classifier += [nn.Sigmoid()] elif activ == 'softmax': classifier += [nn.Softmax()] elif activ == 'none': classifier += [] else: raise AssertionError("Unsupported activation: {}".format(activ)) classifier = nn.Sequential(*classifier) classifier.apply(weights_init_classifier) self.add_block = add_block self.classifier = classifier def forward(self, x): x = self.add_block(x) x = self.classifier(x) return x