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

Method update

domainbed/algorithms.py:943–960  ·  view source on GitHub ↗
(self, minibatches, unlabeled=None)

Source from the content-addressed store, hash-verified

941 self.tau = hparams["tau"]
942
943 def update(self, minibatches, unlabeled=None):
944 mean_loss = 0
945 param_gradients = [[] for _ in self.network.parameters()]
946 for i, (x, y) in enumerate(minibatches):
947 logits = self.network(x)
948
949 env_loss = F.cross_entropy(logits, y)
950 mean_loss += env_loss.item() / len(minibatches)
951
952 env_grads = autograd.grad(env_loss, self.network.parameters())
953 for grads, env_grad in zip(param_gradients, env_grads):
954 grads.append(env_grad)
955
956 self.optimizer.zero_grad()
957 self.mask_grads(self.tau, param_gradients, self.network.parameters())
958 self.optimizer.step()
959
960 return {'loss': mean_loss}
961
962 def mask_grads(self, tau, gradients, params):
963

Callers

nothing calls this directly

Calls 1

mask_gradsMethod · 0.95

Tested by

no test coverage detected