| 35 | |
| 36 | |
| 37 | class AvgDown3D(nn.Module): |
| 38 | def __init__( |
| 39 | self, |
| 40 | in_channels, |
| 41 | out_channels, |
| 42 | factor_t, |
| 43 | factor_s=1, |
| 44 | ): |
| 45 | super().__init__() |
| 46 | self.in_channels = in_channels |
| 47 | self.out_channels = out_channels |
| 48 | self.factor_t = factor_t |
| 49 | self.factor_s = factor_s |
| 50 | self.factor = self.factor_t * self.factor_s * self.factor_s |
| 51 | |
| 52 | assert in_channels * self.factor % out_channels == 0 |
| 53 | self.group_size = in_channels * self.factor // out_channels |
| 54 | |
| 55 | def forward(self, x: torch.Tensor) -> torch.Tensor: |
| 56 | pad_t = (self.factor_t - x.shape[2] % self.factor_t) % self.factor_t |
| 57 | pad = (0, 0, 0, 0, pad_t, 0) |
| 58 | x = F.pad(x, pad) |
| 59 | B, C, T, H, W = x.shape |
| 60 | x = x.view( |
| 61 | B, |
| 62 | C, |
| 63 | T // self.factor_t, |
| 64 | self.factor_t, |
| 65 | H // self.factor_s, |
| 66 | self.factor_s, |
| 67 | W // self.factor_s, |
| 68 | self.factor_s, |
| 69 | ) |
| 70 | x = x.permute(0, 1, 3, 5, 7, 2, 4, 6).contiguous() |
| 71 | x = x.view( |
| 72 | B, |
| 73 | C * self.factor, |
| 74 | T // self.factor_t, |
| 75 | H // self.factor_s, |
| 76 | W // self.factor_s, |
| 77 | ) |
| 78 | x = x.view( |
| 79 | B, |
| 80 | self.out_channels, |
| 81 | self.group_size, |
| 82 | T // self.factor_t, |
| 83 | H // self.factor_s, |
| 84 | W // self.factor_s, |
| 85 | ) |
| 86 | x = x.mean(dim=2) |
| 87 | return x |
| 88 | |
| 89 | |
| 90 | class DupUp3D(nn.Module): |