MCPcopy Create free account
hub / github.com/OpenImagingLab/4DSloMo / BaseNet

Class BaseNet

lpipsPyTorch/modules/networks.py:36–63  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

34
35
36class BaseNet(nn.Module):
37 def __init__(self):
38 super(BaseNet, self).__init__()
39
40 # register buffer
41 self.register_buffer(
42 'mean', torch.Tensor([-.030, -.088, -.188])[None, :, None, None])
43 self.register_buffer(
44 'std', torch.Tensor([.458, .448, .450])[None, :, None, None])
45
46 def set_requires_grad(self, state: bool):
47 for param in chain(self.parameters(), self.buffers()):
48 param.requires_grad = state
49
50 def z_score(self, x: torch.Tensor):
51 return (x - self.mean) / self.std
52
53 def forward(self, x: torch.Tensor):
54 x = self.z_score(x)
55
56 output = []
57 for i, (_, layer) in enumerate(self.layers._modules.items(), 1):
58 x = layer(x)
59 if i in self.target_layers:
60 output.append(normalize_activation(x))
61 if len(output) == len(self.target_layers):
62 break
63 return output
64
65
66class SqueezeNet(BaseNet):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected