MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / DiscriminatorBlock

Class DiscriminatorBlock

sat/sgm/modules/autoencoding/magvit2_pytorch.py:512–546  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

510
511
512class 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
549class Discriminator(Module):

Callers 1

__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected