hugging-apps/echo-memory
0
1# Copyright 2025 StepFun Inc. All Rights Reserved.2# 3# Permission is hereby granted, free of charge, to any person obtaining a copy4# of this software and associated documentation files (the "Software"), to deal5# in the Software without restriction, including without limitation the rights6# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell7# copies of the Software, and to permit persons to whom the Software is8# furnished to do so, subject to the following conditions:9#10# The above copyright notice and this permission notice shall be included in all11# copies or substantial portions of the Software.12# ==============================================================================13import torch14from einops import rearrange15from torch import nn16from torch.nn import functional as F17from tqdm import tqdm18from einops import repeat19 20 21class BaseGroupNorm(nn.GroupNorm):22 def __init__(self, num_groups, num_channels):23 super().__init__(num_groups=num_groups, num_channels=num_channels)24 25 def forward(self, x, zero_pad=False, **kwargs):26 if zero_pad:27 return base_group_norm_with_zero_pad(x, self, **kwargs)28 else:29 return base_group_norm(x, self, **kwargs)30 31 32def base_group_norm(x, norm_layer, act_silu=False, channel_last=False):33 if hasattr(base_group_norm, 'spatial') and base_group_norm.spatial:34 assert channel_last == True35 x_shape = x.shape36 x = x.flatten(0, 1)37 if channel_last:38 # Permute to NCHW format39 x = x.permute(0, 3, 1, 2)40 41 out = F.group_norm(x.contiguous(), norm_layer.num_groups, norm_layer.weight, norm_layer.bias, norm_layer.eps)42 if act_silu:43 out = F.silu(out)44 45 if channel_last:46 # Permute back to NHWC format47 out = out.permute(0, 2, 3, 1)48 49 out = out.view(x_shape)50 else:51 if channel_last:52 # Permute to NCHW format53 x = x.permute(0, 3, 1, 2)54 out = F.group_norm(x.contiguous(), norm_layer.num_groups, norm_layer.weight, norm_layer.bias, norm_layer.eps)55 if act_silu:56 out = F.silu(out)57 if channel_last:58 # Permute back to NHWC format59 out = out.permute(0, 2, 3, 1)60 return out61 62def base_conv2d(x, conv_layer, channel_last=False, residual=None):63 if channel_last:64 x = x.permute(0, 3, 1, 2) # NHWC to NCHW65 out = F.conv2d(x, conv_layer.weight, conv_layer.bias, stride=conv_layer.stride, padding=conv_layer.padding)66 if residual is not None:67 if channel_last:68 residual = residual.permute(0, 3, 1, 2) # NHWC to NCHW69 out += residual70 if channel_last:71 out = out.permute(0, 2, 3, 1) # NCHW to NHWC72 return out73 74def base_conv3d(x, conv_layer, channel_last=False, residual=None, only_return_output=False):75 if only_return_output:76 size = cal_outsize(x.shape, conv_layer.weight.shape, conv_layer.stride, conv_layer.padding)77 return torch.empty(size, device=x.device, dtype=x.dtype)78 if channel_last:79 x = x.permute(0, 4, 1, 2, 3) # NDHWC to NCDHW80 out = F.conv3d(x, conv_layer.weight, conv_layer.bias, stride=conv_layer.stride, padding=conv_layer.padding)81 if residual is not None:82 if channel_last:83 residual = residual.permute(0, 4, 1, 2, 3) # NDHWC to NCDHW84 out += residual85 if channel_last:86 out = out.permute(0, 2, 3, 4, 1) # NCDHW to NDHWC87 return out88 89 90def cal_outsize(input_sizes, kernel_sizes, stride, padding):91 stride_d, stride_h, stride_w = stride92 padding_d, padding_h, padding_w = padding 93 dilation_d, dilation_h, dilation_w = 1, 1, 194 95 in_d = input_sizes[1]96 in_h = input_sizes[2]97 in_w = input_sizes[3]98 in_channel = input_sizes[4]99 100 101 kernel_d = kernel_sizes[2]102 kernel_h = kernel_sizes[3]103 kernel_w = kernel_sizes[4]104 out_channels = kernel_sizes[0]105 106 out_d = calc_out_(in_d, padding_d, dilation_d, kernel_d, stride_d)107 out_h = calc_out_(in_h, padding_h, dilation_h, kernel_h, stride_h)108 out_w = calc_out_(in_w, padding_w, dilation_w, kernel_w, stride_w)109 size = [input_sizes[0], out_d, out_h, out_w, out_channels]110 return size111 112 113 114 115def calc_out_(in_size, padding, dilation, kernel, stride):116 return (in_size + 2 * padding - dilation * (kernel - 1) - 1) // stride + 1117 118 119 120def base_conv3d_channel_last(x, conv_layer, residual=None):121 in_numel = x.numel()122 out_numel = int(x.numel() * conv_layer.out_channels / conv_layer.in_channels)123 if (in_numel >= 2**30) or (out_numel >= 2**30):124 assert conv_layer.stride[0] == 1, "time split asks time stride = 1"125 126 B,T,H,W,C = x.shape127 K = conv_layer.kernel_size[0]128 129 chunks = 4130 chunk_size = T // chunks131 132 if residual is None:133 out_nhwc = base_conv3d(x, conv_layer, channel_last=True, residual=residual, only_return_output=True)134 else:135 out_nhwc = residual136 137 assert B == 1138 outs = []139 for i in range(chunks):140 if i == chunks-1:141 xi = x[:1,chunk_size*i:]142 out_nhwci = out_nhwc[:1,chunk_size*i:]143 else:144 xi = x[:1,chunk_size*i:chunk_size*(i+1)+K-1]145 out_nhwci = out_nhwc[:1,chunk_size*i:chunk_size*(i+1)]146 if residual is not None:147 if i == chunks-1:148 ri = residual[:1,chunk_size*i:]149 else:150 ri = residual[:1,chunk_size*i:chunk_size*(i+1)]151 else:152 ri = None153 out_nhwci.copy_(base_conv3d(xi, conv_layer, channel_last=True, residual=ri))154 else:155 out_nhwc = base_conv3d(x, conv_layer, channel_last=True, residual=residual)156 return out_nhwc157 158 159 160class Upsample2D(nn.Module):161 def __init__(self,162 channels,163 use_conv=False,164 use_conv_transpose=False,165 out_channels=None):166 super().__init__()167 self.channels = channels168 self.out_channels = out_channels or channels169 self.use_conv = use_conv170 self.use_conv_transpose = use_conv_transpose171 172 if use_conv:173 self.conv = nn.Conv2d(self.channels, self.out_channels, 3, padding=1)174 else:175 assert "Not Supported"176 self.conv = nn.ConvTranspose2d(channels, self.out_channels, 4, 2, 1)177 178 def forward(self, x, output_size=None):179 assert x.shape[-1] == self.channels180 181 if self.use_conv_transpose:182 return self.conv(x)183 184 if output_size is None:185 x = F.interpolate(186 x.permute(0,3,1,2).to(memory_format=torch.channels_last),187 scale_factor=2.0, mode='nearest').permute(0,2,3,1).contiguous()188 else:189 x = F.interpolate(190 x.permute(0,3,1,2).to(memory_format=torch.channels_last),191 size=output_size, mode='nearest').permute(0,2,3,1).contiguous()192 193 # x = self.conv(x)194 x = base_conv2d(x, self.conv, channel_last=True)195 return x196 197 198class Downsample2D(nn.Module):199 def __init__(self, channels, use_conv=False, out_channels=None, padding=1):200 super().__init__()201 self.channels = channels202 self.out_channels = out_channels or channels203 self.use_conv = use_conv204 self.padding = padding205 stride = 2206 207 if use_conv:208 self.conv = nn.Conv2d(self.channels, self.out_channels, 3, stride=stride, padding=padding)209 else:210 assert self.channels == self.out_channels211 self.conv = nn.AvgPool2d(kernel_size=stride, stride=stride)212 213 def forward(self, x):214 assert x.shape[-1] == self.channels215 if self.use_conv and self.padding == 0:216 pad = (0, 0, 0, 1, 0, 1)217 x = F.pad(x, pad, mode="constant", value=0)218 219 assert x.shape[-1] == self.channels220 # x = self.conv(x)221 x = base_conv2d(x, self.conv, channel_last=True)222 return x223 224 225 226class CausalConv(nn.Module):227 def __init__(self,228 chan_in,229 chan_out,230 kernel_size,231 **kwargs232 ):233 super().__init__()234 235 if isinstance(kernel_size, int):236 kernel_size = kernel_size if isinstance(kernel_size, tuple) else ((kernel_size,) * 3)237 time_kernel_size, height_kernel_size, width_kernel_size = kernel_size238 239 self.dilation = kwargs.pop('dilation', 1)240 self.stride = kwargs.pop('stride', 1)241 if isinstance(self.stride, int):242 self.stride = (self.stride, 1, 1)243 time_pad = self.dilation * (time_kernel_size - 1) + max((1 - self.stride[0]), 0)244 height_pad = height_kernel_size // 2245 width_pad = width_kernel_size // 2246 self.time_causal_padding = (width_pad, width_pad, height_pad, height_pad, time_pad, 0)247 self.time_uncausal_padding = (width_pad, width_pad, height_pad, height_pad, 0, 0)248 249 self.conv = nn.Conv3d(chan_in, chan_out, kernel_size, stride=self.stride, dilation=self.dilation, **kwargs)250 self.is_first_run = True251 252 def forward(self, x, is_init=True, residual=None):253 x = nn.functional.pad(x,254 self.time_causal_padding if is_init else self.time_uncausal_padding)255 256 x = self.conv(x)257 if residual is not None:258 x.add_(residual)259 return x260 261 262class ChannelDuplicatingPixelUnshuffleUpSampleLayer3D(nn.Module):263 def __init__(264 self,265 in_channels: int,266 out_channels: int,267 factor: int,268 ):269 super().__init__()270 self.in_channels = in_channels271 self.out_channels = out_channels272 self.factor = factor273 assert out_channels * factor**3 % in_channels == 0274 self.repeats = out_channels * factor**3 // in_channels275 276 def forward(self, x: torch.Tensor, is_init=True) -> torch.Tensor:277 x = x.repeat_interleave(self.repeats, dim=1)278 x = x.view(x.size(0), self.out_channels, self.factor, self.factor, self.factor, x.size(2), x.size(3), x.size(4))279 x = x.permute(0, 1, 5, 2, 6, 3, 7, 4).contiguous()280 x = x.view(x.size(0), self.out_channels, x.size(2)*self.factor, x.size(4)*self.factor, x.size(6)*self.factor)281 x = x[:, :, self.factor - 1:, :, :]282 return x283 284class ConvPixelShuffleUpSampleLayer3D(nn.Module):285 def __init__(286 self,287 in_channels: int,288 out_channels: int,289 kernel_size: int,290 factor: int,291 ):292 super().__init__()293 self.factor = factor294 out_ratio = factor**3295 self.conv = CausalConv(296 in_channels,297 out_channels * out_ratio,298 kernel_size=kernel_size299 )300 301 def forward(self, x: torch.Tensor, is_init=True) -> torch.Tensor:302 x = self.conv(x, is_init)303 x = self.pixel_shuffle_3d(x, self.factor)304 return x305 306 @staticmethod307 def pixel_shuffle_3d(x: torch.Tensor, factor: int) -> torch.Tensor:308 batch_size, channels, depth, height, width = x.size()309 new_channels = channels // (factor ** 3)310 new_depth = depth * factor311 new_height = height * factor312 new_width = width * factor313 314 x = x.view(batch_size, new_channels, factor, factor, factor, depth, height, width)315 x = x.permute(0, 1, 5, 2, 6, 3, 7, 4).contiguous()316 x = x.view(batch_size, new_channels, new_depth, new_height, new_width)317 x = x[:, :, factor - 1:, :, :]318 return x319 320class ConvPixelUnshuffleDownSampleLayer3D(nn.Module):321 def __init__(322 self,323 in_channels: int,324 out_channels: int,325 kernel_size: int,326 factor: int,327 ):328 super().__init__()329 self.factor = factor330 out_ratio = factor**3331 assert out_channels % out_ratio == 0332 self.conv = CausalConv(333 in_channels,334 out_channels // out_ratio,335 kernel_size=kernel_size336 )337 338 def forward(self, x: torch.Tensor, is_init=True) -> torch.Tensor:339 x = self.conv(x, is_init)340 x = self.pixel_unshuffle_3d(x, self.factor)341 return x342 343 @staticmethod344 def pixel_unshuffle_3d(x: torch.Tensor, factor: int) -> torch.Tensor:345 pad = (0, 0, 0, 0, factor-1, 0) # (left, right, top, bottom, front, back)346 x = F.pad(x, pad)347 B, C, D, H, W = x.shape348 x = x.view(B, C, D // factor, factor, H // factor, factor, W // factor, factor)349 x = x.permute(0, 1, 3, 5, 7, 2, 4, 6).contiguous()350 x = x.view(B, C * factor**3, D // factor, H // factor, W // factor)351 return x352 353class PixelUnshuffleChannelAveragingDownSampleLayer3D(nn.Module):354 def __init__(355 self,356 in_channels: int,357 out_channels: int,358 factor: int,359 ):360 super().__init__()361 self.in_channels = in_channels362 self.out_channels = out_channels363 self.factor = factor364 assert in_channels * factor**3 % out_channels == 0365 self.group_size = in_channels * factor**3 // out_channels366 367 def forward(self, x: torch.Tensor, is_init=True) -> torch.Tensor:368 pad = (0, 0, 0, 0, self.factor-1, 0) # (left, right, top, bottom, front, back)369 x = F.pad(x, pad)370 B, C, D, H, W = x.shape371 x = x.view(B, C, D // self.factor, self.factor, H // self.factor, self.factor, W // self.factor, self.factor)372 x = x.permute(0, 1, 3, 5, 7, 2, 4, 6).contiguous()373 x = x.view(B, C * self.factor**3, D // self.factor, H // self.factor, W // self.factor)374 x = x.view(B, self.out_channels, self.group_size, D // self.factor, H // self.factor, W // self.factor)375 x = x.mean(dim=2)376 return x377 378 def __init__(379 self,380 in_channels: int,381 out_channels: int,382 factor: int,383 ):384 super().__init__()385 self.in_channels = in_channels386 self.out_channels = out_channels387 self.factor = factor388 assert in_channels * factor**3 % out_channels == 0389 self.group_size = in_channels * factor**3 // out_channels390 391 def forward(self, x: torch.Tensor, is_init=True) -> torch.Tensor:392 pad = (0, 0, 0, 0, self.factor-1, 0) # (left, right, top, bottom, front, back)393 x = F.pad(x, pad)394 B, C, D, H, W = x.shape395 x = x.view(B, C, D // self.factor, self.factor, H // self.factor, self.factor, W // self.factor, self.factor)396 x = x.permute(0, 1, 3, 5, 7, 2, 4, 6).contiguous()397 x = x.view(B, C * self.factor**3, D // self.factor, H // self.factor, W // self.factor)398 x = x.view(B, self.out_channels, self.group_size, D // self.factor, H // self.factor, W // self.factor)399 x = x.mean(dim=2)400 return x401 402 403 404 405def base_group_norm_with_zero_pad(x, norm_layer, act_silu=True, pad_size=2):406 out_shape = list(x.shape)407 out_shape[1] += pad_size408 out = torch.empty(out_shape, dtype=x.dtype, device=x.device)409 out[:, pad_size:] = base_group_norm(x, norm_layer, act_silu=act_silu, channel_last=True)410 out[:, :pad_size] = 0411 return out412 413 414class CausalConvChannelLast(CausalConv):415 def __init__(self,416 chan_in,417 chan_out,418 kernel_size,419 **kwargs420 ):421 super().__init__(422 chan_in, chan_out, kernel_size, **kwargs)423 424 self.time_causal_padding = (0, 0) + self.time_causal_padding425 self.time_uncausal_padding = (0, 0) + self.time_uncausal_padding426 427 def forward(self, x, is_init=True, residual=None):428 if self.is_first_run:429 self.is_first_run = False430 # self.conv.weight = nn.Parameter(self.conv.weight.permute(0,2,3,4,1).contiguous())431 432 x = nn.functional.pad(x,433 self.time_causal_padding if is_init else self.time_uncausal_padding)434 435 x = base_conv3d_channel_last(x, self.conv, residual=residual)436 return x437 438class CausalConvAfterNorm(CausalConv):439 def __init__(self,440 chan_in,441 chan_out,442 kernel_size,443 **kwargs444 ):445 super().__init__(446 chan_in, chan_out, kernel_size, **kwargs)447 448 if self.time_causal_padding == (1, 1, 1, 1, 2, 0):449 self.conv = nn.Conv3d(chan_in, chan_out, kernel_size, stride=self.stride, dilation=self.dilation, padding=(0, 1, 1), **kwargs)450 else:451 self.conv = nn.Conv3d(chan_in, chan_out, kernel_size, stride=self.stride, dilation=self.dilation, **kwargs)452 self.is_first_run = True453 454 def forward(self, x, is_init=True, residual=None):455 if self.is_first_run:456 self.is_first_run = False457 458 if self.time_causal_padding == (1, 1, 1, 1, 2, 0):459 pass460 else:461 x = nn.functional.pad(x, self.time_causal_padding).contiguous()462 463 x = base_conv3d_channel_last(x, self.conv, residual=residual)464 return x465 466class AttnBlock(nn.Module):467 def __init__(self,468 in_channels469 ):470 super().__init__()471 472 self.norm = BaseGroupNorm(num_groups=32, num_channels=in_channels)473 self.q = CausalConvChannelLast(in_channels, in_channels, kernel_size=1)474 self.k = CausalConvChannelLast(in_channels, in_channels, kernel_size=1)475 self.v = CausalConvChannelLast(in_channels, in_channels, kernel_size=1)476 self.proj_out = CausalConvChannelLast(in_channels, in_channels, kernel_size=1)477 478 def attention(self, x, is_init=True):479 x = self.norm(x, act_silu=False, channel_last=True)480 q = self.q(x, is_init)481 k = self.k(x, is_init)482 v = self.v(x, is_init)483 484 b, t, h, w, c = q.shape485 q, k, v = map(lambda x: rearrange(x, "b t h w c -> b 1 (t h w) c"), (q, k, v))486 x = nn.functional.scaled_dot_product_attention(q, k, v, is_causal=True)487 x = rearrange(x, "b 1 (t h w) c -> b t h w c", t=t, h=h, w=w)488 489 return x490 491 def forward(self, x):492 x = x.permute(0,2,3,4,1).contiguous()493 h = self.attention(x)494 x = self.proj_out(h, residual=x)495 x = x.permute(0,4,1,2,3)496 return x497 498class Resnet3DBlock(nn.Module):499 def __init__(self,500 in_channels,501 out_channels=None,502 temb_channels=512,503 conv_shortcut=False,504 ):505 super().__init__()506 507 self.in_channels = in_channels508 out_channels = in_channels if out_channels is None else out_channels509 self.out_channels = out_channels510 511 self.norm1 = BaseGroupNorm(num_groups=32, num_channels=in_channels)512 self.conv1 = CausalConvAfterNorm(in_channels, out_channels, kernel_size=3)513 if temb_channels > 0:514 self.temb_proj = nn.Linear(temb_channels, out_channels)515 516 self.norm2 = BaseGroupNorm(num_groups=32, num_channels=out_channels)517 self.conv2 = CausalConvAfterNorm(out_channels, out_channels, kernel_size=3)518 519 assert conv_shortcut is False520 self.use_conv_shortcut = conv_shortcut521 if self.in_channels != self.out_channels:522 if self.use_conv_shortcut:523 self.conv_shortcut = CausalConvAfterNorm(in_channels, out_channels, kernel_size=3)524 else:525 self.nin_shortcut = CausalConvAfterNorm(in_channels, out_channels, kernel_size=1)526 527 def forward(self, x, temb=None, is_init=True):528 x = x.permute(0,2,3,4,1).contiguous()529 530 h = self.norm1(x, zero_pad=True, act_silu=True, pad_size=2)531 h = self.conv1(h)532 if temb is not None:533 h = h + self.temb_proj(nn.functional.silu(temb))[:, :, None, None]534 535 x = self.nin_shortcut(x) if self.in_channels != self.out_channels else x536 537 h = self.norm2(h, zero_pad=True, act_silu=True, pad_size=2)538 x = self.conv2(h, residual=x)539 540 x = x.permute(0,4,1,2,3)541 return x542 543 544class Downsample3D(nn.Module):545 def __init__(self,546 in_channels,547 with_conv,548 stride549 ):550 super().__init__()551 552 self.with_conv = with_conv553 if with_conv:554 self.conv = CausalConv(in_channels, in_channels, kernel_size=3, stride=stride)555 556 def forward(self, x, is_init=True):557 if self.with_conv:558 x = self.conv(x, is_init)559 else:560 x = nn.functional.avg_pool3d(x, kernel_size=2, stride=2)561 return x562 563class VideoEncoder(nn.Module):564 def __init__(self,565 ch=32,566 ch_mult=(4, 8, 16, 16),567 num_res_blocks=2,568 in_channels=3,569 z_channels=16,570 double_z=True,571 down_sampling_layer=[1, 2],572 resamp_with_conv=True,573 version=1,574 ):575 super().__init__()576 577 temb_ch = 0578 579 self.num_resolutions = len(ch_mult)580 self.num_res_blocks = num_res_blocks581 582 # downsampling583 self.conv_in = CausalConv(in_channels, ch, kernel_size=3)584 self.down_sampling_layer = down_sampling_layer585 586 in_ch_mult = (1,) + tuple(ch_mult)587 self.down = nn.ModuleList()588 for i_level in range(self.num_resolutions):589 block = nn.ModuleList()590 attn = nn.ModuleList()591 block_in = ch * in_ch_mult[i_level]592 block_out = ch * ch_mult[i_level]593 for i_block in range(self.num_res_blocks):594 block.append(595 Resnet3DBlock(in_channels=block_in, out_channels=block_out, temb_channels=temb_ch))596 block_in = block_out597 down = nn.Module()598 down.block = block599 down.attn = attn600 if i_level != self.num_resolutions - 1:601 if i_level in self.down_sampling_layer:602 down.downsample = Downsample3D(block_in, resamp_with_conv, stride=(2, 2, 2))603 else:604 down.downsample = Downsample2D(block_in, resamp_with_conv, padding=0) #DIFF605 self.down.append(down)606 607 # middle608 self.mid = nn.Module()609 self.mid.block_1 = Resnet3DBlock(in_channels=block_in, out_channels=block_in, temb_channels=temb_ch)610 self.mid.attn_1 = AttnBlock(block_in)611 self.mid.block_2 = Resnet3DBlock(in_channels=block_in, out_channels=block_in, temb_channels=temb_ch)612 613 # end614 self.norm_out = nn.GroupNorm(num_groups=32, num_channels=block_in)615 self.version = version616 if version == 2:617 channels = 4 * z_channels * 2 ** 3618 self.conv_patchify = ConvPixelUnshuffleDownSampleLayer3D(block_in, channels, kernel_size=3, factor=2)619 self.shortcut_pathify = PixelUnshuffleChannelAveragingDownSampleLayer3D(block_in, channels, 2)620 self.shortcut_out = PixelUnshuffleChannelAveragingDownSampleLayer3D(channels, 2 * z_channels if double_z else z_channels, 1)621 self.conv_out = CausalConvChannelLast(channels, 2 * z_channels if double_z else z_channels, kernel_size=3)622 else:623 self.conv_out = CausalConvAfterNorm(block_in, 2 * z_channels if double_z else z_channels, kernel_size=3)624 625 @torch.inference_mode()626 def forward(self, x, video_frame_num, is_init=True):627 # timestep embedding628 temb = None629 630 t = video_frame_num631 632 # downsampling633 h = self.conv_in(x, is_init)634 635 # make it real channel last, but behave like normal layout636 h = h.permute(0,2,3,4,1).contiguous().permute(0,4,1,2,3)637 638 for i_level in range(self.num_resolutions):639 for i_block in range(self.num_res_blocks):640 h = self.down[i_level].block[i_block](h, temb, is_init)641 if len(self.down[i_level].attn) > 0:642 h = self.down[i_level].attn[i_block](h)643 644 if i_level != self.num_resolutions - 1:645 if isinstance(self.down[i_level].downsample, Downsample2D):646 _, _, t, _, _ = h.shape647 h = rearrange(h, "b c t h w -> (b t) h w c", t=t)648 h = self.down[i_level].downsample(h)649 h = rearrange(h, "(b t) h w c -> b c t h w", t=t)650 else:651 h = self.down[i_level].downsample(h, is_init)652 653 h = self.mid.block_1(h, temb, is_init)654 h = self.mid.attn_1(h)655 h = self.mid.block_2(h, temb, is_init)656 657 h = h.permute(0,2,3,4,1).contiguous() # b c l h w -> b l h w c658 if self.version == 2:659 h = base_group_norm(h, self.norm_out, act_silu=True, channel_last=True)660 h = h.permute(0,4,1,2,3).contiguous()661 shortcut = self.shortcut_pathify(h, is_init)662 h = self.conv_patchify(h, is_init)663 h = h.add_(shortcut)664 shortcut = self.shortcut_out(h, is_init).permute(0,2,3,4,1)665 h = self.conv_out(h.permute(0,2,3,4,1).contiguous(), is_init)666 h = h.add_(shortcut)667 else:668 h = base_group_norm_with_zero_pad(h, self.norm_out, act_silu=True, pad_size=2)669 h = self.conv_out(h, is_init)670 h = h.permute(0,4,1,2,3) # b l h w c -> b c l h w671 672 h = rearrange(h, "b c t h w -> b t c h w")673 return h674 675 676class Res3DBlockUpsample(nn.Module):677 def __init__(self,678 input_filters,679 num_filters,680 down_sampling_stride,681 down_sampling=False682 ):683 super().__init__()684 685 self.input_filters = input_filters686 self.num_filters = num_filters687 688 self.act_ = nn.SiLU(inplace=True)689 690 self.conv1 = CausalConvChannelLast(num_filters, num_filters, kernel_size=[3, 3, 3])691 self.norm1 = BaseGroupNorm(32, num_filters)692 693 self.conv2 = CausalConvChannelLast(num_filters, num_filters, kernel_size=[3, 3, 3])694 self.norm2 = BaseGroupNorm(32, num_filters)695 696 self.down_sampling = down_sampling697 if down_sampling:698 self.down_sampling_stride = down_sampling_stride699 else:700 self.down_sampling_stride = [1, 1, 1]701 702 if num_filters != input_filters or down_sampling:703 self.conv3 = CausalConvChannelLast(input_filters, num_filters, kernel_size=[1, 1, 1], stride=self.down_sampling_stride)704 self.norm3 = BaseGroupNorm(32, num_filters)705 706 def forward(self, x, is_init=False):707 x = x.permute(0,2,3,4,1).contiguous()708 709 residual = x710 711 h = self.conv1(x, is_init)712 h = self.norm1(h, act_silu=True, channel_last=True)713 714 h = self.conv2(h, is_init)715 h = self.norm2(h, act_silu=False, channel_last=True)716 717 if self.down_sampling or self.num_filters != self.input_filters:718 x = self.conv3(x, is_init)719 x = self.norm3(x, act_silu=False, channel_last=True)720 721 h.add_(x)722 h = self.act_(h)723 if residual is not None:724 h.add_(residual)725 726 h = h.permute(0,4,1,2,3)727 return h728 729class Upsample3D(nn.Module):730 def __init__(self,731 in_channels,732 scale_factor=2733 ):734 super().__init__()735 736 self.scale_factor = scale_factor737 self.conv3d = Res3DBlockUpsample(input_filters=in_channels,738 num_filters=in_channels,739 down_sampling_stride=(1, 1, 1),740 down_sampling=False)741 742 def forward(self, x, is_init=True, is_split=True):743 b, c, t, h, w = x.shape744 745 # x = x.permute(0,2,3,4,1).contiguous().permute(0,4,1,2,3).to(memory_format=torch.channels_last_3d)746 if is_split:747 split_size = c // 8748 x_slices = torch.split(x, split_size, dim=1)749 x = [nn.functional.interpolate(x, scale_factor=self.scale_factor) for x in x_slices]750 x = torch.cat(x, dim=1)751 else:752 x = nn.functional.interpolate(x, scale_factor=self.scale_factor)753 754 x = self.conv3d(x, is_init)755 return x756 757class VideoDecoder(nn.Module):758 def __init__(self,759 ch=128,760 z_channels=16,761 out_channels=3,762 ch_mult=(1, 2, 4, 4),763 num_res_blocks=2,764 temporal_up_layers=[2, 3],765 temporal_downsample=4,766 resamp_with_conv=True,767 version=1,768 ):769 super().__init__()770 771 temb_ch = 0772 773 self.num_resolutions = len(ch_mult)774 self.num_res_blocks = num_res_blocks775 self.temporal_downsample = temporal_downsample776 777 block_in = ch * ch_mult[self.num_resolutions - 1]778 self.version = version779 if version == 2:780 channels = 4 * z_channels * 2 ** 3781 self.conv_in = CausalConv(z_channels, channels, kernel_size=3)782 self.shortcut_in = ChannelDuplicatingPixelUnshuffleUpSampleLayer3D(z_channels, channels, 1)783 self.conv_unpatchify = ConvPixelShuffleUpSampleLayer3D(channels, block_in, kernel_size=3, factor=2)784 self.shortcut_unpathify = ChannelDuplicatingPixelUnshuffleUpSampleLayer3D(channels, block_in, 2)785 else:786 self.conv_in = CausalConv(z_channels, block_in, kernel_size=3)787 788 # middle789 self.mid = nn.Module()790 self.mid.block_1 = Resnet3DBlock(in_channels=block_in, out_channels=block_in, temb_channels=temb_ch)791 self.mid.attn_1 = AttnBlock(block_in)792 self.mid.block_2 = Resnet3DBlock(in_channels=block_in, out_channels=block_in, temb_channels=temb_ch)793 794 # upsampling795 self.up_id = len(temporal_up_layers)796 self.video_frame_num = 1797 self.cur_video_frame_num = self.video_frame_num // 2 ** self.up_id + 1798 self.up = nn.ModuleList()799 for i_level in reversed(range(self.num_resolutions)):800 block = nn.ModuleList()801 attn = nn.ModuleList()802 block_out = ch * ch_mult[i_level]803 for i_block in range(self.num_res_blocks + 1):804 block.append(805 Resnet3DBlock(in_channels=block_in, out_channels=block_out, temb_channels=temb_ch))806 block_in = block_out807 up = nn.Module()808 up.block = block809 up.attn = attn810 if i_level != 0:811 if i_level in temporal_up_layers:812 up.upsample = Upsample3D(block_in)813 self.cur_video_frame_num = self.cur_video_frame_num * 2814 else:815 up.upsample = Upsample2D(block_in, resamp_with_conv)816 self.up.insert(0, up) # prepend to get consistent order817 818 # end819 self.norm_out = nn.GroupNorm(num_groups=32, num_channels=block_in)820 self.conv_out = CausalConvAfterNorm(block_in, out_channels, kernel_size=3)821 822 @torch.inference_mode()823 def forward(self, z, is_init=True):824 z = rearrange(z, "b t c h w -> b c t h w")825 826 h = self.conv_in(z, is_init=is_init)827 if self.version == 2:828 shortcut = self.shortcut_in(z, is_init=is_init)829 h = h.add_(shortcut)830 shortcut = self.shortcut_unpathify(h, is_init=is_init)831 h = self.conv_unpatchify(h, is_init=is_init)832 h = h.add_(shortcut)833 834 temb = None835 836 h = h.permute(0,2,3,4,1).contiguous().permute(0,4,1,2,3)837 h = self.mid.block_1(h, temb, is_init=is_init)838 h = self.mid.attn_1(h)839 h = h.permute(0,2,3,4,1).contiguous().permute(0,4,1,2,3)840 h = self.mid.block_2(h, temb, is_init=is_init)841 842 # upsampling843 for i_level in reversed(range(self.num_resolutions)):844 for i_block in range(self.num_res_blocks + 1):845 h = h.permute(0,2,3,4,1).contiguous().permute(0,4,1,2,3)846 h = self.up[i_level].block[i_block](h, temb, is_init=is_init)847 if len(self.up[i_level].attn) > 0:848 h = self.up[i_level].attn[i_block](h)849 if i_level != 0:850 if isinstance(self.up[i_level].upsample, Upsample2D) or (hasattr(self.up[i_level].upsample, "module") and isinstance(self.up[i_level].upsample.module, Upsample2D)):851 B = h.size(0)852 h = h.permute(0,2,3,4,1).flatten(0,1)853 h = self.up[i_level].upsample(h)854 h = h.unflatten(0, (B, -1)).permute(0,4,1,2,3)855 else:856 h = self.up[i_level].upsample(h, is_init=is_init)857 858 # end859 h = h.permute(0,2,3,4,1) # b c l h w -> b l h w c860 self.norm_out.to(dtype=h.dtype, device=h.device) # To be updated861 h = base_group_norm_with_zero_pad(h, self.norm_out, act_silu=True, pad_size=2)862 h = self.conv_out(h)863 h = h.permute(0,4,1,2,3)864 865 if is_init:866 h = h[:, :, (self.temporal_downsample - 1):]867 return h868 869 870 871def rms_norm(input, normalized_shape, eps=1e-6):872 dtype = input.dtype873 input = input.to(torch.float32)874 variance = input.pow(2).flatten(-len(normalized_shape)).mean(-1)[(...,) + (None,) * len(normalized_shape)]875 input = input * torch.rsqrt(variance + eps)876 return input.to(dtype)877 878class DiagonalGaussianDistribution(object):879 def __init__(self, parameters, deterministic=False, rms_norm_mean=False, only_return_mean=False):880 self.parameters = parameters881 self.mean, self.logvar = torch.chunk(parameters, 2, dim=-3) #N,[X],C,H,W882 self.logvar = torch.clamp(self.logvar, -30.0, 20.0)883 self.std = torch.exp(0.5 * self.logvar)884 self.var = torch.exp(self.logvar)885 self.deterministic = deterministic886 if self.deterministic:887 self.var = self.std = torch.zeros_like(888 self.mean,889 device=self.parameters.device,890 dtype=self.parameters.dtype)891 if rms_norm_mean:892 self.mean = rms_norm(self.mean, self.mean.size()[1:])893 self.only_return_mean = only_return_mean894 895 def sample(self, generator=None):896 # make sure sample is on the same device897 # as the parameters and has same dtype898 sample = torch.randn(899 self.mean.shape, generator=generator, device=self.parameters.device)900 sample = sample.to(dtype=self.parameters.dtype)901 x = self.mean + self.std * sample902 if self.only_return_mean:903 return self.mean904 else:905 return x906 907 908class StepVideoVAE(nn.Module):909 def __init__(self,910 in_channels=3,911 out_channels=3,912 z_channels=64,913 num_res_blocks=2,914 model_path=None,915 weight_dict={},916 world_size=1,917 version=2,918 ):919 super().__init__()920 921 self.frame_len = 17922 self.latent_len = 3 if version == 2 else 5923 924 base_group_norm.spatial = True if version == 2 else False925 926 self.encoder = VideoEncoder(927 in_channels=in_channels,928 z_channels=z_channels,929 num_res_blocks=num_res_blocks,930 version=version,931 )932 933 self.decoder = VideoDecoder(934 z_channels=z_channels,935 out_channels=out_channels,936 num_res_blocks=num_res_blocks,937 version=version,938 )939 940 if model_path is not None:941 weight_dict = self.init_from_ckpt(model_path)942 if len(weight_dict) != 0:943 self.load_from_dict(weight_dict)944 self.convert_channel_last()945 946 self.world_size = world_size947 948 def init_from_ckpt(self, model_path):949 from safetensors import safe_open950 p = {}951 with safe_open(model_path, framework="pt", device="cpu") as f:952 for k in f.keys():953 tensor = f.get_tensor(k)954 if k.startswith("decoder.conv_out."):955 k = k.replace("decoder.conv_out.", "decoder.conv_out.conv.")956 p[k] = tensor957 return p958 959 def load_from_dict(self, p):960 self.load_state_dict(p)961 962 def convert_channel_last(self):963 #Conv2d NCHW->NHWC964 pass965 966 def naive_encode(self, x, is_init_image=True):967 b, l, c, h, w = x.size()968 x = rearrange(x, 'b l c h w -> b c l h w').contiguous()969 z = self.encoder(x, l, True) # 下采样[1, 4, 8, 16, 16]970 return z971 972 @torch.inference_mode()973 def encode(self, x):974 # b (nc cf) c h w -> (b nc) cf c h w -> encode -> (b nc) cf c h w -> b (nc cf) c h w975 chunks = list(x.split(self.frame_len, dim=1))976 for i in range(len(chunks)):977 chunks[i] = self.naive_encode(chunks[i], True)978 z = torch.cat(chunks, dim=1)979 980 posterior = DiagonalGaussianDistribution(z)981 return posterior.sample()982 983 def decode_naive(self, z, is_init=True):984 z = z.to(next(self.decoder.parameters()).dtype)985 dec = self.decoder(z, is_init)986 return dec987 988 @torch.inference_mode()989 def decode_original(self, z):990 # b (nc cf) c h w -> (b nc) cf c h w -> decode -> (b nc) c cf h w -> b (nc cf) c h w991 chunks = list(z.split(self.latent_len, dim=1))992 993 if self.world_size > 1:994 chunks_total_num = len(chunks)995 max_num_per_rank = (chunks_total_num + self.world_size - 1) // self.world_size996 rank = torch.distributed.get_rank()997 chunks_ = chunks[max_num_per_rank * rank : max_num_per_rank * (rank + 1)]998 if len(chunks_) < max_num_per_rank:999 chunks_.extend(chunks[:max_num_per_rank-len(chunks_)])1000 chunks = chunks_1001 1002 for i in range(len(chunks)):1003 chunks[i] = self.decode_naive(chunks[i], True).permute(0,2,1,3,4)1004 x = torch.cat(chunks, dim=1)1005 1006 if self.world_size > 1:1007 x_ = torch.empty([x.size(0), (self.world_size * max_num_per_rank) * self.frame_len, *x.shape[2:]], dtype=x.dtype, device=x.device)1008 torch.distributed.all_gather_into_tensor(x_, x)1009 x = x_[:, : chunks_total_num * self.frame_len]1010 1011 x = self.mix(x)1012 return x1013 1014 def mix(self, x, smooth_scale = 0.6):1015 remain_scale = smooth_scale1016 mix_scale = 1. - remain_scale1017 front = slice(self.frame_len - 1, x.size(1) - 1, self.frame_len)1018 back = slice(self.frame_len, x.size(1), self.frame_len)1019 x[:, front], x[:, back] = (1020 x[:, front] * remain_scale + x[:, back] * mix_scale,1021 x[:, back] * remain_scale + x[:, front] * mix_scale1022 )1023 return x1024 1025 def single_decode(self, hidden_states, device):1026 chunks = list(hidden_states.split(self.latent_len, dim=1))1027 for i in range(len(chunks)):1028 chunks[i] = self.decode_naive(chunks[i].to(device), True).permute(0,2,1,3,4).cpu()1029 x = torch.cat(chunks, dim=1)1030 return x1031 1032 def build_1d_mask(self, length, left_bound, right_bound, border_width):1033 x = torch.ones((length,))1034 if not left_bound:1035 x[:border_width] = (torch.arange(border_width) + 1) / border_width1036 if not right_bound:1037 x[-border_width:] = torch.flip((torch.arange(border_width) + 1) / border_width, dims=(0,))1038 return x1039 1040 def build_mask(self, data, is_bound, border_width):1041 _, _, _, H, W = data.shape1042 h = self.build_1d_mask(H, is_bound[0], is_bound[1], border_width[0])1043 w = self.build_1d_mask(W, is_bound[2], is_bound[3], border_width[1])1044 1045 h = repeat(h, "H -> H W", H=H, W=W)1046 w = repeat(w, "W -> H W", H=H, W=W)1047 1048 mask = torch.stack([h, w]).min(dim=0).values1049 mask = rearrange(mask, "H W -> 1 1 1 H W")1050 return mask1051 1052 def tiled_decode(self, hidden_states, device, tile_size=(34, 34), tile_stride=(16, 16)):1053 B, T, C, H, W = hidden_states.shape1054 size_h, size_w = tile_size1055 stride_h, stride_w = tile_stride1056 1057 # Split tasks1058 tasks = []1059 for t in range(0, T, 3):1060 for h in range(0, H, stride_h):1061 if (h-stride_h >= 0 and h-stride_h+size_h >= H): continue1062 for w in range(0, W, stride_w):1063 if (w-stride_w >= 0 and w-stride_w+size_w >= W): continue1064 t_, h_, w_ = t + 3, h + size_h, w + size_w1065 tasks.append((t, t_, h, h_, w, w_))1066 1067 # Run1068 data_device = "cpu"1069 computation_device = device1070 1071 weight = torch.zeros((1, 1, T//3*17, H * 16, W * 16), dtype=hidden_states.dtype, device=data_device)1072 values = torch.zeros((B, 3, T//3*17, H * 16, W * 16), dtype=hidden_states.dtype, device=data_device)1073 1074 for t, t_, h, h_, w, w_ in tqdm(tasks, desc="VAE decoding"):1075 hidden_states_batch = hidden_states[:, t:t_, :, h:h_, w:w_].to(computation_device)1076 hidden_states_batch = self.decode_naive(hidden_states_batch, True).to(data_device)1077 1078 mask = self.build_mask(1079 hidden_states_batch,1080 is_bound=(h==0, h_>=H, w==0, w_>=W),1081 border_width=((size_h - stride_h) * 16, (size_w - stride_w) * 16)1082 ).to(dtype=hidden_states.dtype, device=data_device)1083 1084 target_t = t // 3 * 171085 target_h = h * 161086 target_w = w * 161087 values[1088 :,1089 :,1090 target_t: target_t + hidden_states_batch.shape[2],1091 target_h: target_h + hidden_states_batch.shape[3],1092 target_w: target_w + hidden_states_batch.shape[4],1093 ] += hidden_states_batch * mask1094 weight[1095 :,1096 :,1097 target_t: target_t + hidden_states_batch.shape[2],1098 target_h: target_h + hidden_states_batch.shape[3],1099 target_w: target_w + hidden_states_batch.shape[4],1100 ] += mask1101 return values / weight1102 1103 def decode(self, hidden_states, device, tiled=False, tile_size=(34, 34), tile_stride=(16, 16), smooth_scale=0.6):1104 hidden_states = hidden_states.to("cpu")1105 if tiled:1106 video = self.tiled_decode(hidden_states, device, tile_size, tile_stride)1107 else:1108 video = self.single_decode(hidden_states, device)1109 video = self.mix(video, smooth_scale=smooth_scale)1110 return video1111 1112 @staticmethod1113 def state_dict_converter():1114 return StepVideoVAEStateDictConverter()1115 1116 1117class StepVideoVAEStateDictConverter:1118 def __init__(self):1119 super().__init__()1120 1121 def from_diffusers(self, state_dict):1122 return self.from_civitai(state_dict)1123 1124 def from_civitai(self, state_dict):1125 state_dict_ = {}1126 for name, param in state_dict.items():1127 if name.startswith("decoder.conv_out."):1128 name_ = name.replace("decoder.conv_out.", "decoder.conv_out.conv.")1129 else:1130 name_ = name1131 state_dict_[name_] = param1132 return state_dict_1133 