(self)
| 7 | |
| 8 | class SD3VAEDecoder(torch.nn.Module): |
| 9 | def __init__(self): |
| 10 | super().__init__() |
| 11 | self.scaling_factor = 1.5305 # Different from SD 1.x |
| 12 | self.shift_factor = 0.0609 # Different from SD 1.x |
| 13 | self.conv_in = torch.nn.Conv2d(16, 512, kernel_size=3, padding=1) # Different from SD 1.x |
| 14 | |
| 15 | self.blocks = torch.nn.ModuleList([ |
| 16 | # UNetMidBlock2D |
| 17 | ResnetBlock(512, 512, eps=1e-6), |
| 18 | VAEAttentionBlock(1, 512, 512, 1, eps=1e-6), |
| 19 | ResnetBlock(512, 512, eps=1e-6), |
| 20 | # UpDecoderBlock2D |
| 21 | ResnetBlock(512, 512, eps=1e-6), |
| 22 | ResnetBlock(512, 512, eps=1e-6), |
| 23 | ResnetBlock(512, 512, eps=1e-6), |
| 24 | UpSampler(512), |
| 25 | # UpDecoderBlock2D |
| 26 | ResnetBlock(512, 512, eps=1e-6), |
| 27 | ResnetBlock(512, 512, eps=1e-6), |
| 28 | ResnetBlock(512, 512, eps=1e-6), |
| 29 | UpSampler(512), |
| 30 | # UpDecoderBlock2D |
| 31 | ResnetBlock(512, 256, eps=1e-6), |
| 32 | ResnetBlock(256, 256, eps=1e-6), |
| 33 | ResnetBlock(256, 256, eps=1e-6), |
| 34 | UpSampler(256), |
| 35 | # UpDecoderBlock2D |
| 36 | ResnetBlock(256, 128, eps=1e-6), |
| 37 | ResnetBlock(128, 128, eps=1e-6), |
| 38 | ResnetBlock(128, 128, eps=1e-6), |
| 39 | ]) |
| 40 | |
| 41 | self.conv_norm_out = torch.nn.GroupNorm(num_channels=128, num_groups=32, eps=1e-6) |
| 42 | self.conv_act = torch.nn.SiLU() |
| 43 | self.conv_out = torch.nn.Conv2d(128, 3, kernel_size=3, padding=1) |
| 44 | |
| 45 | def tiled_forward(self, sample, tile_size=64, tile_stride=32): |
| 46 | hidden_states = TileWorker().tiled_forward( |
nothing calls this directly
no test coverage detected