Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
hunyuan_video_vae_encoder.py308 linesDownload Raw Back to models
1import torch2import torch.nn as nn3import torch.nn.functional as F4from einops import rearrange, repeat5import numpy as np6from tqdm import tqdm7from .hunyuan_video_vae_decoder import CausalConv3d, ResnetBlockCausal3D, UNetMidBlockCausal3D8 9 10class DownsampleCausal3D(nn.Module):11 12    def __init__(self, channels, out_channels, kernel_size=3, bias=True, stride=2):13        super().__init__()14        self.conv = CausalConv3d(channels, out_channels, kernel_size, stride=stride, bias=bias)15 16    def forward(self, hidden_states):17        hidden_states = self.conv(hidden_states)18        return hidden_states19 20 21class DownEncoderBlockCausal3D(nn.Module):22 23    def __init__(24        self,25        in_channels,26        out_channels,27        dropout=0.0,28        num_layers=1,29        eps=1e-6,30        num_groups=32,31        add_downsample=True,32        downsample_stride=2,33    ):34 35        super().__init__()36        resnets = []37        for i in range(num_layers):38            cur_in_channel = in_channels if i == 0 else out_channels39            resnets.append(40                ResnetBlockCausal3D(41                    in_channels=cur_in_channel,42                    out_channels=out_channels,43                    groups=num_groups,44                    dropout=dropout,45                    eps=eps,46                ))47        self.resnets = nn.ModuleList(resnets)48 49        self.downsamplers = None50        if add_downsample:51            self.downsamplers = nn.ModuleList([DownsampleCausal3D(52                out_channels,53                out_channels,54                stride=downsample_stride,55            )])56 57    def forward(self, hidden_states):58        for resnet in self.resnets:59            hidden_states = resnet(hidden_states)60 61        if self.downsamplers is not None:62            for downsampler in self.downsamplers:63                hidden_states = downsampler(hidden_states)64 65        return hidden_states66 67 68class EncoderCausal3D(nn.Module):69 70    def __init__(71        self,72        in_channels: int = 3,73        out_channels: int = 16,74        eps=1e-6,75        dropout=0.0,76        block_out_channels=[128, 256, 512, 512],77        layers_per_block=2,78        num_groups=32,79        time_compression_ratio: int = 4,80        spatial_compression_ratio: int = 8,81        gradient_checkpointing=False,82    ):83        super().__init__()84        self.conv_in = CausalConv3d(in_channels, block_out_channels[0], kernel_size=3, stride=1)85        self.down_blocks = nn.ModuleList([])86 87        # down88        output_channel = block_out_channels[0]89        for i in range(len(block_out_channels)):90            input_channel = output_channel91            output_channel = block_out_channels[i]92            is_final_block = i == len(block_out_channels) - 193            num_spatial_downsample_layers = int(np.log2(spatial_compression_ratio))94            num_time_downsample_layers = int(np.log2(time_compression_ratio))95 96            add_spatial_downsample = bool(i < num_spatial_downsample_layers)97            add_time_downsample = bool(i >= (len(block_out_channels) - 1 - num_time_downsample_layers) and not is_final_block)98 99            downsample_stride_HW = (2, 2) if add_spatial_downsample else (1, 1)100            downsample_stride_T = (2,) if add_time_downsample else (1,)101            downsample_stride = tuple(downsample_stride_T + downsample_stride_HW)102            down_block = DownEncoderBlockCausal3D(103                in_channels=input_channel,104                out_channels=output_channel,105                dropout=dropout,106                num_layers=layers_per_block,107                eps=eps,108                num_groups=num_groups,109                add_downsample=bool(add_spatial_downsample or add_time_downsample),110                downsample_stride=downsample_stride,111            )112            self.down_blocks.append(down_block)113 114        # mid115        self.mid_block = UNetMidBlockCausal3D(116            in_channels=block_out_channels[-1],117            dropout=dropout,118            eps=eps,119            num_groups=num_groups,120            attention_head_dim=block_out_channels[-1],121        )122        # out123        self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[-1], num_groups=num_groups, eps=eps)124        self.conv_act = nn.SiLU()125        self.conv_out = CausalConv3d(block_out_channels[-1], 2 * out_channels, kernel_size=3)126 127        self.gradient_checkpointing = gradient_checkpointing128 129    def forward(self, hidden_states):130        hidden_states = self.conv_in(hidden_states)131        if self.training and self.gradient_checkpointing:132 133            def create_custom_forward(module):134 135                def custom_forward(*inputs):136                    return module(*inputs)137 138                return custom_forward139 140            # down141            for down_block in self.down_blocks:142                torch.utils.checkpoint.checkpoint(143                    create_custom_forward(down_block),144                    hidden_states,145                    use_reentrant=False,146                )147            # middle148            hidden_states = torch.utils.checkpoint.checkpoint(149                create_custom_forward(self.mid_block),150                hidden_states,151                use_reentrant=False,152            )153        else:154            # down155            for down_block in self.down_blocks:156                hidden_states = down_block(hidden_states)157            # middle158            hidden_states = self.mid_block(hidden_states)159        # post-process160        hidden_states = self.conv_norm_out(hidden_states)161        hidden_states = self.conv_act(hidden_states)162        hidden_states = self.conv_out(hidden_states)163 164        return hidden_states165 166 167class HunyuanVideoVAEEncoder(nn.Module):168 169    def __init__(170        self,171        in_channels=3,172        out_channels=16,173        eps=1e-6,174        dropout=0.0,175        block_out_channels=[128, 256, 512, 512],176        layers_per_block=2,177        num_groups=32,178        time_compression_ratio=4,179        spatial_compression_ratio=8,180        gradient_checkpointing=False,181    ):182        super().__init__()183        self.encoder = EncoderCausal3D(184            in_channels=in_channels,185            out_channels=out_channels,186            eps=eps,187            dropout=dropout,188            block_out_channels=block_out_channels,189            layers_per_block=layers_per_block,190            num_groups=num_groups,191            time_compression_ratio=time_compression_ratio,192            spatial_compression_ratio=spatial_compression_ratio,193            gradient_checkpointing=gradient_checkpointing,194        )195        self.quant_conv = nn.Conv3d(2 * out_channels, 2 * out_channels, kernel_size=1)196        self.scaling_factor = 0.476986197 198 199    def forward(self, images):200        latents = self.encoder(images)201        latents = self.quant_conv(latents)202        latents = latents[:, :16]203        latents = latents * self.scaling_factor204        return latents205    206 207    def build_1d_mask(self, length, left_bound, right_bound, border_width):208        x = torch.ones((length,))209        if not left_bound:210            x[:border_width] = (torch.arange(border_width) + 1) / border_width211        if not right_bound:212            x[-border_width:] = torch.flip((torch.arange(border_width) + 1) / border_width, dims=(0,))213        return x214    215 216    def build_mask(self, data, is_bound, border_width):217        _, _, T, H, W = data.shape218        t = self.build_1d_mask(T, is_bound[0], is_bound[1], border_width[0])219        h = self.build_1d_mask(H, is_bound[2], is_bound[3], border_width[1])220        w = self.build_1d_mask(W, is_bound[4], is_bound[5], border_width[2])221 222        t = repeat(t, "T -> T H W", T=T, H=H, W=W)223        h = repeat(h, "H -> T H W", T=T, H=H, W=W)224        w = repeat(w, "W -> T H W", T=T, H=H, W=W)225 226        mask = torch.stack([t, h, w]).min(dim=0).values227        mask = rearrange(mask, "T H W -> 1 1 T H W")228        return mask229    230 231    def tile_forward(self, hidden_states, tile_size, tile_stride):232        B, C, T, H, W = hidden_states.shape233        size_t, size_h, size_w = tile_size234        stride_t, stride_h, stride_w = tile_stride235 236        # Split tasks237        tasks = []238        for t in range(0, T, stride_t):239            if (t-stride_t >= 0 and t-stride_t+size_t >= T): continue240            for h in range(0, H, stride_h):241                if (h-stride_h >= 0 and h-stride_h+size_h >= H): continue242                for w in range(0, W, stride_w):243                    if (w-stride_w >= 0 and w-stride_w+size_w >= W): continue244                    t_, h_, w_ = t + size_t, h + size_h, w + size_w245                    tasks.append((t, t_, h, h_, w, w_))246 247        # Run248        torch_dtype = self.quant_conv.weight.dtype249        data_device = hidden_states.device250        computation_device = self.quant_conv.weight.device251 252        weight = torch.zeros((1, 1,  (T - 1) // 4 + 1, H // 8, W // 8), dtype=torch_dtype, device=data_device)253        values = torch.zeros((B, 16, (T - 1) // 4 + 1, H // 8, W // 8), dtype=torch_dtype, device=data_device)254 255        for t, t_, h, h_, w, w_ in tqdm(tasks, desc="VAE encoding"):256            hidden_states_batch = hidden_states[:, :, t:t_, h:h_, w:w_].to(computation_device)257            hidden_states_batch = self.forward(hidden_states_batch).to(data_device)258            if t > 0:259                hidden_states_batch = hidden_states_batch[:, :, 1:]260 261            mask = self.build_mask(262                hidden_states_batch,263                is_bound=(t==0, t_>=T, h==0, h_>=H, w==0, w_>=W),264                border_width=((size_t - stride_t) // 4, (size_h - stride_h) // 8, (size_w - stride_w) // 8)265            ).to(dtype=torch_dtype, device=data_device)266 267            target_t = 0 if t==0 else t // 4 + 1268            target_h = h // 8269            target_w = w // 8270            values[271                :,272                :,273                target_t: target_t + hidden_states_batch.shape[2],274                target_h: target_h + hidden_states_batch.shape[3],275                target_w: target_w + hidden_states_batch.shape[4],276            ] += hidden_states_batch * mask277            weight[278                :,279                :,280                target_t: target_t + hidden_states_batch.shape[2],281                target_h: target_h + hidden_states_batch.shape[3],282                target_w: target_w + hidden_states_batch.shape[4],283            ] += mask284        return values / weight285 286 287    def encode_video(self, latents, tile_size=(65, 256, 256), tile_stride=(48, 192, 192)):288        latents = latents.to(self.quant_conv.weight.dtype)289        return self.tile_forward(latents, tile_size=tile_size, tile_stride=tile_stride)290 291 292    @staticmethod293    def state_dict_converter():294        return HunyuanVideoVAEEncoderStateDictConverter()295 296 297class HunyuanVideoVAEEncoderStateDictConverter:298 299    def __init__(self):300        pass301 302    def from_diffusers(self, state_dict):303        state_dict_ = {}304        for name in state_dict:305            if name.startswith('encoder.') or name.startswith('quant_conv.'):306                state_dict_[name] = state_dict[name]307        return state_dict_308