Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
sd3_vae_encoder.py96 linesDownload Raw Back to root
1import torch2from .sd_unet import ResnetBlock, DownSampler3from .sd_vae_encoder import VAEAttentionBlock, SDVAEEncoderStateDictConverter4from .tiler import TileWorker5from einops import rearrange6 7 8class SD3VAEEncoder(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(3, 128, kernel_size=3, padding=1)14 15        self.blocks = torch.nn.ModuleList([16            # DownEncoderBlock2D17            ResnetBlock(128, 128, eps=1e-6),18            ResnetBlock(128, 128, eps=1e-6),19            DownSampler(128, padding=0, extra_padding=True),20            # DownEncoderBlock2D21            ResnetBlock(128, 256, eps=1e-6),22            ResnetBlock(256, 256, eps=1e-6),23            DownSampler(256, padding=0, extra_padding=True),24            # DownEncoderBlock2D25            ResnetBlock(256, 512, eps=1e-6),26            ResnetBlock(512, 512, eps=1e-6),27            DownSampler(512, padding=0, extra_padding=True),28            # DownEncoderBlock2D29            ResnetBlock(512, 512, eps=1e-6),30            ResnetBlock(512, 512, eps=1e-6),31            # UNetMidBlock2D32            ResnetBlock(512, 512, eps=1e-6),33            VAEAttentionBlock(1, 512, 512, 1, eps=1e-6),34            ResnetBlock(512, 512, eps=1e-6),35        ])36 37        self.conv_norm_out = torch.nn.GroupNorm(num_channels=512, num_groups=32, eps=1e-6)38        self.conv_act = torch.nn.SiLU()39        self.conv_out = torch.nn.Conv2d(512, 32, kernel_size=3, padding=1)40 41    def tiled_forward(self, sample, tile_size=64, tile_stride=32):42        hidden_states = TileWorker().tiled_forward(43            lambda x: self.forward(x),44            sample,45            tile_size,46            tile_stride,47            tile_device=sample.device,48            tile_dtype=sample.dtype49        )50        return hidden_states51 52    def forward(self, sample, tiled=False, tile_size=64, tile_stride=32, **kwargs):53        # For VAE Decoder, we do not need to apply the tiler on each layer.54        if tiled:55            return self.tiled_forward(sample, tile_size=tile_size, tile_stride=tile_stride)56        57        # 1. pre-process58        hidden_states = self.conv_in(sample)59        time_emb = None60        text_emb = None61        res_stack = None62 63        # 2. blocks64        for i, block in enumerate(self.blocks):65            hidden_states, time_emb, text_emb, res_stack = block(hidden_states, time_emb, text_emb, res_stack)66        67        # 3. output68        hidden_states = self.conv_norm_out(hidden_states)69        hidden_states = self.conv_act(hidden_states)70        hidden_states = self.conv_out(hidden_states)71        hidden_states = hidden_states[:, :16]72        hidden_states = (hidden_states - self.shift_factor) * self.scaling_factor73 74        return hidden_states75    76    def encode_video(self, sample, batch_size=8):77        B = sample.shape[0]78        hidden_states = []79 80        for i in range(0, sample.shape[2], batch_size):81 82            j = min(i + batch_size, sample.shape[2])83            sample_batch = rearrange(sample[:,:,i:j], "B C T H W -> (B T) C H W")84 85            hidden_states_batch = self(sample_batch)86            hidden_states_batch = rearrange(hidden_states_batch, "(B T) C H W -> B C T H W", B=B)87 88            hidden_states.append(hidden_states_batch)89        90        hidden_states = torch.concat(hidden_states, dim=2)91        return hidden_states92    93    @staticmethod94    def state_dict_converter():95        return SDVAEEncoderStateDictConverter()96