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

Method update

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

Source from the content-addressed store, hash-verified

816 return x.view(*sizes)
817
818 def update(self, minibatches, unlabeled=None):
819 all_x = torch.cat([x for x, y in minibatches])
820 all_y = torch.cat([y for x, y in minibatches])
821
822 # learn content
823 self.optimizer_f.zero_grad()
824 self.optimizer_c.zero_grad()
825 loss_c = F.cross_entropy(self.forward_c(all_x), all_y)
826 loss_c.backward()
827 self.optimizer_f.step()
828 self.optimizer_c.step()
829
830 # learn style
831 self.optimizer_s.zero_grad()
832 loss_s = F.cross_entropy(self.forward_s(all_x), all_y)
833 loss_s.backward()
834 self.optimizer_s.step()
835
836 # learn adversary
837 self.optimizer_f.zero_grad()
838 loss_adv = -F.log_softmax(self.forward_s(all_x), dim=1).mean(1).mean()
839 loss_adv = loss_adv * self.weight_adv
840 loss_adv.backward()
841 self.optimizer_f.step()
842
843 return {'loss_c': loss_c.item(), 'loss_s': loss_s.item(), 'loss_adv': loss_adv.item()}
844
845 def predict(self, x):
846 return self.network_c(self.network_f(x))

Callers

nothing calls this directly

Calls 3

forward_cMethod · 0.95
forward_sMethod · 0.95
meanMethod · 0.80

Tested by

no test coverage detected