Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
sd3_vae_decoder.py81 linesDownload Raw Back to models
1import torch2from .sd_vae_decoder import VAEAttentionBlock, SDVAEDecoderStateDictConverter3from .sd_unet import ResnetBlock, UpSampler4from .tiler import TileWorker5 6 7 8class SD3VAEDecoder(torch.nn.Module):9    def __init__(self):10        super().__init__()11        self.scaling_factor = 1.5305 # Different from SD 1.x12        self.shift_factor = 0.0609 # Different from SD 1.x13        self.conv_in = torch.nn.Conv2d(16, 512, kernel_size=3, padding=1) # Different from SD 1.x14 15        self.blocks = torch.nn.ModuleList([16            # UNetMidBlock2D17            ResnetBlock(512, 512, eps=1e-6),18            VAEAttentionBlock(1, 512, 512, 1, eps=1e-6),19            ResnetBlock(512, 512, eps=1e-6),20            # UpDecoderBlock2D21            ResnetBlock(512, 512, eps=1e-6),22            ResnetBlock(512, 512, eps=1e-6),23            ResnetBlock(512, 512, eps=1e-6),24            UpSampler(512),25            # UpDecoderBlock2D26            ResnetBlock(512, 512, eps=1e-6),27            ResnetBlock(512, 512, eps=1e-6),28            ResnetBlock(512, 512, eps=1e-6),29            UpSampler(512),30            # UpDecoderBlock2D31            ResnetBlock(512, 256, eps=1e-6),32            ResnetBlock(256, 256, eps=1e-6),33            ResnetBlock(256, 256, eps=1e-6),34            UpSampler(256),35            # UpDecoderBlock2D36            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(47            lambda x: self.forward(x),48            sample,49            tile_size,50            tile_stride,51            tile_device=sample.device,52            tile_dtype=sample.dtype53        )54        return hidden_states55 56    def forward(self, sample, tiled=False, tile_size=64, tile_stride=32, **kwargs):57        # For VAE Decoder, we do not need to apply the tiler on each layer.58        if tiled:59            return self.tiled_forward(sample, tile_size=tile_size, tile_stride=tile_stride)60 61        # 1. pre-process62        hidden_states = sample / self.scaling_factor + self.shift_factor63        hidden_states = self.conv_in(hidden_states)64        time_emb = None65        text_emb = None66        res_stack = None67 68        # 2. blocks69        for i, block in enumerate(self.blocks):70            hidden_states, time_emb, text_emb, res_stack = block(hidden_states, time_emb, text_emb, res_stack)71        72        # 3. output73        hidden_states = self.conv_norm_out(hidden_states)74        hidden_states = self.conv_act(hidden_states)75        hidden_states = self.conv_out(hidden_states)76 77        return hidden_states78    79    @staticmethod80    def state_dict_converter():81        return SDVAEDecoderStateDictConverter()