(self)
| 43 | |
| 44 | class 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( |
no test coverage detected