| 274 | |
| 275 | # Also ConvResblock |
| 276 | class Downsample(nn.Module): |
| 277 | def __init__(self, in_channels=320) -> None: |
| 278 | super().__init__() |
| 279 | self.f_t = nn.Linear(1280, in_channels * 2) |
| 280 | |
| 281 | self.gn_1 = nn.GroupNorm(32, in_channels) |
| 282 | self.f_1 = nn.Conv2d(in_channels, in_channels, kernel_size=3, padding=1) |
| 283 | self.gn_2 = nn.GroupNorm(32, in_channels) |
| 284 | |
| 285 | self.f_2 = nn.Conv2d(in_channels, in_channels, kernel_size=3, padding=1) |
| 286 | |
| 287 | def forward(self, x, t) -> torch.Tensor: |
| 288 | x_skip = x |
| 289 | |
| 290 | t = self.f_t(F.silu(t)) |
| 291 | t_1, t_2 = t.chunk(2, dim=1) |
| 292 | t_1 = t_1.unsqueeze(2).unsqueeze(3) + 1 |
| 293 | t_2 = t_2.unsqueeze(2).unsqueeze(3) |
| 294 | |
| 295 | gn_1 = F.silu(self.gn_1(x)) |
| 296 | avg_pool2d = F.avg_pool2d(gn_1, kernel_size=(2, 2), stride=None) |
| 297 | |
| 298 | f_1 = self.f_1(avg_pool2d) |
| 299 | gn_2 = self.gn_2(f_1) |
| 300 | |
| 301 | f_2 = self.f_2(F.silu(t_2 + (t_1 * gn_2))) |
| 302 | |
| 303 | return f_2 + F.avg_pool2d(x_skip, kernel_size=(2, 2), stride=None) |
| 304 | |
| 305 | |
| 306 | # Also ConvResblock |