(
self,
in_channels: int,
out_channels: int,
factor_t,
factor_s=1,
)
| 89 | |
| 90 | class DupUp3D(nn.Module): |
| 91 | def __init__( |
| 92 | self, |
| 93 | in_channels: int, |
| 94 | out_channels: int, |
| 95 | factor_t, |
| 96 | factor_s=1, |
| 97 | ): |
| 98 | super().__init__() |
| 99 | self.in_channels = in_channels |
| 100 | self.out_channels = out_channels |
| 101 | |
| 102 | self.factor_t = factor_t |
| 103 | self.factor_s = factor_s |
| 104 | self.factor = self.factor_t * self.factor_s * self.factor_s |
| 105 | |
| 106 | assert out_channels * self.factor % in_channels == 0 |
| 107 | self.repeats = out_channels * self.factor // in_channels |
| 108 | |
| 109 | def forward(self, x: torch.Tensor, first_chunk=False) -> torch.Tensor: |
| 110 | x = x.repeat_interleave(self.repeats, dim=1) |
no outgoing calls
no test coverage detected