hugging-apps/echo-memory
0
1import torch2from einops import rearrange, repeat3from .tiler import TileWorker2Dto3D4 5 6 7class Downsample3D(torch.nn.Module):8 def __init__(9 self,10 in_channels: int,11 out_channels: int,12 kernel_size: int = 3,13 stride: int = 2,14 padding: int = 0,15 compress_time: bool = False,16 ):17 super().__init__()18 19 self.conv = torch.nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, stride=stride, padding=padding)20 self.compress_time = compress_time21 22 def forward(self, x: torch.Tensor, xq: torch.Tensor) -> torch.Tensor:23 if self.compress_time:24 batch_size, channels, frames, height, width = x.shape25 26 # (batch_size, channels, frames, height, width) -> (batch_size, height, width, channels, frames) -> (batch_size * height * width, channels, frames)27 x = x.permute(0, 3, 4, 1, 2).reshape(batch_size * height * width, channels, frames)28 29 if x.shape[-1] % 2 == 1:30 x_first, x_rest = x[..., 0], x[..., 1:]31 if x_rest.shape[-1] > 0:32 # (batch_size * height * width, channels, frames - 1) -> (batch_size * height * width, channels, (frames - 1) // 2)33 x_rest = torch.nn.functional.avg_pool1d(x_rest, kernel_size=2, stride=2)34 35 x = torch.cat([x_first[..., None], x_rest], dim=-1)36 # (batch_size * height * width, channels, (frames // 2) + 1) -> (batch_size, height, width, channels, (frames // 2) + 1) -> (batch_size, channels, (frames // 2) + 1, height, width)37 x = x.reshape(batch_size, height, width, channels, x.shape[-1]).permute(0, 3, 4, 1, 2)38 else:39 # (batch_size * height * width, channels, frames) -> (batch_size * height * width, channels, frames // 2)40 x = torch.nn.functional.avg_pool1d(x, kernel_size=2, stride=2)41 # (batch_size * height * width, channels, frames // 2) -> (batch_size, height, width, channels, frames // 2) -> (batch_size, channels, frames // 2, height, width)42 x = x.reshape(batch_size, height, width, channels, x.shape[-1]).permute(0, 3, 4, 1, 2)43 44 # Pad the tensor45 pad = (0, 1, 0, 1)46 x = torch.nn.functional.pad(x, pad, mode="constant", value=0)47 batch_size, channels, frames, height, width = x.shape48 # (batch_size, channels, frames, height, width) -> (batch_size, frames, channels, height, width) -> (batch_size * frames, channels, height, width)49 x = x.permute(0, 2, 1, 3, 4).reshape(batch_size * frames, channels, height, width)50 x = self.conv(x)51 # (batch_size * frames, channels, height, width) -> (batch_size, frames, channels, height, width) -> (batch_size, channels, frames, height, width)52 x = x.reshape(batch_size, frames, x.shape[1], x.shape[2], x.shape[3]).permute(0, 2, 1, 3, 4)53 return x54 55 56 57class Upsample3D(torch.nn.Module):58 def __init__(59 self,60 in_channels: int,61 out_channels: int,62 kernel_size: int = 3,63 stride: int = 1,64 padding: int = 1,65 compress_time: bool = False,66 ) -> None:67 super().__init__()68 self.conv = torch.nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, stride=stride, padding=padding)69 self.compress_time = compress_time70 71 def forward(self, inputs: torch.Tensor, xq: torch.Tensor) -> torch.Tensor:72 if self.compress_time:73 if inputs.shape[2] > 1 and inputs.shape[2] % 2 == 1:74 # split first frame75 x_first, x_rest = inputs[:, :, 0], inputs[:, :, 1:]76 77 x_first = torch.nn.functional.interpolate(x_first, scale_factor=2.0)78 x_rest = torch.nn.functional.interpolate(x_rest, scale_factor=2.0)79 x_first = x_first[:, :, None, :, :]80 inputs = torch.cat([x_first, x_rest], dim=2)81 elif inputs.shape[2] > 1:82 inputs = torch.nn.functional.interpolate(inputs, scale_factor=2.0)83 else:84 inputs = inputs.squeeze(2)85 inputs = torch.nn.functional.interpolate(inputs, scale_factor=2.0)86 inputs = inputs[:, :, None, :, :]87 else:88 # only interpolate 2D89 b, c, t, h, w = inputs.shape90 inputs = inputs.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w)91 inputs = torch.nn.functional.interpolate(inputs, scale_factor=2.0)92 inputs = inputs.reshape(b, t, c, *inputs.shape[2:]).permute(0, 2, 1, 3, 4)93 94 b, c, t, h, w = inputs.shape95 inputs = inputs.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w)96 inputs = self.conv(inputs)97 inputs = inputs.reshape(b, t, *inputs.shape[1:]).permute(0, 2, 1, 3, 4)98 99 return inputs100 101 102 103class CogVideoXSpatialNorm3D(torch.nn.Module):104 def __init__(self, f_channels, zq_channels, groups):105 super().__init__()106 self.norm_layer = torch.nn.GroupNorm(num_channels=f_channels, num_groups=groups, eps=1e-6, affine=True)107 self.conv_y = torch.nn.Conv3d(zq_channels, f_channels, kernel_size=1, stride=1)108 self.conv_b = torch.nn.Conv3d(zq_channels, f_channels, kernel_size=1, stride=1)109 110 111 def forward(self, f: torch.Tensor, zq: torch.Tensor) -> torch.Tensor:112 if f.shape[2] > 1 and f.shape[2] % 2 == 1:113 f_first, f_rest = f[:, :, :1], f[:, :, 1:]114 f_first_size, f_rest_size = f_first.shape[-3:], f_rest.shape[-3:]115 z_first, z_rest = zq[:, :, :1], zq[:, :, 1:]116 z_first = torch.nn.functional.interpolate(z_first, size=f_first_size)117 z_rest = torch.nn.functional.interpolate(z_rest, size=f_rest_size)118 zq = torch.cat([z_first, z_rest], dim=2)119 else:120 zq = torch.nn.functional.interpolate(zq, size=f.shape[-3:])121 122 norm_f = self.norm_layer(f)123 new_f = norm_f * self.conv_y(zq) + self.conv_b(zq)124 return new_f125 126 127 128class Resnet3DBlock(torch.nn.Module):129 def __init__(self, in_channels, out_channels, spatial_norm_dim, groups, eps=1e-6, use_conv_shortcut=False):130 super().__init__()131 self.nonlinearity = torch.nn.SiLU()132 if spatial_norm_dim is None:133 self.norm1 = torch.nn.GroupNorm(num_channels=in_channels, num_groups=groups, eps=eps)134 self.norm2 = torch.nn.GroupNorm(num_channels=out_channels, num_groups=groups, eps=eps)135 else:136 self.norm1 = CogVideoXSpatialNorm3D(in_channels, spatial_norm_dim, groups)137 self.norm2 = CogVideoXSpatialNorm3D(out_channels, spatial_norm_dim, groups)138 139 self.conv1 = CachedConv3d(in_channels, out_channels, kernel_size=3, padding=(0, 1, 1))140 141 self.conv2 = CachedConv3d(out_channels, out_channels, kernel_size=3, padding=(0, 1, 1))142 143 if in_channels != out_channels:144 if use_conv_shortcut:145 self.conv_shortcut = CachedConv3d(in_channels, out_channels, kernel_size=3, padding=(0, 1, 1))146 else:147 self.conv_shortcut = torch.nn.Conv3d(in_channels, out_channels, kernel_size=1)148 else:149 self.conv_shortcut = lambda x: x150 151 152 def forward(self, hidden_states, zq):153 residual = hidden_states154 155 hidden_states = self.norm1(hidden_states, zq) if isinstance(self.norm1, CogVideoXSpatialNorm3D) else self.norm1(hidden_states)156 hidden_states = self.nonlinearity(hidden_states)157 hidden_states = self.conv1(hidden_states)158 159 hidden_states = self.norm2(hidden_states, zq) if isinstance(self.norm2, CogVideoXSpatialNorm3D) else self.norm2(hidden_states)160 hidden_states = self.nonlinearity(hidden_states)161 hidden_states = self.conv2(hidden_states)162 163 hidden_states = hidden_states + self.conv_shortcut(residual)164 165 return hidden_states166 167 168 169class CachedConv3d(torch.nn.Conv3d):170 def __init__(self, in_channels, out_channels, kernel_size, stride=1, padding=0):171 super().__init__(in_channels, out_channels, kernel_size=kernel_size, stride=stride, padding=padding)172 self.cached_tensor = None173 174 175 def clear_cache(self):176 self.cached_tensor = None177 178 179 def forward(self, input: torch.Tensor, use_cache = True) -> torch.Tensor:180 if use_cache:181 if self.cached_tensor is None:182 self.cached_tensor = torch.concat([input[:, :, :1]] * 2, dim=2)183 input = torch.concat([self.cached_tensor, input], dim=2)184 self.cached_tensor = input[:, :, -2:]185 return super().forward(input)186 187 188 189class CogVAEDecoder(torch.nn.Module):190 def __init__(self):191 super().__init__()192 self.scaling_factor = 0.7193 self.conv_in = CachedConv3d(16, 512, kernel_size=3, stride=1, padding=(0, 1, 1))194 195 self.blocks = torch.nn.ModuleList([196 Resnet3DBlock(512, 512, 16, 32),197 Resnet3DBlock(512, 512, 16, 32),198 Resnet3DBlock(512, 512, 16, 32),199 Resnet3DBlock(512, 512, 16, 32),200 Resnet3DBlock(512, 512, 16, 32),201 Resnet3DBlock(512, 512, 16, 32),202 Upsample3D(512, 512, compress_time=True),203 Resnet3DBlock(512, 256, 16, 32),204 Resnet3DBlock(256, 256, 16, 32),205 Resnet3DBlock(256, 256, 16, 32),206 Resnet3DBlock(256, 256, 16, 32),207 Upsample3D(256, 256, compress_time=True),208 Resnet3DBlock(256, 256, 16, 32),209 Resnet3DBlock(256, 256, 16, 32),210 Resnet3DBlock(256, 256, 16, 32),211 Resnet3DBlock(256, 256, 16, 32),212 Upsample3D(256, 256, compress_time=False),213 Resnet3DBlock(256, 128, 16, 32),214 Resnet3DBlock(128, 128, 16, 32),215 Resnet3DBlock(128, 128, 16, 32),216 Resnet3DBlock(128, 128, 16, 32),217 ])218 219 self.norm_out = CogVideoXSpatialNorm3D(128, 16, 32)220 self.conv_act = torch.nn.SiLU()221 self.conv_out = CachedConv3d(128, 3, kernel_size=3, stride=1, padding=(0, 1, 1))222 223 224 def forward(self, sample):225 sample = sample / self.scaling_factor226 hidden_states = self.conv_in(sample)227 228 for block in self.blocks:229 hidden_states = block(hidden_states, sample)230 231 hidden_states = self.norm_out(hidden_states, sample)232 hidden_states = self.conv_act(hidden_states)233 hidden_states = self.conv_out(hidden_states)234 235 return hidden_states236 237 238 def decode_video(self, sample, tiled=True, tile_size=(60, 90), tile_stride=(30, 45), progress_bar=lambda x:x):239 if tiled:240 B, C, T, H, W = sample.shape241 return TileWorker2Dto3D().tiled_forward(242 forward_fn=lambda x: self.decode_small_video(x),243 model_input=sample,244 tile_size=tile_size, tile_stride=tile_stride,245 tile_device=sample.device, tile_dtype=sample.dtype,246 computation_device=sample.device, computation_dtype=sample.dtype,247 scales=(3/16, (T//2*8+T%2)/T, 8, 8),248 progress_bar=progress_bar249 )250 else:251 return self.decode_small_video(sample)252 253 254 def decode_small_video(self, sample):255 B, C, T, H, W = sample.shape256 computation_device = self.conv_in.weight.device257 computation_dtype = self.conv_in.weight.dtype258 value = []259 for i in range(T//2):260 tl = i*2 + T%2 - (T%2 and i==0)261 tr = i*2 + 2 + T%2262 model_input = sample[:, :, tl: tr, :, :].to(dtype=computation_dtype, device=computation_device)263 model_output = self.forward(model_input).to(dtype=sample.dtype, device=sample.device)264 value.append(model_output)265 value = torch.concat(value, dim=2)266 for name, module in self.named_modules():267 if isinstance(module, CachedConv3d):268 module.clear_cache()269 return value270 271 272 @staticmethod273 def state_dict_converter():274 return CogVAEDecoderStateDictConverter()275 276 277 278class CogVAEEncoder(torch.nn.Module):279 def __init__(self):280 super().__init__()281 self.scaling_factor = 0.7282 self.conv_in = CachedConv3d(3, 128, kernel_size=3, stride=1, padding=(0, 1, 1))283 284 self.blocks = torch.nn.ModuleList([285 Resnet3DBlock(128, 128, None, 32),286 Resnet3DBlock(128, 128, None, 32),287 Resnet3DBlock(128, 128, None, 32),288 Downsample3D(128, 128, compress_time=True),289 Resnet3DBlock(128, 256, None, 32),290 Resnet3DBlock(256, 256, None, 32),291 Resnet3DBlock(256, 256, None, 32),292 Downsample3D(256, 256, compress_time=True),293 Resnet3DBlock(256, 256, None, 32),294 Resnet3DBlock(256, 256, None, 32),295 Resnet3DBlock(256, 256, None, 32),296 Downsample3D(256, 256, compress_time=False),297 Resnet3DBlock(256, 512, None, 32),298 Resnet3DBlock(512, 512, None, 32),299 Resnet3DBlock(512, 512, None, 32),300 Resnet3DBlock(512, 512, None, 32),301 Resnet3DBlock(512, 512, None, 32),302 ])303 304 self.norm_out = torch.nn.GroupNorm(32, 512, eps=1e-06, affine=True)305 self.conv_act = torch.nn.SiLU()306 self.conv_out = CachedConv3d(512, 32, kernel_size=3, stride=1, padding=(0, 1, 1))307 308 309 def forward(self, sample):310 hidden_states = self.conv_in(sample)311 312 for block in self.blocks:313 hidden_states = block(hidden_states, sample)314 315 hidden_states = self.norm_out(hidden_states)316 hidden_states = self.conv_act(hidden_states)317 hidden_states = self.conv_out(hidden_states)[:, :16]318 hidden_states = hidden_states * self.scaling_factor319 320 return hidden_states321 322 323 def encode_video(self, sample, tiled=True, tile_size=(60, 90), tile_stride=(30, 45), progress_bar=lambda x:x):324 if tiled:325 B, C, T, H, W = sample.shape326 return TileWorker2Dto3D().tiled_forward(327 forward_fn=lambda x: self.encode_small_video(x),328 model_input=sample,329 tile_size=(i * 8 for i in tile_size), tile_stride=(i * 8 for i in tile_stride),330 tile_device=sample.device, tile_dtype=sample.dtype,331 computation_device=sample.device, computation_dtype=sample.dtype,332 scales=(16/3, (T//4+T%2)/T, 1/8, 1/8),333 progress_bar=progress_bar334 )335 else:336 return self.encode_small_video(sample)337 338 339 def encode_small_video(self, sample):340 B, C, T, H, W = sample.shape341 computation_device = self.conv_in.weight.device342 computation_dtype = self.conv_in.weight.dtype343 value = []344 for i in range(T//8):345 t = i*8 + T%2 - (T%2 and i==0)346 t_ = i*8 + 8 + T%2347 model_input = sample[:, :, t: t_, :, :].to(dtype=computation_dtype, device=computation_device)348 model_output = self.forward(model_input).to(dtype=sample.dtype, device=sample.device)349 value.append(model_output)350 value = torch.concat(value, dim=2)351 for name, module in self.named_modules():352 if isinstance(module, CachedConv3d):353 module.clear_cache()354 return value355 356 357 @staticmethod358 def state_dict_converter():359 return CogVAEEncoderStateDictConverter()360 361 362 363class CogVAEEncoderStateDictConverter:364 def __init__(self):365 pass366 367 368 def from_diffusers(self, state_dict):369 rename_dict = {370 "encoder.conv_in.conv.weight": "conv_in.weight",371 "encoder.conv_in.conv.bias": "conv_in.bias",372 "encoder.down_blocks.0.downsamplers.0.conv.weight": "blocks.3.conv.weight",373 "encoder.down_blocks.0.downsamplers.0.conv.bias": "blocks.3.conv.bias",374 "encoder.down_blocks.1.downsamplers.0.conv.weight": "blocks.7.conv.weight",375 "encoder.down_blocks.1.downsamplers.0.conv.bias": "blocks.7.conv.bias",376 "encoder.down_blocks.2.downsamplers.0.conv.weight": "blocks.11.conv.weight",377 "encoder.down_blocks.2.downsamplers.0.conv.bias": "blocks.11.conv.bias",378 "encoder.norm_out.weight": "norm_out.weight",379 "encoder.norm_out.bias": "norm_out.bias",380 "encoder.conv_out.conv.weight": "conv_out.weight",381 "encoder.conv_out.conv.bias": "conv_out.bias",382 }383 prefix_dict = {384 "encoder.down_blocks.0.resnets.0.": "blocks.0.",385 "encoder.down_blocks.0.resnets.1.": "blocks.1.",386 "encoder.down_blocks.0.resnets.2.": "blocks.2.",387 "encoder.down_blocks.1.resnets.0.": "blocks.4.",388 "encoder.down_blocks.1.resnets.1.": "blocks.5.",389 "encoder.down_blocks.1.resnets.2.": "blocks.6.",390 "encoder.down_blocks.2.resnets.0.": "blocks.8.",391 "encoder.down_blocks.2.resnets.1.": "blocks.9.",392 "encoder.down_blocks.2.resnets.2.": "blocks.10.",393 "encoder.down_blocks.3.resnets.0.": "blocks.12.",394 "encoder.down_blocks.3.resnets.1.": "blocks.13.",395 "encoder.down_blocks.3.resnets.2.": "blocks.14.",396 "encoder.mid_block.resnets.0.": "blocks.15.",397 "encoder.mid_block.resnets.1.": "blocks.16.",398 }399 suffix_dict = {400 "norm1.norm_layer.weight": "norm1.norm_layer.weight",401 "norm1.norm_layer.bias": "norm1.norm_layer.bias",402 "norm1.conv_y.conv.weight": "norm1.conv_y.weight",403 "norm1.conv_y.conv.bias": "norm1.conv_y.bias",404 "norm1.conv_b.conv.weight": "norm1.conv_b.weight",405 "norm1.conv_b.conv.bias": "norm1.conv_b.bias",406 "norm2.norm_layer.weight": "norm2.norm_layer.weight",407 "norm2.norm_layer.bias": "norm2.norm_layer.bias",408 "norm2.conv_y.conv.weight": "norm2.conv_y.weight",409 "norm2.conv_y.conv.bias": "norm2.conv_y.bias",410 "norm2.conv_b.conv.weight": "norm2.conv_b.weight",411 "norm2.conv_b.conv.bias": "norm2.conv_b.bias",412 "conv1.conv.weight": "conv1.weight",413 "conv1.conv.bias": "conv1.bias",414 "conv2.conv.weight": "conv2.weight",415 "conv2.conv.bias": "conv2.bias",416 "conv_shortcut.weight": "conv_shortcut.weight",417 "conv_shortcut.bias": "conv_shortcut.bias",418 "norm1.weight": "norm1.weight",419 "norm1.bias": "norm1.bias",420 "norm2.weight": "norm2.weight",421 "norm2.bias": "norm2.bias",422 }423 state_dict_ = {}424 for name, param in state_dict.items():425 if name in rename_dict:426 state_dict_[rename_dict[name]] = param427 else:428 for prefix in prefix_dict:429 if name.startswith(prefix):430 suffix = name[len(prefix):]431 state_dict_[prefix_dict[prefix] + suffix_dict[suffix]] = param432 return state_dict_433 434 435 def from_civitai(self, state_dict):436 return self.from_diffusers(state_dict)437 438 439 440class CogVAEDecoderStateDictConverter:441 def __init__(self):442 pass443 444 445 def from_diffusers(self, state_dict):446 rename_dict = {447 "decoder.conv_in.conv.weight": "conv_in.weight",448 "decoder.conv_in.conv.bias": "conv_in.bias",449 "decoder.up_blocks.0.upsamplers.0.conv.weight": "blocks.6.conv.weight",450 "decoder.up_blocks.0.upsamplers.0.conv.bias": "blocks.6.conv.bias",451 "decoder.up_blocks.1.upsamplers.0.conv.weight": "blocks.11.conv.weight",452 "decoder.up_blocks.1.upsamplers.0.conv.bias": "blocks.11.conv.bias",453 "decoder.up_blocks.2.upsamplers.0.conv.weight": "blocks.16.conv.weight",454 "decoder.up_blocks.2.upsamplers.0.conv.bias": "blocks.16.conv.bias",455 "decoder.norm_out.norm_layer.weight": "norm_out.norm_layer.weight",456 "decoder.norm_out.norm_layer.bias": "norm_out.norm_layer.bias",457 "decoder.norm_out.conv_y.conv.weight": "norm_out.conv_y.weight",458 "decoder.norm_out.conv_y.conv.bias": "norm_out.conv_y.bias",459 "decoder.norm_out.conv_b.conv.weight": "norm_out.conv_b.weight",460 "decoder.norm_out.conv_b.conv.bias": "norm_out.conv_b.bias",461 "decoder.conv_out.conv.weight": "conv_out.weight",462 "decoder.conv_out.conv.bias": "conv_out.bias"463 }464 prefix_dict = {465 "decoder.mid_block.resnets.0.": "blocks.0.",466 "decoder.mid_block.resnets.1.": "blocks.1.",467 "decoder.up_blocks.0.resnets.0.": "blocks.2.",468 "decoder.up_blocks.0.resnets.1.": "blocks.3.",469 "decoder.up_blocks.0.resnets.2.": "blocks.4.",470 "decoder.up_blocks.0.resnets.3.": "blocks.5.",471 "decoder.up_blocks.1.resnets.0.": "blocks.7.",472 "decoder.up_blocks.1.resnets.1.": "blocks.8.",473 "decoder.up_blocks.1.resnets.2.": "blocks.9.",474 "decoder.up_blocks.1.resnets.3.": "blocks.10.",475 "decoder.up_blocks.2.resnets.0.": "blocks.12.",476 "decoder.up_blocks.2.resnets.1.": "blocks.13.",477 "decoder.up_blocks.2.resnets.2.": "blocks.14.",478 "decoder.up_blocks.2.resnets.3.": "blocks.15.",479 "decoder.up_blocks.3.resnets.0.": "blocks.17.",480 "decoder.up_blocks.3.resnets.1.": "blocks.18.",481 "decoder.up_blocks.3.resnets.2.": "blocks.19.",482 "decoder.up_blocks.3.resnets.3.": "blocks.20.",483 }484 suffix_dict = {485 "norm1.norm_layer.weight": "norm1.norm_layer.weight",486 "norm1.norm_layer.bias": "norm1.norm_layer.bias",487 "norm1.conv_y.conv.weight": "norm1.conv_y.weight",488 "norm1.conv_y.conv.bias": "norm1.conv_y.bias",489 "norm1.conv_b.conv.weight": "norm1.conv_b.weight",490 "norm1.conv_b.conv.bias": "norm1.conv_b.bias",491 "norm2.norm_layer.weight": "norm2.norm_layer.weight",492 "norm2.norm_layer.bias": "norm2.norm_layer.bias",493 "norm2.conv_y.conv.weight": "norm2.conv_y.weight",494 "norm2.conv_y.conv.bias": "norm2.conv_y.bias",495 "norm2.conv_b.conv.weight": "norm2.conv_b.weight",496 "norm2.conv_b.conv.bias": "norm2.conv_b.bias",497 "conv1.conv.weight": "conv1.weight",498 "conv1.conv.bias": "conv1.bias",499 "conv2.conv.weight": "conv2.weight",500 "conv2.conv.bias": "conv2.bias",501 "conv_shortcut.weight": "conv_shortcut.weight",502 "conv_shortcut.bias": "conv_shortcut.bias",503 }504 state_dict_ = {}505 for name, param in state_dict.items():506 if name in rename_dict:507 state_dict_[rename_dict[name]] = param508 else:509 for prefix in prefix_dict:510 if name.startswith(prefix):511 suffix = name[len(prefix):]512 state_dict_[prefix_dict[prefix] + suffix_dict[suffix]] = param513 return state_dict_514 515 516 def from_civitai(self, state_dict):517 return self.from_diffusers(state_dict)518 519 