MCPcopy Create free account
hub / github.com/alexrame/fishr / update

Method update

domainbed/algorithms.py:987–1015  ·  view source on GitHub ↗
(self, minibatches, unlabeled=False)

Source from the content-addressed store, hash-verified

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
1018class SelfReg(ERM):

Callers

nothing calls this directly

Calls 1

sumMethod · 0.80

Tested by

no test coverage detected