| 985 | super(IGA, self).__init__(in_features, num_classes, num_domains, hparams) |
| 986 | |
| 987 | def update(self, minibatches, unlabeled=False): |
| 988 | total_loss = 0 |
| 989 | grads = [] |
| 990 | for i, (x, y) in enumerate(minibatches): |
| 991 | logits = self.network(x) |
| 992 | |
| 993 | env_loss = F.cross_entropy(logits, y) |
| 994 | total_loss += env_loss |
| 995 | |
| 996 | env_grad = autograd.grad(env_loss, self.network.parameters(), create_graph=True) |
| 997 | |
| 998 | grads.append(env_grad) |
| 999 | |
| 1000 | mean_loss = total_loss / len(minibatches) |
| 1001 | mean_grad = autograd.grad(mean_loss, self.network.parameters(), retain_graph=True) |
| 1002 | |
| 1003 | # compute trace penalty |
| 1004 | penalty_value = 0 |
| 1005 | for grad in grads: |
| 1006 | for g, mean_g in zip(grad, mean_grad): |
| 1007 | penalty_value += (g - mean_g).pow(2).sum() |
| 1008 | |
| 1009 | objective = mean_loss + self.hparams['penalty'] * penalty_value |
| 1010 | |
| 1011 | self.optimizer.zero_grad() |
| 1012 | objective.backward() |
| 1013 | self.optimizer.step() |
| 1014 | |
| 1015 | return {'loss': mean_loss.item(), 'penalty': penalty_value.item()} |
| 1016 | |
| 1017 | |
| 1018 | class SelfReg(ERM): |