(self, input_shape, num_classes, num_domains, hparams)
| 1170 | "Invariant Gradients variances for Out-of-distribution Generalization" |
| 1171 | |
| 1172 | def __init__(self, input_shape, num_classes, num_domains, hparams): |
| 1173 | assert backpack is not None, "Install backpack with: 'pip install backpack-for-pytorch==1.3.0'" |
| 1174 | super(Fishr, self).__init__(input_shape, num_classes, num_domains, hparams) |
| 1175 | self.num_domains = num_domains |
| 1176 | |
| 1177 | self.featurizer = networks.Featurizer(input_shape, self.hparams) |
| 1178 | self.classifier = extend( |
| 1179 | networks.Classifier( |
| 1180 | self.featurizer.n_outputs, |
| 1181 | num_classes, |
| 1182 | self.hparams['nonlinear_classifier'], |
| 1183 | ) |
| 1184 | ) |
| 1185 | self.network = nn.Sequential(self.featurizer, self.classifier) |
| 1186 | |
| 1187 | self.register_buffer("update_count", torch.tensor([0])) |
| 1188 | self.bce_extended = extend(nn.CrossEntropyLoss(reduction='none')) |
| 1189 | self.ema_per_domain = [ |
| 1190 | MovingAverage(ema=self.hparams["ema"], oneminusema_correction=True) |
| 1191 | for _ in range(self.num_domains) |
| 1192 | ] |
| 1193 | self._init_optimizer() |
| 1194 | |
| 1195 | def _init_optimizer(self): |
| 1196 | self.optimizer = torch.optim.Adam( |
nothing calls this directly
no test coverage detected