(self, x)
| 82 | self.weight.data.mul_(0.5) |
| 83 | |
| 84 | def forward(self, x): |
| 85 | B, T, *spatial_dims = x.shape |
| 86 | out = super().forward(x.reshape(B * T, *spatial_dims)) |
| 87 | BT, *spatial_dims = out.shape |
| 88 | out = out.view(B, T, *spatial_dims).contiguous() |
| 89 | return out |
| 90 | |
| 91 | |
| 92 | # x = torch.randn(1, 2, 3, 4, 5) |