(self, in_channel, out_channel, kernel_size, style_dim, demodulate=True, upsample=False,
downsample=False, blur_kernel=[1, 3, 3, 1], )
| 190 | |
| 191 | class ModulatedConv2d(nn.Module): |
| 192 | def __init__(self, in_channel, out_channel, kernel_size, style_dim, demodulate=True, upsample=False, |
| 193 | downsample=False, blur_kernel=[1, 3, 3, 1], ): |
| 194 | super().__init__() |
| 195 | |
| 196 | self.eps = 1e-8 |
| 197 | self.kernel_size = kernel_size |
| 198 | self.in_channel = in_channel |
| 199 | self.out_channel = out_channel |
| 200 | self.upsample = upsample |
| 201 | self.downsample = downsample |
| 202 | |
| 203 | if upsample: |
| 204 | factor = 2 |
| 205 | p = (len(blur_kernel) - factor) - (kernel_size - 1) |
| 206 | pad0 = (p + 1) // 2 + factor - 1 |
| 207 | pad1 = p // 2 + 1 |
| 208 | |
| 209 | self.blur = Blur(blur_kernel, pad=(pad0, pad1), upsample_factor=factor) |
| 210 | |
| 211 | if downsample: |
| 212 | factor = 2 |
| 213 | p = (len(blur_kernel) - factor) + (kernel_size - 1) |
| 214 | pad0 = (p + 1) // 2 |
| 215 | pad1 = p // 2 |
| 216 | |
| 217 | self.blur = Blur(blur_kernel, pad=(pad0, pad1)) |
| 218 | |
| 219 | fan_in = in_channel * kernel_size ** 2 |
| 220 | self.scale = 1 / math.sqrt(fan_in) |
| 221 | self.padding = kernel_size // 2 |
| 222 | |
| 223 | self.weight = nn.Parameter(torch.randn(1, out_channel, in_channel, kernel_size, kernel_size)) |
| 224 | |
| 225 | self.modulation = EqualLinear(style_dim, in_channel, bias_init=1) |
| 226 | self.demodulate = demodulate |
| 227 | |
| 228 | def __repr__(self): |
| 229 | return ( |
nothing calls this directly
no test coverage detected