| 736 | """ |
| 737 | |
| 738 | def __init__(self, input_shape, num_classes, num_domains, hparams): |
| 739 | super(SagNet, self).__init__(input_shape, num_classes, num_domains, hparams) |
| 740 | # featurizer network |
| 741 | self.network_f = networks.Featurizer(input_shape, self.hparams) |
| 742 | # content network |
| 743 | self.network_c = networks.Classifier( |
| 744 | self.network_f.n_outputs, num_classes, self.hparams['nonlinear_classifier'] |
| 745 | ) |
| 746 | # style network |
| 747 | self.network_s = networks.Classifier( |
| 748 | self.network_f.n_outputs, num_classes, self.hparams['nonlinear_classifier'] |
| 749 | ) |
| 750 | |
| 751 | # # This commented block of code implements something closer to the |
| 752 | # # original paper, but is specific to ResNet and puts in disadvantage |
| 753 | # # the other algorithms. |
| 754 | # resnet_c = networks.Featurizer(input_shape, self.hparams) |
| 755 | # resnet_s = networks.Featurizer(input_shape, self.hparams) |
| 756 | # # featurizer network |
| 757 | # self.network_f = torch.nn.Sequential( |
| 758 | # resnet_c.network.conv1, |
| 759 | # resnet_c.network.bn1, |
| 760 | # resnet_c.network.relu, |
| 761 | # resnet_c.network.maxpool, |
| 762 | # resnet_c.network.layer1, |
| 763 | # resnet_c.network.layer2, |
| 764 | # resnet_c.network.layer3) |
| 765 | # # content network |
| 766 | # self.network_c = torch.nn.Sequential( |
| 767 | # resnet_c.network.layer4, |
| 768 | # resnet_c.network.avgpool, |
| 769 | # networks.Flatten(), |
| 770 | # resnet_c.network.fc) |
| 771 | # # style network |
| 772 | # self.network_s = torch.nn.Sequential( |
| 773 | # resnet_s.network.layer4, |
| 774 | # resnet_s.network.avgpool, |
| 775 | # networks.Flatten(), |
| 776 | # resnet_s.network.fc) |
| 777 | |
| 778 | def opt(p): |
| 779 | return torch.optim.Adam(p, lr=hparams["lr"], weight_decay=hparams["weight_decay"]) |
| 780 | |
| 781 | self.optimizer_f = opt(self.network_f.parameters()) |
| 782 | self.optimizer_c = opt(self.network_c.parameters()) |
| 783 | self.optimizer_s = opt(self.network_s.parameters()) |
| 784 | self.weight_adv = hparams["sag_w_adv"] |
| 785 | |
| 786 | def forward_c(self, x): |
| 787 | # learning content network on randomized style |