| 510 | |
| 511 | |
| 512 | class DiscriminatorBlock(Module): |
| 513 | def __init__(self, input_channels, filters, downsample=True, antialiased_downsample=True): |
| 514 | super().__init__() |
| 515 | self.conv_res = nn.Conv2d(input_channels, filters, 1, stride=(2 if downsample else 1)) |
| 516 | |
| 517 | self.net = nn.Sequential( |
| 518 | nn.Conv2d(input_channels, filters, 3, padding=1), |
| 519 | leaky_relu(), |
| 520 | nn.Conv2d(filters, filters, 3, padding=1), |
| 521 | leaky_relu(), |
| 522 | ) |
| 523 | |
| 524 | self.maybe_blur = Blur() if antialiased_downsample else None |
| 525 | |
| 526 | self.downsample = ( |
| 527 | nn.Sequential( |
| 528 | Rearrange("b c (h p1) (w p2) -> b (c p1 p2) h w", p1=2, p2=2), nn.Conv2d(filters * 4, filters, 1) |
| 529 | ) |
| 530 | if downsample |
| 531 | else None |
| 532 | ) |
| 533 | |
| 534 | def forward(self, x): |
| 535 | res = self.conv_res(x) |
| 536 | |
| 537 | x = self.net(x) |
| 538 | |
| 539 | if exists(self.downsample): |
| 540 | if exists(self.maybe_blur): |
| 541 | x = self.maybe_blur(x, space_only=True) |
| 542 | |
| 543 | x = self.downsample(x) |
| 544 | |
| 545 | x = (x + res) * (2**-0.5) |
| 546 | return x |
| 547 | |
| 548 | |
| 549 | class Discriminator(Module): |