Team Ai
Modelpublic

MiniMaxAI/MiniMax-H3

sourceHugging Faceotherupdated 2mo agoView on Hugging Face
6klikes3.6mdownloads
vae_vit.py381 linesDownload Raw Back to video_vae
1# SPDX-License-Identifier: Apache-2.02# ViT3D decoder for the MiniMax H3 visual VAE (inference-only bundle).3import torch4import torch.nn as nn5import torch.distributed as dist6from diffusers.configuration_utils import ConfigMixin, register_to_config7from diffusers.models.modeling_utils import ModelMixin8from diffusers.utils import logging9 10from .attention import maybe_checkpoint11from .base_module import TransformerBlock, RotaryEmbeddingND12from .flash import make_block_causal_mask_mod13from .func import create_token_ids14from .parallel import get_subseq, gather_subseq, get_parallel_state15 16logger = logging.get_logger(__name__)17 18 19def _linear_with_module_dtype(linear, tensor, out_dtype=None):20    weight = getattr(linear, "weight", None)21    target_dtype = getattr(weight, "dtype", tensor.dtype)22    output = linear(tensor.to(target_dtype))23    if out_dtype is not None and output.dtype != out_dtype:24        output = output.to(out_dtype)25    return output26 27 28def _make_seq_len_mask_mod(seq_len, base_mask_mod=None):29    if base_mask_mod is None:30        def mask_mod(b, h, q_idx, kv_idx, seqlen_info, aux_tensors):31            return (q_idx < seq_len) & (kv_idx < seq_len)32 33        mask_mod.block_sparse_cache_key = ("seq_len", seq_len)34        return mask_mod35 36    def mask_mod(b, h, q_idx, kv_idx, seqlen_info, aux_tensors):37        return (38            (q_idx < seq_len)39            & (kv_idx < seq_len)40            & base_mask_mod(b, h, q_idx, kv_idx, seqlen_info, aux_tensors)41        )42 43    base_cache_key = getattr(base_mask_mod, "block_sparse_cache_key", None)44    if base_cache_key is not None:45        mask_mod.block_sparse_cache_key = ("seq_len", seq_len, base_cache_key)46    if hasattr(base_mask_mod, "use_fast_sampling"):47        mask_mod.use_fast_sampling = base_mask_mod.use_fast_sampling48    return mask_mod49 50 51 52 53 54 55def _pack_tensors_3d(tensors, patch_size, patch_size_t):56    batch_size, num_channels_tensors, temporal, height, width = tensors.shape57 58    tensors = tensors.view(59        batch_size,60        num_channels_tensors,61        temporal // patch_size_t,62        patch_size_t,63        height // patch_size,64        patch_size,65        width // patch_size,66        patch_size,67    )68    tensors = tensors.permute(0, 2, 4, 6, 1, 3, 5, 7)69    tensors = tensors.reshape(70        batch_size,71        (temporal // patch_size_t) * (height // patch_size) * (width // patch_size),72        num_channels_tensors * patch_size_t * patch_size * patch_size,73    )74    return tensors75 76 77def _unpack_tensors_3d(tensors, patch_size, patch_size_t, temporal, height, width):78    batch_size, num_patches, channels = tensors.shape79    num_channels_tensors = channels // (patch_size_t * patch_size * patch_size)80 81    tensors = tensors.view(82        batch_size,83        temporal // patch_size_t,84        height // patch_size,85        width // patch_size,86        num_channels_tensors,87        patch_size_t,88        patch_size,89        patch_size,90    )91    tensors = tensors.permute(0, 4, 1, 5, 2, 6, 3, 7).contiguous()92    tensors = tensors.reshape(batch_size, num_channels_tensors, temporal, height, width)93    return tensors94 95 96class ViTBase(ModelMixin, ConfigMixin):97    """Base class for ViT Encoder and Decoder with common functionality."""98 99    _supports_gradient_checkpointing = True100    _no_split_modules = ["TransformerBlock"]101    gradient_checkpointing_mode = "full"102 103    def _set_gradient_checkpointing(self, module, value=False):104        if hasattr(module, "gradient_checkpointing"):105            module.gradient_checkpointing = value106 107    def set_spatial_parallel(self, enabled):108        self.spatial_parallel = enabled109        if hasattr(self, "transformer_blocks"):110            for block in self.transformer_blocks:111                block.attn.spatial_parallel = enabled112 113    def _init_weights(self):114        def basic_init(m):115            if isinstance(m, nn.Linear):116                nn.init.xavier_uniform_(m.weight)117                if m.bias is not None:118                    nn.init.constant_(m.bias, 0)119 120        self.apply(basic_init)121 122    def init_mask_config(self, dim, is_3d=False):123        self._mask_dim = dim124        self._mask_is_3d = is_3d125        self.register_buffer("mask_token", torch.zeros(1, 1, dim))126 127    def set_mask_config(self, mask_config):128        self.mask_prob = mask_config.get("mask_prob", 0.0)129        self.mask_enabled = self.mask_prob > 0130        self.mask_style = mask_config.get("mask_style", "replace")131        if self.mask_enabled and self.mask_style == "drop" and self.mask_prob < 1.0:132            logger.warning("mask_style='drop' with mask_prob < 1.0")133        if self._mask_is_3d:134            self.temporal_scale_range = mask_config.get("temporal_scale_range", (0.3, 0.5))135            self.spatial_scale_range = mask_config.get("spatial_scale_range", (0.1, 0.25))136            self.min_mask_ratio = mask_config.get("min_mask_ratio", 0.75)137            self.max_mask_ratio = mask_config.get("max_mask_ratio", 0.95)138        else:139            self.spatial_scale_range = mask_config.get("spatial_scale_range", (0.15, 0.15))140            self.min_mask_ratio = mask_config.get("min_mask_ratio", 0.5)141            self.max_mask_ratio = mask_config.get("max_mask_ratio", 0.75)142        self.aspect_ratio_range = mask_config.get("aspect_ratio_range", (0.75, 1.5))143        self.max_retries = mask_config.get("max_retries", 100)144        if self.mask_enabled and self.mask_style == "drop" and getattr(self, "t_causal", False):145            logger.warning("mask_style='drop' with t_causal may cause issues")146        if self.mask_enabled and "mask_token" in self._buffers:147            del self._buffers["mask_token"]148            self.mask_token = nn.Parameter(torch.randn(1, 1, self._mask_dim) * 0.02)149 150    def init_suffix_tokens(self, dim, num_register_tokens, has_cls_token=True):151        self.num_register_tokens = num_register_tokens152        if num_register_tokens > 0:153            self.register_tokens = nn.Parameter(torch.randn(1, num_register_tokens, dim) * 0.02)154        else:155            self.register_tokens = None156        if has_cls_token:157            self.cls_token = nn.Parameter(torch.randn(1, 1, dim) * 0.02)158 159    def apply_mask_preprocess(self, hidden_states, img_ids, patch_dims, num_suffix):160        if self.training and self.mask_enabled:161            raise NotImplementedError(162                "mask modeling is not supported in this inference-only bundle"163            )164        return hidden_states, img_ids165 166    def forward_transformer_blocks(self, hidden_states, rotary_pos_emb, pack_info=None):167        if pack_info is None:168            pack_info = {}169        for block in self.transformer_blocks:170            hidden_states = maybe_checkpoint(171                self, block, hidden_states, rotary_pos_emb, pack_info172            )173        return hidden_states174 175    def _pad_for_sp(self, hidden_states, img_ids, pack_info=None):176        if pack_info is None:177            pack_info = {}178        if not self.spatial_parallel:179            return hidden_states, img_ids, pack_info, 0180 181        seq_len = hidden_states.shape[1]182        sp_size = get_parallel_state().get("sp_size", 1)183        pad_len = (-seq_len) % sp_size184        if pad_len == 0:185            return hidden_states, img_ids, pack_info, 0186 187        hidden_states = torch.nn.functional.pad(hidden_states, (0, 0, 0, pad_len))188        img_ids = torch.nn.functional.pad(img_ids, (0, 0, 0, pad_len))189 190        pack_info = dict(pack_info)191        base_mask_mod = pack_info.get("mask_mod")192        pack_info["mask_mod"] = _make_seq_len_mask_mod(seq_len, base_mask_mod)193        pack_info.pop("block_sparse", None)194        return hidden_states, img_ids, pack_info, pad_len195 196    @staticmethod197    def _unpad_for_sp(hidden_states, pad_len):198        if pad_len == 0:199            return hidden_states200        return hidden_states[:, :-pad_len, :]201 202    def apply_mask_postprocess(self, hidden_states, num_patches):203        if self.training and self.mask_enabled and self.mask_style == "drop":204            raise NotImplementedError(205                "mask modeling is not supported in this inference-only bundle"206            )207        return hidden_states208 209 210 211 212 213 214 215 216class ViT3DDecoder(ViTBase):217    """Vision Transformer Video Decoder using TransformerBlock."""218 219    @register_to_config220    def __init__(221        self,222        patch_size: int = 16,223        patch_size_t: int = 4,224        t_causal: bool = False,225        in_channels: int = 16,226        out_channels: int = 3,227        num_layers: int = 24,228        heads: int = 16,229        dim_head: int = 64,230        norm_type: str = "layer_norm",231        norm_affine: bool = True,232        qk_norm_type: str = None,233        qk_norm_affine: bool = False,234        ffn_activation_fn: str = "gelu",235        ffn_use_gated: bool = False,236        rope_theta: float = 100.0,237        rope_dim_ratio: float = 1.0,238        bias: bool = True,239        eps: float = 1e-5,240        num_register_tokens: int = 4,241        mask_config: dict = {},242        **kwargs,243    ):244        super().__init__()245 246        dim = heads * dim_head247        rope_apply_dim = int(dim_head * rope_dim_ratio)248 249        self.pos_embed = RotaryEmbeddingND(rope_apply_dim, rope_theta, n_dim=3, use_angle=True)250 251        self.x_embedder = nn.Linear(in_channels, dim)252 253        self.init_suffix_tokens(dim, num_register_tokens, has_cls_token=False)254 255        self.t_causal = t_causal256 257        self.transformer_blocks = nn.ModuleList(258            [259                TransformerBlock(260                    heads=heads,261                    dim_head=dim_head,262                    norm_type=norm_type,263                    norm_affine=norm_affine,264                    qk_norm_type=qk_norm_type,265                    qk_norm_affine=qk_norm_affine,266                    ffn_activation_fn=ffn_activation_fn,267                    ffn_use_gated=ffn_use_gated,268                    bias=bias,269                    eps=eps,270                    **kwargs,271                )272                for _ in range(num_layers)273            ]274        )275 276        self.spatial_parallel = False277        for block in self.transformer_blocks:278            block.attn.spatial_parallel = False279 280        self.norm_out = nn.LayerNorm(dim, elementwise_affine=norm_affine, eps=eps)281        patch_dim = out_channels * patch_size_t * patch_size * patch_size282        self.proj_out = nn.Linear(dim, patch_dim)283 284        self.init_mask_config(dim, is_3d=True)285        self.set_mask_config(mask_config)286 287        self._init_weights()288        self.gradient_checkpointing = False289 290        if len(kwargs) > 0 and (not dist.is_initialized() or dist.get_rank() == 0):291            logger.warning(f"Unused kwargs: {kwargs}")292 293    def forward(self, x: torch.Tensor) -> torch.Tensor:294        self.loss_info = {}295 296        B, C, latent_T, latent_H, latent_W = x.shape297        patch_size = self.config.patch_size298        patch_size_t = self.config.patch_size_t299        num_suffix = 1 + self.num_register_tokens300 301        hidden_states = _pack_tensors_3d(x, 1, 1)302        latent_size = (latent_T, latent_H, latent_W)303 304        with torch.autocast("cuda", enabled=False):305            hidden_states = _linear_with_module_dtype(self.x_embedder, hidden_states, hidden_states.dtype)306 307        num_patches = hidden_states.shape[1]308 309        tokens = [hidden_states]310 311        if self.register_tokens is not None:312            register_tokens = self.register_tokens.expand(B, -1, -1)313            tokens.append(register_tokens)314 315        cls_token = torch.zeros_like(hidden_states[:, 0:1, :])316        tokens.append(cls_token)317        hidden_states = torch.cat(tokens, dim=1)318 319        patch_dims = [latent_T, latent_H, latent_W]320        img_ids = create_token_ids(latent_size, x.device, x.dtype).expand(B, -1, -1)321        suffix_ids = torch.zeros((B, num_suffix, 3), device=x.device, dtype=img_ids.dtype)322        img_ids = torch.cat([img_ids, suffix_ids], dim=1)323 324        hidden_states, img_ids = self.apply_mask_preprocess(hidden_states, img_ids, patch_dims, num_suffix)325 326        pack_info = {}327        if self.t_causal:328            spatial_size = latent_H * latent_W329            mask_mod = make_block_causal_mask_mod(330                num_tokens=num_patches,331                block_size=spatial_size,332                suffix=True,333            )334            pack_info["mask_mod"] = mask_mod335 336        hidden_states, img_ids, pack_info, sp_pad_len = self._pad_for_sp(hidden_states, img_ids, pack_info)337 338        rotary_pos_emb = self.pos_embed(img_ids)339 340        if self.spatial_parallel:341            hidden_states = get_subseq(hidden_states)342 343        for block in self.transformer_blocks:344            hidden_states = maybe_checkpoint(345                self, block, hidden_states, rotary_pos_emb, pack_info346            )347 348        if self.spatial_parallel:349            hidden_states = gather_subseq(hidden_states)350        hidden_states = self._unpad_for_sp(hidden_states, sp_pad_len)351 352        hidden_states = self.norm_out(hidden_states)353 354        hidden_states = self.apply_mask_postprocess(hidden_states, num_patches)355 356        with torch.autocast("cuda", enabled=False):357            output = _linear_with_module_dtype(self.proj_out, hidden_states, hidden_states.dtype)358 359        output = output[:, :num_patches, :]360 361        video_t = latent_size[0] * patch_size_t362        video_h = latent_size[1] * patch_size363        video_w = latent_size[2] * patch_size364        output = _unpack_tensors_3d(output, patch_size, patch_size_t, video_t, video_h, video_w)365 366        return output367 368 369 370 371 372 373 374 375 376 377 378 379 380 381