| 34 | |
| 35 | |
| 36 | class 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 | |
| 66 | class SqueezeNet(BaseNet): |
nothing calls this directly
no outgoing calls
no test coverage detected