| 4 | |
| 5 | class SimpleAdapter(nn.Module): |
| 6 | def __init__(self, in_dim, out_dim, kernel_size, stride, downscale_factor=8, num_residual_blocks=1): |
| 7 | super(SimpleAdapter, self).__init__() |
| 8 | |
| 9 | # Pixel Unshuffle: reduce spatial dimensions by a factor of 8 |
| 10 | self.pixel_unshuffle = nn.PixelUnshuffle(downscale_factor=downscale_factor) |
| 11 | |
| 12 | # Convolution: reduce spatial dimensions by a factor |
| 13 | # of 2 (without overlap) |
| 14 | self.conv = nn.Conv2d(in_dim * downscale_factor * downscale_factor, out_dim, kernel_size=kernel_size, stride=stride, padding=0) |
| 15 | |
| 16 | # Residual blocks for feature extraction |
| 17 | self.residual_blocks = nn.Sequential( |
| 18 | *[ResidualBlock(out_dim) for _ in range(num_residual_blocks)] |
| 19 | ) |
| 20 | |
| 21 | def forward(self, x): |
| 22 | # Reshape to merge the frame dimension into batch |