Team Ai
Modelpublic

MiniMaxAI/MiniMax-H3

sourceHugging Faceotherupdated 2mo agoView on Hugging Face
6klikes3.6mdownloads
vae_processor.py235 linesDownload Raw Back to video_vae
1# SPDX-License-Identifier: Apache-2.02# Tensor pre/post-processing for the MiniMax H3 visual VAE.3import math4import numpy as np5import torch6from diffusers.utils import logging7from einops import rearrange8 9from .normalize import get_normalize_transform, get_denormalize_transform10 11logger = logging.get_logger(__name__)  # pylint: disable=invalid-name12 13 14class VAEProcessor:15 16    def __init__(17        self,18        *,19        vae_ratio,20        vae_ratio_t,21        clip_length,22        frame_overlap,23        token_overlap,24        tokens_chunk_size,25        isolated_last_frame,26        latent_patch_size,27        crop_mode,28        pixel_norm_type="imagenet",29        transform=None,30        transform_rev=None,31        use_3d_conv=False,32    ):33        self.vae_ratio = vae_ratio34        self.vae_ratio_t = vae_ratio_t35        self.clip_length = clip_length36        self.frame_overlap = frame_overlap37        self.token_overlap = token_overlap38        self.tokens_chunk_size = tokens_chunk_size39        self.isolated_last_frame = isolated_last_frame40        self.latent_patch_size = latent_patch_size41        self.crop_mode = crop_mode42        self.transform = transform or get_normalize_transform(pixel_norm_type)43        self.transform_rev = transform_rev or get_denormalize_transform(pixel_norm_type)44        self.use_3d_conv = use_3d_conv45 46    def _ensure_list(self, data):47        return data if isinstance(data, list) else [data]48 49    def _align_to_total_patch_size(self, h, w):50        total_patch_size = self.latent_patch_size * self.vae_ratio51        new_h = (h // total_patch_size) * total_patch_size52        new_w = (w // total_patch_size) * total_patch_size53        return new_h, new_w54 55    def _crop_to_align(self, tensor, new_h, new_w, is_video=False):56        if is_video:57            _, _, _, h, w = tensor.shape58        else:59            _, _, h, w = tensor.shape60 61        if self.crop_mode == "center":62            top = (h - new_h) // 263            left = (w - new_w) // 264        else:65            top = 066            left = 067 68        if is_video:69            return tensor[:, :, :, top : top + new_h, left : left + new_w]70        else:71            return tensor[:, :, top : top + new_h, left : left + new_w]72 73    def _align_target_token(self, T, mode):74        intra_tail = self.clip_length % self.vae_ratio_t75        min_frames = intra_tail or self.vae_ratio_t76        full_chunks = T // self.clip_length77        remainder = T % self.clip_length78 79        if remainder == 0:80            return max(T, min_frames)81 82        if mode == "pad":83            aligned_r = (84                math.ceil((remainder - intra_tail) / self.vae_ratio_t) * self.vae_ratio_t85                + intra_tail86            )87            if aligned_r > self.clip_length:88                return (full_chunks + 1) * self.clip_length + intra_tail89            return full_chunks * self.clip_length + aligned_r90        else:  # trim91            k = (remainder - intra_tail) // self.vae_ratio_t92            if k >= 0:93                target = full_chunks * self.clip_length + k * self.vae_ratio_t + intra_tail94                return max(target, min_frames)95            elif full_chunks > 0:96                return full_chunks * self.clip_length97            else:98                return min_frames99 100    def _align_target(self, T, mode, granularity):101        if granularity == "chunk":102            step = self.clip_length103            tail = self.frame_overlap104            if self.isolated_last_frame:105                tail += 1106 107            k = math.ceil((T - tail) / step) if mode == "pad" else (T - tail) // step108            return max(k, 1) * step + tail109 110        isolated_extra = 1 if self.isolated_last_frame else 0111        return self._align_target_token(T - isolated_extra, mode) + isolated_extra112 113    def align_video_length(self, video_length, mode="pad", granularity="chunk"):114        target = self._align_target(video_length, mode, granularity)115        delta = target - video_length116        if delta > 0 and mode == "trim":117            raise ValueError(118                f"Cannot trim {video_length} frames to valid length {target}: "119                f"not enough frames (granularity={granularity})"120            )121        return delta122 123    def align_video_length_2pass(self, video_length):124        """Return the leading/trailing frame pads and trailing latent drop.125 126        This is the continuation-prefix (2-pass) alignment.  The caller temporarily disables the model's normal token127        drop and keeps these mirrored processor fields at zero.128        """129        if self.isolated_last_frame:130            raise ValueError(131                "align_video_length_2pass does not support isolated_last_frame"132            )133        if self.token_overlap != 0 or self.frame_overlap != 0:134            raise ValueError(135                "align_video_length_2pass requires token_drop=0 alignment"136            )137 138        leading = self.align_video_length(139            video_length, mode="pad", granularity="token"140        )141        token_aligned = video_length + leading142        trailing = self.align_video_length(143            token_aligned, mode="pad", granularity="chunk"144        )145 146        if trailing > 0:147            intra_tail = self.clip_length % self.vae_ratio_t148            full_chunks = token_aligned // self.clip_length149            remainder = token_aligned % self.clip_length150            real_tokens = full_chunks * self.tokens_chunk_size151            if remainder > 0:152                real_tokens += (153                    (remainder - intra_tail) // self.vae_ratio_t + 1154                )155            drop_tokens = (156                self.get_latent_length(token_aligned + trailing) - real_tokens157            )158        else:159            drop_tokens = 0160 161        return leading, trailing, drop_tokens162 163    def get_suitable_video_length(self, video_length, verbose=False):164        used_frame_length = video_length + self.align_video_length(165            video_length, mode="trim", granularity="chunk"166        )167        if verbose:168            logger.info(169                f"Pick first {used_frame_length} frames from {video_length}-frame video"170            )171        return used_frame_length172 173    def get_latent_length(self, video_length):174        tail_frame = self.frame_overlap175        tail_token = self.token_overlap176        if self.isolated_last_frame:177            tail_frame += 1178            tail_token += 1179 180        video_length = self.get_suitable_video_length(video_length)181        latent_length = (182            int((video_length - tail_frame) // self.clip_length)183            * self.tokens_chunk_size184            + tail_token185        )186        return latent_length187 188 189 190    def transform_tensor(self, tensor):191        B, T = None, None192        if tensor.ndim == 5:193            if tensor.shape[2] == 3:194                tensor = tensor.transpose(1, 2)195            B, _, T, _, _ = tensor.shape196            tensor = rearrange(tensor, "b c t h w -> (b t) c h w")197        elif tensor.ndim == 4:198            if tensor.shape[0] == 3:199                tensor = tensor.transpose(0, 1)200        elif tensor.ndim == 3:201            tensor = tensor.unsqueeze(0)202        else:203            raise ValueError(f"Unsupported tensor shape: {tensor.shape}")204 205        tensor = self.transform(tensor)206 207        if B is not None and T is not None:208            tensor = rearrange(tensor, "(b t) c h w -> b c t h w", b=B, t=T)209 210        return tensor.contiguous()211 212    def revert_tensor(self, tensor):213        B, T = None, None214        if self.use_3d_conv:215            tensor = tensor.unsqueeze(2) if tensor.ndim == 4 else tensor216            B, _, T, _, _ = tensor.shape217            tensor = rearrange(tensor, "b c t h w -> (b t) c h w")218        tensor_rev = self.transform_rev(tensor).clamp(0, 1)219        if B is not None:220            tensor_rev = rearrange(tensor_rev, "(b t) c h w -> b c t h w", b=B, t=T)221        return tensor_rev.contiguous()222 223    @staticmethod224    def convert_numpy_to_tensor(numpy_array, device=None):225        if isinstance(numpy_array, list):226            numpy_array = np.stack(numpy_array, axis=0)227        numpy_array = numpy_array.astype(np.float32)228        tensor = torch.from_numpy(numpy_array)229        tensor = tensor.permute(0, 3, 1, 2)230        tensor = tensor / 255.0231        if device is not None:232            tensor = tensor.to(device)233        return tensor234 235