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

Method randomize

domainbed/algorithms.py:794–816  ·  view source on GitHub ↗
(self, x, what="style", eps=1e-5)

Source from the content-addressed store, hash-verified

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

Callers 2

forward_cMethod · 0.95
forward_sMethod · 0.95

Calls 1

meanMethod · 0.80

Tested by

no test coverage detected