| 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)) |