Skip to content

preprocess_minimax_h3_ref2va

Precompute a raw MiniMax H3 Ref2VA manifest into training Parquet shards.

Classes

Functions:

fastvideo.pipelines.preprocess.preprocess_minimax_h3_ref2va.encode_ref2va_conditioning

encode_ref2va_conditioning(caption: str, references: list[MiniMaxH3PreparedReference], model_path: Path, model_index: dict[str, Any], fastvideo_args: FastVideoArgs) -> tuple[Tensor, Tensor]

Encode the exact ordered presentation without padding or truncation.

Source code in fastvideo/pipelines/preprocess/preprocess_minimax_h3_ref2va.py
def encode_ref2va_conditioning(
    caption: str,
    references: list[MiniMaxH3PreparedReference],
    model_path: Path,
    model_index: dict[str, Any],
    fastvideo_args: FastVideoArgs,
) -> tuple[torch.Tensor, torch.Tensor]:
    """Encode the exact ordered presentation without padding or truncation."""
    print("Loading MiniMax H3 tokenizer, processor, and Qwen3-VL encoder")
    tokenizer = _load_component("tokenizer", model_path, model_index, fastvideo_args)
    processor = _load_component("processor", model_path, model_index, fastvideo_args)
    conditioner = _load_component("text_encoder", model_path, model_index, fastvideo_args)
    stage = MiniMaxH3ConditioningStage(
        conditioner=conditioner,
        tokenizer=tokenizer,
        processor=processor,
        ref2va=bool(references),
    )
    batch = ForwardBatch(data_type="video", prompt=caption, references=references)
    if not references:
        # The Ref stage intentionally rejects an empty list. Prompt-only rows
        # use the exact T2VA tokenizer path and contain only text tags.
        batch.extra[MINIMAX_H3_KEYFRAMES_KEY] = []
    batch = stage.forward(batch, fastvideo_args)
    if len(batch.prompt_embeds) != 1:
        raise RuntimeError("MiniMax H3 conditioning must return exactly one embedding")
    text_embedding = batch.prompt_embeds[0].squeeze(0).float().cpu().contiguous()
    text_token_tags = batch.extra.get(MINIMAX_H3_TEXT_TOKEN_TAGS_KEY)
    if not isinstance(text_token_tags, torch.Tensor):
        raise RuntimeError("MiniMax H3 conditioning did not return text token tags")
    text_token_tags = text_token_tags.to(dtype=torch.long, device="cpu").contiguous()
    if text_embedding.ndim != 2 or text_embedding.shape[1] != 5120 or text_embedding.shape[0] == 0:
        raise ValueError(f"Unexpected Qwen embedding shape: {tuple(text_embedding.shape)}")
    if text_token_tags.shape != text_embedding.shape[:1]:
        raise ValueError("Qwen text token tags do not align with its hidden states")
    if not bool(((text_token_tags == 0) | (text_token_tags == 1)).all()):
        raise ValueError("Qwen text token tags may contain only vision=0 and text=1")

    dynamic_length = int(text_embedding.shape[0])
    print(f"Qwen Ref2VA conditioning shape: {tuple(text_embedding.shape)}; preserving all {dynamic_length} tokens")
    del batch, stage, conditioner, processor, tokenizer
    gc.collect()
    torch.cuda.empty_cache()
    return text_embedding, text_token_tags

fastvideo.pipelines.preprocess.preprocess_minimax_h3_ref2va.encode_ref_audio_anchor

encode_ref_audio_anchor(references: list[MiniMaxH3PreparedReference], model_path: Path, model_index: dict[str, Any], fastvideo_args: FastVideoArgs) -> Tensor

Encode clean channel-major audio anchors in ordered-reference order.

Source code in fastvideo/pipelines/preprocess/preprocess_minimax_h3_ref2va.py
def encode_ref_audio_anchor(
    references: list[MiniMaxH3PreparedReference],
    model_path: Path,
    model_index: dict[str, Any],
    fastvideo_args: FastVideoArgs,
) -> torch.Tensor:
    """Encode clean channel-major audio anchors in ordered-reference order."""
    if not any(reference.has_audio for reference in references):
        return torch.empty((0, MINIMAX_H3_REF2VA_AUDIO_ROW_WIDTH), dtype=torch.float32)

    print("Loading MiniMax H3 audio VAE for Ref2VA audio anchors")
    audio_vae = _load_component("audio_vae", model_path, model_index, fastvideo_args)
    device = torch.device("cuda:0")
    latent_channels = int(audio_vae.latent_channels)
    if int(audio_vae.sampling_rate) != AUDIO_SAMPLE_RATE:
        raise ValueError(f"Audio VAE sampling rate must be {AUDIO_SAMPLE_RATE}, got {audio_vae.sampling_rate}")
    rows: list[torch.Tensor] = []
    with torch.no_grad():
        for reference in references:
            if not reference.has_audio:
                continue
            if reference.waveform is None:
                raise ValueError("Audio-bearing reference is missing its prepared waveform")
            posterior = audio_vae.encode(reference.waveform.to(device=device, dtype=torch.float32)[:, None]).latent_dist
            latents = audio_vae.normalize_latents(posterior.mode().float()).cpu().transpose(1, 2)
            if latents.ndim != 3 or latents.shape[0] != 2 or latents.shape[2] != latent_channels:
                raise ValueError(f"Unexpected reference audio latent shape: {tuple(latents.shape)}")
            reference.num_audio_latents = int(latents.shape[1])
            rows.append(latents.reshape(-1, latent_channels).float().contiguous())
            del posterior, latents
    anchor = torch.cat(rows).float().contiguous()
    if latent_channels != MINIMAX_H3_REF2VA_AUDIO_ROW_WIDTH or anchor.shape[1] != latent_channels:
        raise ValueError(
            f"Ref2VA audio anchor width must be {MINIMAX_H3_REF2VA_AUDIO_ROW_WIDTH}, got {anchor.shape[1]}")
    del rows, audio_vae
    gc.collect()
    torch.cuda.empty_cache()
    print(f"Ref2VA audio anchor shape: {tuple(anchor.shape)}")
    return anchor

fastvideo.pipelines.preprocess.preprocess_minimax_h3_ref2va.encode_ref_visual_anchor

encode_ref_visual_anchor(references: list[MiniMaxH3PreparedReference], model_path: Path, model_index: dict[str, Any], fastvideo_args: FastVideoArgs, patch_size: tuple[int, int, int]) -> Tensor

Encode and cache official 0.999-noised ordered visual condition rows.

Source code in fastvideo/pipelines/preprocess/preprocess_minimax_h3_ref2va.py
def encode_ref_visual_anchor(
    references: list[MiniMaxH3PreparedReference],
    model_path: Path,
    model_index: dict[str, Any],
    fastvideo_args: FastVideoArgs,
    patch_size: tuple[int, int, int],
) -> torch.Tensor:
    """Encode and cache official 0.999-noised ordered visual condition rows."""
    visual_references = [reference for reference in references if reference.media_type != "audio"]
    if not visual_references:
        return torch.empty((0, MINIMAX_H3_REF2VA_VISUAL_ROW_WIDTH), dtype=torch.float32)

    print("Loading MiniMax H3 video VAE for Ref2VA visual anchors")
    vae = _load_component("vae", model_path, model_index, fastvideo_args)
    device = torch.device("cuda:0")
    clean_rows: list[torch.Tensor] = []
    latent_channels = int(vae.latent_channels)
    with torch.no_grad():
        for reference in references:
            if reference.media_type == "audio":
                continue
            if reference.media_type == "image":
                if reference.image is None:
                    raise ValueError("Prepared image reference is missing pixels")
                pixels = torch.from_numpy(np.asarray(reference.image).copy()).permute(2, 0, 1)[None, :, None]
                pixels = pixels.to(device=device, dtype=torch.float32).div_(255.0)
                posterior = vae.encode_keyframe(vae.normalize_pixels(pixels)).latent_dist
            else:
                if reference.frames is None:
                    raise ValueError("Prepared video reference is missing frames")
                frames = reference.frames[:trim_reference_num_frames(reference.frames.shape[0])]
                pixels = torch.from_numpy(frames.copy()).permute(3, 0, 1, 2)[None]
                pixels = pixels.to(device=device, dtype=torch.float32).div_(255.0)
                posterior = vae.encode(vae.normalize_pixels(pixels)).latent_dist

            # The fp16 round trip before latent normalization is part of the
            # released Ref2VA condition encoding path.
            latents = vae.normalize_latents(_sample_visual_posterior(posterior).to(torch.float16).float()).cpu()
            if latents.ndim != 5 or latents.shape[0] != 1 or latents.shape[1] != latent_channels:
                raise ValueError(f"Unexpected reference visual latent shape: {tuple(latents.shape)}")
            reference.num_latent_frames = int(latents.shape[2])
            reference.latent_height = int(latents.shape[3])
            reference.latent_width = int(latents.shape[4])
            clean_rows.append(patchify_video_latents(latents, patch_size).float().contiguous())
            del posterior, latents, pixels

        clean_anchor = torch.cat(clean_rows).to(device=device, dtype=torch.float32)
        shapes = tuple((reference.num_latent_frames, reference.latent_height, reference.latent_width)
                       for reference in references if reference.media_type != "audio")
        noise_generator = torch.Generator("cpu").manual_seed(MINIMAX_H3_KEYFRAME_ENCODE_SEED)
        noise = keyframe_condition_noise(
            shapes,
            patch_size,
            latent_channels,
            generator=noise_generator,
            device=device,
            dtype=torch.float32,
        )
        anchor = MiniMaxH3Scheduler(shift=12.0).scale_noise(
            clean_anchor,
            MINIMAX_H3_KEYFRAME_NOISE_AUG,
            noise,
        ).float().cpu().contiguous()

    expected_width = latent_channels * int(np.prod(patch_size))
    if expected_width != MINIMAX_H3_REF2VA_VISUAL_ROW_WIDTH or anchor.shape[1] != expected_width:
        raise ValueError(
            f"Ref2VA visual anchor width must be {MINIMAX_H3_REF2VA_VISUAL_ROW_WIDTH}, got {anchor.shape[1]}")
    del clean_anchor, clean_rows, noise, vae
    gc.collect()
    torch.cuda.empty_cache()
    print(f"Ref2VA visual anchor shape: {tuple(anchor.shape)} at fixed clean-time 0.999")
    return anchor