| 1680 | """ |
| 1681 | |
| 1682 | def __init__( |
| 1683 | self, |
| 1684 | f_channels: int, |
| 1685 | zq_channels: int, |
| 1686 | ): |
| 1687 | super().__init__() |
| 1688 | self.norm_layer = nn.GroupNorm(num_channels=f_channels, num_groups=32, eps=1e-6, affine=True) |
| 1689 | self.conv_y = nn.Conv2d(zq_channels, f_channels, kernel_size=1, stride=1, padding=0) |
| 1690 | self.conv_b = nn.Conv2d(zq_channels, f_channels, kernel_size=1, stride=1, padding=0) |
| 1691 | |
| 1692 | def forward(self, f: torch.FloatTensor, zq: torch.FloatTensor) -> torch.FloatTensor: |
| 1693 | f_size = f.shape[-2:] |