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

Method update

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

Source from the content-addressed store, hash-verified

367 self.register_buffer('update_count', torch.tensor([0]))
368
369 def update(self, minibatches, unlabeled=None):
370 if self.update_count >= self.hparams["vrex_penalty_anneal_iters"]:
371 penalty_weight = self.hparams["vrex_lambda"]
372 else:
373 penalty_weight = 1.0
374
375 nll = 0.
376
377 all_x = torch.cat([x for x, y in minibatches])
378 all_logits = self.network(all_x)
379 all_logits_idx = 0
380 losses = torch.zeros(len(minibatches))
381 for i, (x, y) in enumerate(minibatches):
382 logits = all_logits[all_logits_idx:all_logits_idx + x.shape[0]]
383 all_logits_idx += x.shape[0]
384 nll = F.cross_entropy(logits, y)
385 losses[i] = nll
386
387 mean = losses.mean()
388 penalty = ((losses - mean)**2).mean()
389 loss = mean + penalty_weight * penalty
390
391 if self.update_count == self.hparams['vrex_penalty_anneal_iters']:
392 # Reset Adam (like IRM), because it doesn't like the sharp jump in
393 # gradient magnitudes that happens at this step.
394 self.optimizer = torch.optim.Adam(
395 self.network.parameters(),
396 lr=self.hparams["lr"],
397 weight_decay=self.hparams['weight_decay']
398 )
399
400 self.optimizer.zero_grad()
401 loss.backward()
402 self.optimizer.step()
403
404 self.update_count += 1
405 return {'loss': loss.item(), 'nll': nll.item(), 'penalty': penalty.item()}
406
407
408class Mixup(ERM):

Callers

nothing calls this directly

Calls 1

meanMethod · 0.80

Tested by

no test coverage detected