(self, f: torch.Tensor, zq: torch.Tensor)
| 125 | self.conv_b = nn.Conv2d(zq_channels, f_channels, kernel_size=1, stride=1, padding=0) |
| 126 | |
| 127 | def forward(self, f: torch.Tensor, zq: torch.Tensor) -> torch.Tensor: |
| 128 | f_size = f.shape[-2:] |
| 129 | zq = F.interpolate(zq, size=f_size, mode="nearest") |
| 130 | norm_f = self.norm_layer(f) |
| 131 | new_f = norm_f * self.conv_y(zq) + self.conv_b(zq) |
| 132 | return new_f |
nothing calls this directly
no outgoing calls
no test coverage detected