MCPcopy Create free account
hub / github.com/OpenImagingLab/FlashVSR / __init__

Method __init__

diffsynth/models/sd_vae_decoder.py:45–79  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

43
44class SDVAEDecoder(torch.nn.Module):
45 def __init__(self):
46 super().__init__()
47 self.scaling_factor = 0.18215
48 self.post_quant_conv = torch.nn.Conv2d(4, 4, kernel_size=1)
49 self.conv_in = torch.nn.Conv2d(4, 512, kernel_size=3, padding=1)
50
51 self.blocks = torch.nn.ModuleList([
52 # UNetMidBlock2D
53 ResnetBlock(512, 512, eps=1e-6),
54 VAEAttentionBlock(1, 512, 512, 1, eps=1e-6),
55 ResnetBlock(512, 512, eps=1e-6),
56 # UpDecoderBlock2D
57 ResnetBlock(512, 512, eps=1e-6),
58 ResnetBlock(512, 512, eps=1e-6),
59 ResnetBlock(512, 512, eps=1e-6),
60 UpSampler(512),
61 # UpDecoderBlock2D
62 ResnetBlock(512, 512, eps=1e-6),
63 ResnetBlock(512, 512, eps=1e-6),
64 ResnetBlock(512, 512, eps=1e-6),
65 UpSampler(512),
66 # UpDecoderBlock2D
67 ResnetBlock(512, 256, eps=1e-6),
68 ResnetBlock(256, 256, eps=1e-6),
69 ResnetBlock(256, 256, eps=1e-6),
70 UpSampler(256),
71 # UpDecoderBlock2D
72 ResnetBlock(256, 128, eps=1e-6),
73 ResnetBlock(128, 128, eps=1e-6),
74 ResnetBlock(128, 128, eps=1e-6),
75 ])
76
77 self.conv_norm_out = torch.nn.GroupNorm(num_channels=128, num_groups=32, eps=1e-5)
78 self.conv_act = torch.nn.SiLU()
79 self.conv_out = torch.nn.Conv2d(128, 3, kernel_size=3, padding=1)
80
81 def tiled_forward(self, sample, tile_size=64, tile_stride=32):
82 hidden_states = TileWorker().tiled_forward(

Callers 1

__init__Method · 0.45

Calls 3

ResnetBlockClass · 0.85
UpSamplerClass · 0.85
VAEAttentionBlockClass · 0.70

Tested by

no test coverage detected