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

Method update

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

Source from the content-addressed store, hash-verified

1124 self.register_buffer('update_count', torch.tensor([0]))
1125
1126 def update(self, minibatches, unlabeled=None):
1127
1128 mean_loss = 0
1129 param_gradients = [[] for _ in self.network.parameters()]
1130 for i, (x, y) in enumerate(minibatches):
1131 logits = self.network(x)
1132
1133 env_loss = F.cross_entropy(logits, y)
1134 mean_loss += env_loss.item() / len(minibatches)
1135 env_grads = autograd.grad(env_loss, self.network.parameters(), retain_graph=True)
1136 for grads, env_grad in zip(param_gradients, env_grads):
1137 grads.append(env_grad)
1138
1139 self.optimizer.zero_grad()
1140 # gradient masking applied here
1141 self.mask_grads(param_gradients, self.network.parameters())
1142 self.optimizer.step()
1143 self.update_count += 1
1144
1145 return {'loss': mean_loss}
1146
1147 def mask_grads(self, gradients, params):
1148 '''

Callers

nothing calls this directly

Calls 1

mask_gradsMethod · 0.95

Tested by

no test coverage detected