| 201 | """Domain-Adversarial Neural Networks (abstract class)""" |
| 202 | |
| 203 | def __init__(self, input_shape, num_classes, num_domains, hparams, conditional, class_balance): |
| 204 | |
| 205 | super(AbstractDANN, self).__init__(input_shape, num_classes, num_domains, hparams) |
| 206 | |
| 207 | self.register_buffer('update_count', torch.tensor([0])) |
| 208 | self.conditional = conditional |
| 209 | self.class_balance = class_balance |
| 210 | |
| 211 | # Algorithms |
| 212 | self.featurizer = networks.Featurizer(input_shape, self.hparams) |
| 213 | self.classifier = networks.Classifier( |
| 214 | self.featurizer.n_outputs, num_classes, self.hparams['nonlinear_classifier'] |
| 215 | ) |
| 216 | self.discriminator = networks.MLP(self.featurizer.n_outputs, num_domains, self.hparams) |
| 217 | self.class_embeddings = nn.Embedding(num_classes, self.featurizer.n_outputs) |
| 218 | |
| 219 | # Optimizers |
| 220 | self.disc_opt = torch.optim.Adam( |
| 221 | (list(self.discriminator.parameters()) + list(self.class_embeddings.parameters())), |
| 222 | lr=self.hparams["lr_d"], |
| 223 | weight_decay=self.hparams['weight_decay_d'], |
| 224 | betas=(self.hparams['beta1'], 0.9) |
| 225 | ) |
| 226 | |
| 227 | self.gen_opt = torch.optim.Adam( |
| 228 | (list(self.featurizer.parameters()) + list(self.classifier.parameters())), |
| 229 | lr=self.hparams["lr_g"], |
| 230 | weight_decay=self.hparams['weight_decay_g'], |
| 231 | betas=(self.hparams['beta1'], 0.9) |
| 232 | ) |
| 233 | |
| 234 | def update(self, minibatches, unlabeled=None): |
| 235 | device = "cuda" if minibatches[0][0].is_cuda else "cpu" |