| 21 | |
| 22 | |
| 23 | class CenterPadding(nn.Module): |
| 24 | def __init__(self, multiple): |
| 25 | super().__init__() |
| 26 | self.multiple = multiple |
| 27 | |
| 28 | def _get_pad(self, size): |
| 29 | new_size = math.ceil(size / self.multiple) * self.multiple |
| 30 | pad_size = new_size - size |
| 31 | pad_size_left = pad_size // 2 |
| 32 | pad_size_right = pad_size - pad_size_left |
| 33 | return pad_size_left, pad_size_right |
| 34 | |
| 35 | @torch.inference_mode() |
| 36 | def forward(self, x): |
| 37 | pads = list(itertools.chain.from_iterable(self._get_pad(m) for m in x.shape[:1:-1])) |
| 38 | output = F.pad(x, pads) |
| 39 | return output |
no outgoing calls
no test coverage detected