Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
stepvideo_vae.py1133 linesDownload Raw Back to models
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