| 68 | return input * torch.rsqrt(torch.mean(input ** 2, dim=2, keepdim=True) + 1e-8) |
| 69 | |
| 70 | class Upsample(nn.Module): |
| 71 | def __init__(self, kernel, factor=2): |
| 72 | super().__init__() |
| 73 | |
| 74 | self.factor = factor |
| 75 | kernel = make_kernel(kernel) * (factor ** 2) |
| 76 | self.register_buffer('kernel', kernel) |
| 77 | |
| 78 | p = kernel.shape[0] - factor |
| 79 | |
| 80 | pad0 = (p + 1) // 2 + factor - 1 |
| 81 | pad1 = p // 2 |
| 82 | |
| 83 | self.pad = (pad0, pad1) |
| 84 | |
| 85 | def forward(self, input): |
| 86 | return upfirdn2d(input, self.kernel, up=self.factor, down=1, pad=self.pad) |
| 87 | |
| 88 | |
| 89 | class Downsample(nn.Module): |