| 33 | |
| 34 | class SpatialNorm(nn.Module): |
| 35 | def __init__( |
| 36 | self, |
| 37 | f_channels, |
| 38 | zq_channels, |
| 39 | norm_layer=nn.GroupNorm, |
| 40 | freeze_norm_layer=False, |
| 41 | add_conv=False, |
| 42 | **norm_layer_params, |
| 43 | ): |
| 44 | super().__init__() |
| 45 | self.norm_layer = norm_layer(num_channels=f_channels, **norm_layer_params) |
| 46 | if freeze_norm_layer: |
| 47 | for p in self.norm_layer.parameters: |
| 48 | p.requires_grad = False |
| 49 | self.add_conv = add_conv |
| 50 | if self.add_conv: |
| 51 | self.conv = nn.Conv2d(zq_channels, zq_channels, kernel_size=3, stride=1, padding=1) |
| 52 | self.conv_y = nn.Conv2d(zq_channels, f_channels, kernel_size=1, stride=1, padding=0) |
| 53 | self.conv_b = nn.Conv2d(zq_channels, f_channels, kernel_size=1, stride=1, padding=0) |
| 54 | |
| 55 | def forward(self, f, zq): |
| 56 | f_size = f.shape[-2:] |