hugging-apps/echo-memory
0
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 