(self, in0, in1, retPerLayer=False)
| 63 | self.lins+=[self.lin5,self.lin6] |
| 64 | |
| 65 | def forward(self, in0, in1, retPerLayer=False): |
| 66 | # v0.0 - original release had a bug, where input was not scaled |
| 67 | in0_input, in1_input = (self.scaling_layer(in0), self.scaling_layer(in1)) if self.version=='0.1' else (in0, in1) |
| 68 | outs0, outs1 = self.net.forward(in0_input), self.net.forward(in1_input) |
| 69 | feats0, feats1, diffs = {}, {}, {} |
| 70 | |
| 71 | for kk in range(self.L): |
| 72 | feats0[kk], feats1[kk] = util.normalize_tensor(outs0[kk]), util.normalize_tensor(outs1[kk]) |
| 73 | diffs[kk] = (feats0[kk]-feats1[kk])**2 |
| 74 | |
| 75 | if(self.lpips): |
| 76 | if(self.spatial): |
| 77 | res = [upsample(self.lins[kk].model(diffs[kk]), out_H=in0.shape[2]) for kk in range(self.L)] |
| 78 | else: |
| 79 | res = [spatial_average(self.lins[kk].model(diffs[kk]), keepdim=True) for kk in range(self.L)] |
| 80 | else: |
| 81 | if(self.spatial): |
| 82 | res = [upsample(diffs[kk].sum(dim=1,keepdim=True), out_H=in0.shape[2]) for kk in range(self.L)] |
| 83 | else: |
| 84 | res = [spatial_average(diffs[kk].sum(dim=1,keepdim=True), keepdim=True) for kk in range(self.L)] |
| 85 | |
| 86 | val = res[0] |
| 87 | for l in range(1,self.L): |
| 88 | val += res[l] |
| 89 | |
| 90 | if(retPerLayer): |
| 91 | return (val, res) |
| 92 | else: |
| 93 | return val |
| 94 | |
| 95 | class ScalingLayer(nn.Module): |
| 96 | def __init__(self): |
nothing calls this directly
no test coverage detected