CoolFace
Modelpublic

1ST-PLACE-WINNER/MiniMax-H3

sourceHugging Faceotherupdated 2d agoView on Hugging Face
0likes77downloads
vae_cnn.py305 linesDownload Raw Back to video_vae
1# SPDX-License-Identifier: Apache-2.02# 3D causal CNN encoder for the MiniMax H3 visual VAE (inference-only bundle).3import os4import torch.nn as nn5import torch.nn.functional as F6 7from .attention import maybe_checkpoint8from .conv import SpatialParallelConv3d9from .norm import get_spatial_norm_3d10from .parallel import get_parallel_state, exchange_strides11from .norm import get_group_norm_3d12 13 14 15 16 17 18 19 20 21 22# ============================================================================23# 3D CNN Components24# ============================================================================25 26 27def norm_silu(x, norm, cond=None):28    if cond is None:29        return F.silu(norm(x))30    else:31        return F.silu(norm(x, cond))32 33 34class Downsample3D(nn.Module):35    def __init__(36        self,37        in_channels,38        out_channels,39        time_stride=1,40        space_stride=2,41        padding_mode="zeros",42        padding_mode_t=None,43        causal=True,44    ):45        super().__init__()46        self.time_stride = time_stride47        self.space_stride = space_stride48 49        assert time_stride in [1, 2]50        assert space_stride in [1, 2, 3]51 52        self.conv = SpatialParallelConv3d(53            in_channels,54            out_channels,55            kernel_size=3,56            padding=(1, 0, 0),57            stride=(time_stride, space_stride, space_stride),58            padding_mode=padding_mode,59            padding_mode_t=padding_mode_t,60            causal=causal,61        )62        self.causal = self.conv.causal63        self.pad_mode = self.conv.pad_mode64 65    def forward(self, x):66        if self.space_stride == 2:67            if getattr(self.conv, "spatial_parallel", False):68                state = get_parallel_state()69                x = exchange_strides(70                    x,71                    self.pad_mode,72                    state["sp_rank"],73                    state["sp_size"],74                    state["sp_process_group"],75                    self.conv.chunk_dim,76                )77            else:78                pad = (0, 1, 0, 1, 0, 0)79                x = F.pad(x, pad, mode=self.pad_mode)80        return self.conv(x)81 82 83class ResnetBlock3D(nn.Module):84    def __init__(85        self,86        in_channels,87        out_channels=None,88        zq_ch=None,89        padding_mode="zeros",90        padding_mode_t=None,91        causal=True,92        use_t_isolated_gn=False,93    ):94        super().__init__()95        self.in_channels = in_channels96        out_channels = in_channels if out_channels is None else out_channels97        self.out_channels = out_channels98 99        self.use_fused_norm = (100            os.environ.get("MINIMAX_H3_USE_FUSED_NORM", "false").lower() == "true"101        )102 103        if zq_ch is None:104            self.norm1 = get_group_norm_3d(in_channels, use_t_isolated_gn=use_t_isolated_gn)105            self.norm2 = get_group_norm_3d(out_channels, use_t_isolated_gn=use_t_isolated_gn)106        else:107            self.norm1 = get_spatial_norm_3d(108                in_channels,109                zq_ch,110                padding_mode=padding_mode,111                padding_mode_t=padding_mode_t,112                causal=causal,113                use_t_isolated_gn=use_t_isolated_gn,114            )115            self.norm2 = get_spatial_norm_3d(116                out_channels,117                zq_ch,118                padding_mode=padding_mode,119                padding_mode_t=padding_mode_t,120                causal=causal,121                use_t_isolated_gn=use_t_isolated_gn,122            )123 124        self.conv1 = SpatialParallelConv3d(125            in_channels,126            out_channels,127            kernel_size=3,128            padding=1,129            padding_mode=padding_mode,130            padding_mode_t=padding_mode_t,131            causal=causal,132        )133 134        self.conv2 = SpatialParallelConv3d(135            out_channels,136            out_channels,137            kernel_size=3,138            padding=1,139            padding_mode=padding_mode,140            padding_mode_t=padding_mode_t,141            causal=causal,142        )143 144        if self.in_channels != self.out_channels:145            self.nin_shortcut = SpatialParallelConv3d(146                in_channels,147                out_channels,148                kernel_size=1,149                padding_mode=padding_mode,150                padding_mode_t=padding_mode_t,151                causal=causal,152            )153 154    def forward(self, x, zq=None):155        h = x156 157        if self.use_fused_norm:158            h = self.norm1(h, zq)159        else:160            h = norm_silu(h, self.norm1, zq)161 162        h = self.conv1(h)163 164        if self.use_fused_norm:165            h = self.norm2(h, zq)166        else:167            h = norm_silu(h, self.norm2, zq)168 169        h = self.conv2(h)170 171        if self.in_channels != self.out_channels:172            x = self.nin_shortcut(x)173 174        return x + h175 176 177class EncoderFCN3D(nn.Module):178    def __init__(179        self,180        ch,181        ch_mult,182        space_down,183        time_down,184        num_res_blocks,185        in_channels,186        z_channels,187        double_z=False,188        zq_ch=None,189        padding_mode="zeros",190        padding_mode_t=None,191        causal=True,192        use_t_isolated_gn=False,193    ):194        super().__init__()195        self.ch = ch196        self.num_levels = len(ch_mult)197 198        if isinstance(num_res_blocks, int):199            self.num_res_blocks = [num_res_blocks] * self.num_levels200        else:201            self.num_res_blocks = num_res_blocks202 203        self.space_down_factors = space_down204        self.time_down_factors = time_down205        self.in_channels = in_channels206 207        self.use_fused_norm = (208            os.environ.get("MINIMAX_H3_USE_FUSED_NORM", "false").lower() == "true"209        )210 211        block_mid = [ch * ch_mult[i] for i in range(self.num_levels)]212        block_in = [block_mid[0]] + block_mid[:-1]213        block_out = block_mid214 215        conv_kwargs = dict(216            padding_mode=padding_mode,217            padding_mode_t=padding_mode_t,218            causal=causal,219        )220 221        self.conv_in = SpatialParallelConv3d(222            in_channels, block_in[0], kernel_size=3, padding=1, **conv_kwargs223        )224 225        self.down = nn.ModuleList()226        for i_level in range(self.num_levels):227            down = nn.Module()228 229            down.block = nn.ModuleList()230            for i in range(self.num_res_blocks[i_level]):231                down.block.append(232                    ResnetBlock3D(233                        in_channels=block_in[i_level] if i == 0 else block_mid[i_level],234                        out_channels=block_mid[i_level],235                        zq_ch=zq_ch,236                        use_t_isolated_gn=use_t_isolated_gn,237                        **conv_kwargs,238                    )239                )240 241            if space_down[i_level] * time_down[i_level] > 1:242                down.downsample = Downsample3D(243                    block_mid[i_level],244                    block_out[i_level],245                    time_stride=time_down[i_level],246                    space_stride=space_down[i_level],247                    **conv_kwargs,248                )249            else:250                if block_out[i_level] != block_mid[i_level]:251                    down.downsample = SpatialParallelConv3d(252                        block_mid[i_level],253                        block_out[i_level],254                        kernel_size=1,255                        **conv_kwargs,256                    )257 258            self.down.append(down)259 260        if zq_ch is None:261            self.norm_out = get_group_norm_3d(262                block_out[-1], use_t_isolated_gn=use_t_isolated_gn263            )264        else:265            self.norm_out = get_spatial_norm_3d(266                block_out[-1],267                zq_ch,268                use_t_isolated_gn=use_t_isolated_gn,269                **conv_kwargs,270            )271 272        self.conv_out = SpatialParallelConv3d(273            block_out[-1],274            2 * z_channels if double_z else z_channels,275            kernel_size=3,276            padding=1,277            **conv_kwargs,278        )279 280        self.gradient_checkpointing = False281 282    def _set_gradient_checkpointing(self, module, value=False):283        if hasattr(module, "gradient_checkpointing"):284            module.gradient_checkpointing = value285 286    def forward(self, x, zq=None):287        h = self.conv_in(x)288        for i_level in range(self.num_levels):289            for i_block in range(self.num_res_blocks[i_level]):290                h = maybe_checkpoint(self, self.down[i_level].block[i_block], h, zq)291            if hasattr(self.down[i_level], "downsample"):292                h = self.down[i_level].downsample(h)293 294        if self.use_fused_norm:295            h = self.norm_out(h, zq)296        else:297            h = norm_silu(h, self.norm_out, zq)298 299        h = self.conv_out(h)300        return h301 302 303 304 305