r"""A 2D K-upsampling layer. Parameters: pad_mode (`str`, *optional*, default to `"reflect"`): the padding mode to use.
| 325 | |
| 326 | |
| 327 | class KUpsample2D(nn.Module): |
| 328 | r"""A 2D K-upsampling layer. |
| 329 | |
| 330 | Parameters: |
| 331 | pad_mode (`str`, *optional*, default to `"reflect"`): the padding mode to use. |
| 332 | """ |
| 333 | |
| 334 | def __init__(self, pad_mode: str = "reflect"): |
| 335 | super().__init__() |
| 336 | self.pad_mode = pad_mode |
| 337 | kernel_1d = torch.tensor([[1 / 8, 3 / 8, 3 / 8, 1 / 8]]) * 2 |
| 338 | self.pad = kernel_1d.shape[1] // 2 - 1 |
| 339 | self.register_buffer("kernel", kernel_1d.T @ kernel_1d, persistent=False) |
| 340 | |
| 341 | def forward(self, inputs: torch.Tensor) -> torch.Tensor: |
| 342 | inputs = F.pad(inputs, ((self.pad + 1) // 2,) * 4, self.pad_mode) |
| 343 | weight = inputs.new_zeros( |
| 344 | [ |
| 345 | inputs.shape[1], |
| 346 | inputs.shape[1], |
| 347 | self.kernel.shape[0], |
| 348 | self.kernel.shape[1], |
| 349 | ] |
| 350 | ) |
| 351 | indices = torch.arange(inputs.shape[1], device=inputs.device) |
| 352 | kernel = self.kernel.to(weight)[None, :].expand(inputs.shape[1], -1, -1) |
| 353 | weight[indices, indices] = kernel |
| 354 | return F.conv_transpose2d(inputs, weight, stride=2, padding=self.pad * 2 + 1) |
| 355 | |
| 356 | |
| 357 | class CogVideoXUpsample3D(nn.Module): |