| 792 | return self.network_s(self.randomize(self.network_f(x), "content")) |
| 793 | |
| 794 | def randomize(self, x, what="style", eps=1e-5): |
| 795 | device = "cuda" if x.is_cuda else "cpu" |
| 796 | sizes = x.size() |
| 797 | alpha = torch.rand(sizes[0], 1).to(device) |
| 798 | |
| 799 | if len(sizes) == 4: |
| 800 | x = x.view(sizes[0], sizes[1], -1) |
| 801 | alpha = alpha.unsqueeze(-1) |
| 802 | |
| 803 | mean = x.mean(-1, keepdim=True) |
| 804 | var = x.var(-1, keepdim=True) |
| 805 | |
| 806 | x = (x - mean) / (var + eps).sqrt() |
| 807 | |
| 808 | idx_swap = torch.randperm(sizes[0]) |
| 809 | if what == "style": |
| 810 | mean = alpha * mean + (1 - alpha) * mean[idx_swap] |
| 811 | var = alpha * var + (1 - alpha) * var[idx_swap] |
| 812 | else: |
| 813 | x = x[idx_swap].detach() |
| 814 | |
| 815 | x = x * (var + eps).sqrt() + mean |
| 816 | return x.view(*sizes) |
| 817 | |
| 818 | def update(self, minibatches, unlabeled=None): |
| 819 | all_x = torch.cat([x for x, y in minibatches]) |