Skip to content

minimax_h3_ref2va

MiniMax H3 joint text-to-video-and-audio training (full-tuning and LoRA tuning) plugin

Classes

fastvideo.train.models.minimax_h3.minimax_h3_ref2va.MiniMaxH3Ref2VALoraModel

MiniMaxH3Ref2VALoraModel(*, init_from: str, training_config: TrainingConfig, trainable: bool = True, disable_custom_init_weights: bool = False, enable_gradient_checkpointing_type: str | None = None, transformer_override_safetensor: str | None = None, lora: LoraConfig | dict[str, Any] | None = None, expected_lora_layers: int | None = None, attention_backend: AttentionBackendEnum | str | None = TORCH_SDPA)

Bases: MiniMaxH3LoraModel, MiniMaxH3Ref2VAModel

Use Ref2VA data semantics with LoRA-only transformer tuning.

Source code in fastvideo/train/models/minimax_h3/minimax_h3.py
def __init__(
    self,
    *,
    init_from: str,
    training_config: TrainingConfig,
    trainable: bool = True,
    disable_custom_init_weights: bool = False,
    enable_gradient_checkpointing_type: str | None = None,
    transformer_override_safetensor: str | None = None,
    lora: LoraConfig | dict[str, Any] | None = None,
    expected_lora_layers: int | None = None,
    attention_backend: AttentionBackendEnum | str | None = AttentionBackendEnum.TORCH_SDPA,
) -> None:
    lora_config = LoraConfig.coerce(lora)
    if lora_config is None or not lora_config.enable:
        raise ValueError("MiniMaxH3LoraModel requires models.student.lora.enable=true")
    if not lora_config.target_modules:
        raise ValueError("MiniMaxH3LoraModel requires a non-empty "
                         "models.student.lora.target_modules list")
    if not trainable:
        raise ValueError("MiniMaxH3LoraModel requires trainable=true")

    expected_layers = (int(expected_lora_layers) if expected_lora_layers is not None else None)
    if expected_layers is not None and expected_layers <= 0:
        raise ValueError("expected_lora_layers must be positive")

    super().__init__(
        init_from=init_from,
        training_config=training_config,
        trainable=trainable,
        disable_custom_init_weights=disable_custom_init_weights,
        enable_gradient_checkpointing_type=enable_gradient_checkpointing_type,
        transformer_override_safetensor=transformer_override_safetensor,
        attention_backend=attention_backend,
        lora=lora_config,
    )

    if (expected_layers is not None and self._num_lora_layers != expected_layers):
        raise ValueError("Unexpected MiniMax H3 LoRA layer count: "
                         f"expected {expected_layers}, got {self._num_lora_layers}")

    unexpected_trainable = [
        name for name, parameter in self.transformer.named_parameters()
        if parameter.requires_grad and name.rsplit(".", maxsplit=1)[-1] not in {"lora_A", "lora_B"}
    ]
    if unexpected_trainable:
        raise ValueError("MiniMax H3 LoRA enabled non-LoRA trainable parameters: "
                         f"{unexpected_trainable[:10]}")

fastvideo.train.models.minimax_h3.minimax_h3_ref2va.MiniMaxH3Ref2VAModel

MiniMaxH3Ref2VAModel(*, init_from: str, training_config: TrainingConfig, trainable: bool = True, disable_custom_init_weights: bool = False, enable_gradient_checkpointing_type: str | None = None, transformer_override_safetensor: str | None = None, attention_backend: AttentionBackendEnum | str | None = TORCH_SDPA, lora: LoraConfig | dict[str, Any] | None = None)

Bases: MiniMaxH3Model

Adapt the Ref2VA transformer partition to target-only joint flow loss.

Source code in fastvideo/train/models/minimax_h3/minimax_h3.py
def __init__(
    self,
    *,
    init_from: str,
    training_config: TrainingConfig,
    trainable: bool = True,
    disable_custom_init_weights: bool = False,
    enable_gradient_checkpointing_type: str | None = None,
    transformer_override_safetensor: str | None = None,
    attention_backend: AttentionBackendEnum | str | None = AttentionBackendEnum.TORCH_SDPA,
    lora: LoraConfig | dict[str, Any] | None = None,
) -> None:
    """Validate the single-document T2VA contract and load the transformer."""
    super().__init__(
        trainable=trainable,
        lora=lora,
        attention_backend=attention_backend,
    )
    # PyTorch scaled dot product attention (SDPA) provides dense attention
    # without adding another attention-kernel dependency to H3 training.
    if self.attention_backend != AttentionBackendEnum.TORCH_SDPA:
        raise ValueError("MiniMaxH3Model requires the TORCH_SDPA attention backend")
    if training_config.pipeline_config is None:
        raise ValueError("MiniMaxH3Model requires a resolved MiniMax H3 pipeline config")
    # Packed row indices describe one text-video-audio document without a
    # batch offset, so each data-parallel replica consumes one sample.
    if int(training_config.data.train_batch_size) != 1:
        raise ValueError("MiniMaxH3Model requires training.data.train_batch_size=1")
    # Classifier-free guidance (CFG) dropout replaces text embeddings with
    # zeros, but H3 training does not define a zero-vector branch.
    if float(training_config.data.training_cfg_rate) != 0.0:
        raise ValueError("MiniMaxH3Model requires training.data.training_cfg_rate=0.0")
    # Joint supervision requires paired video and stereo-audio latents from
    # every parquet row.
    if str(training_config.data.preprocessed_data_type) != "t2va":
        raise ValueError("MiniMaxH3Model requires training.data.preprocessed_data_type='t2va'")

    # FastVideo's Fully Sharded Data Parallel loading path requires one BF16
    # parameter dtype, including modules that H3 inference keeps in FP32.
    training_config.pipeline_config.dit_config.uniform_parameter_dtype = True  # type: ignore[attr-defined]

    self._init_from = str(init_from)
    self.training_config = training_config
    self.transformer = self._load_transformer(
        trainable=trainable,
        disable_custom_init_weights=disable_custom_init_weights,
        enable_gradient_checkpointing_type=resolve_checkpointing_type(
            enable_gradient_checkpointing_type,
            training_config,
        ),
        transformer_override_safetensor=transformer_override_safetensor,
    )
    self.noise_scheduler = MiniMaxH3Scheduler(shift=_VIDEO_SCHEDULER_SHIFT)
    self.audio_noise_scheduler = MiniMaxH3Scheduler(shift=_AUDIO_SCHEDULER_SHIFT)
    self.dataloader: Any = None
    self.validator: Any = None
    self.start_step = 0
    self.sp_group: Any = None

Methods:

fastvideo.train.models.minimax_h3.minimax_h3_ref2va.MiniMaxH3Ref2VAModel.init_preprocessors
init_preprocessors(training_config: TrainingConfig) -> None

Build the batch-one loader without truncating Ref2VA Qwen tokens.

Source code in fastvideo/train/models/minimax_h3/minimax_h3_ref2va.py
def init_preprocessors(self, training_config: TrainingConfig) -> None:
    """Build the batch-one loader without truncating Ref2VA Qwen tokens."""
    self.sp_group = get_sp_group()
    _dataset, self.dataloader = build_minimax_h3_ref2va_dataloader(
        training_config.data.data_path,
        int(training_config.data.train_batch_size),
        int(training_config.data.dataloader_num_workers),
        drop_last=True,
        seed=int(training_config.data.seed or 0),
    )
    self.start_step = 0
fastvideo.train.models.minimax_h3.minimax_h3_ref2va.MiniMaxH3Ref2VAModel.predict_noise
predict_noise(noisy_latents: Tensor, timestep: Tensor, batch: TrainingBatch, *, conditional: bool, cfg_uncond: dict[str, Any] | None = None, attn_kind: Literal['dense', 'vsa'] = 'dense') -> NoisePrediction

Prepend fixed reference rows and return only target video/audio flow.

Source code in fastvideo/train/models/minimax_h3/minimax_h3_ref2va.py
def predict_noise(
    self,
    noisy_latents: torch.Tensor,
    timestep: torch.Tensor,
    batch: TrainingBatch,
    *,
    conditional: bool,
    cfg_uncond: dict[str, Any] | None = None,
    attn_kind: Literal["dense", "vsa"] = "dense",
) -> NoisePrediction:
    """Prepend fixed reference rows and return only target video/audio flow."""
    del timestep
    if not conditional or cfg_uncond is not None:
        raise ValueError("MiniMaxH3Ref2VAModel predicts one conditional sample")
    if attn_kind != "dense":
        raise ValueError("MiniMaxH3Ref2VAModel supports dense attention for training")
    layout = batch.minimax_h3_layout
    if not isinstance(layout, MiniMaxH3PackedLayout):
        raise RuntimeError("prepare_batch() must set TrainingBatch.minimax_h3_layout")
    if batch.audio_noisy_model_input is None or batch.encoder_hidden_states is None:
        raise RuntimeError("prepare_batch() must set audio and text transformer inputs")
    if batch.timesteps is None or batch.audio_timesteps is None:
        raise RuntimeError("prepare_batch() must set video and audio timesteps")
    extras = batch.input_kwargs
    if not isinstance(extras, dict):
        raise RuntimeError("prepare_batch() must set Ref2VA anchor inputs")
    visual_anchor = extras.get(_REF_VISUAL_ANCHOR_KEY)
    audio_anchor = extras.get(_REF_AUDIO_ANCHOR_KEY)
    if not isinstance(visual_anchor, torch.Tensor) or not isinstance(audio_anchor, torch.Tensor):
        raise RuntimeError("prepare_batch() must set both Ref2VA anchor tensors")

    dtype = torch.bfloat16
    device = self.device
    target_video_bcthw = noisy_latents.permute(0, 2, 1, 3, 4).to(device=device, dtype=dtype)
    target_video_rows = patchify_video_latents(target_video_bcthw, tuple(self.transformer.patch_size))
    target_audio = batch.audio_noisy_model_input.to(device=device, dtype=dtype)
    num_target_audio_latents = int(target_audio.shape[-1])
    target_audio_rows = target_audio.permute(0, 1, 3, 2).reshape(
        -1,
        MINIMAX_H3_REF2VA_AUDIO_ROW_WIDTH,
    )
    if target_video_rows.shape[1] != MINIMAX_H3_REF2VA_VISUAL_ROW_WIDTH:
        raise ValueError(f"Unexpected target video row width: {target_video_rows.shape[1]}")

    video_rows = torch.cat((visual_anchor.to(device=device, dtype=dtype), target_video_rows), dim=0)
    audio_rows = torch.cat((audio_anchor.to(device=device, dtype=dtype), target_audio_rows), dim=0)
    if video_rows.shape[0] != layout.video_indices.numel():
        raise ValueError("Packed Ref2VA video row count does not match its layout")
    if audio_rows.shape[0] != layout.audio_indices.numel():
        raise ValueError("Packed Ref2VA audio row count does not match its layout")

    video_timestep = float(batch.timesteps[0].item())
    audio_timestep = float(batch.audio_timesteps[0].item())
    unique_timesteps, timestep_indices = build_row_timesteps(
        layout,
        video_timestep=video_timestep,
        audio_timestep=audio_timestep,
        condition_video_timestep=max(video_timestep, MINIMAX_H3_KEYFRAME_NOISE_AUG),
        condition_audio_timestep=1.0,
    )
    unique_timesteps = unique_timesteps.to(device)
    timestep_indices = timestep_indices.to(device)

    with torch.autocast(device.type, dtype=dtype), set_forward_context(
            current_timestep=unique_timesteps,
            attn_metadata=None,
    ):
        video_velocity, audio_velocity = self.transformer(
            hidden_states=video_rows[None],
            audio_hidden_states=audio_rows[None],
            encoder_hidden_states=batch.encoder_hidden_states,
            timestep=unique_timesteps,
            timestep_indices=timestep_indices,
            token_tags=layout.token_tags.to(device),
            position_ids=layout.position_ids.to(device),
            video_indices=layout.video_indices.to(device),
            audio_indices=layout.audio_indices.to(device),
            text_indices=layout.text_indices.to(device),
        )

    if video_velocity.ndim != 3 or video_velocity.shape[1] != video_rows.shape[0]:
        raise ValueError(f"Unexpected Ref2VA video output shape: {tuple(video_velocity.shape)}")
    if audio_velocity.ndim != 3 or audio_velocity.shape[1] != audio_rows.shape[0]:
        raise ValueError(f"Unexpected Ref2VA audio output shape: {tuple(audio_velocity.shape)}")
    target_video_velocity = video_velocity[:, layout.num_condition_video_rows:]
    target_audio_velocity = audio_velocity[:, layout.num_condition_audio_rows:]

    _, channels, num_video_latents, latent_height, latent_width = target_video_bcthw.shape
    video_prediction = unpatchify_video_tokens(
        target_video_velocity,
        num_video_latents,
        latent_height,
        latent_width,
        channels,
        tuple(self.transformer.patch_size),
    ).permute(0, 2, 1, 3, 4)
    audio_prediction = unpack_audio_tokens(target_audio_velocity[0], num_target_audio_latents)[None]
    return -video_prediction, -audio_prediction
fastvideo.train.models.minimax_h3.minimax_h3_ref2va.MiniMaxH3Ref2VAModel.prepare_batch
prepare_batch(raw_batch: dict[str, Any], *, generator: Generator, latents_source: Literal['data', 'zeros'] = 'data') -> TrainingBatch

Reuse target noise preparation, then install ordered Ref2VA conditions.

Source code in fastvideo/train/models/minimax_h3/minimax_h3_ref2va.py
def prepare_batch(
    self,
    raw_batch: dict[str, Any],
    *,
    generator: torch.Generator,
    latents_source: Literal["data", "zeros"] = "data",
) -> TrainingBatch:
    """Reuse target noise preparation, then install ordered Ref2VA conditions."""
    batch = super().prepare_batch(
        raw_batch,
        generator=generator,
        latents_source=latents_source,
    )
    references = _restore_prepared_references(batch.infos)
    text_token_tags = _valid_text_token_tags(raw_batch)
    if batch.encoder_hidden_states is None or batch.encoder_hidden_states.shape[1] != text_token_tags.numel():
        raise ValueError("Filtered text_token_tags do not align with encoder_hidden_states")

    if batch.raw_latent_shape is None or len(batch.raw_latent_shape) != 5:
        raise RuntimeError("Parent MiniMax H3 preparation did not preserve target video geometry")
    if batch.audio_latents is None:
        raise RuntimeError("Parent MiniMax H3 preparation did not preserve target audio latents")
    _, _, num_video_latents, latent_height, latent_width = batch.raw_latent_shape
    num_audio_latents = int(batch.audio_latents.shape[-1])
    patch_size = tuple(self.transformer.patch_size)

    if references:
        layout = build_ref2va_packed_sequence(
            text_token_tags,
            references,
            num_video_latents,
            latent_height,
            latent_width,
            num_audio_latents,
            patch_size,
        )
    else:
        if bool((text_token_tags != MINIMAX_H3_TEXT_TAG).any()):
            raise ValueError("A zero-reference batch cannot contain visual Qwen token tags")
        layout = build_packed_sequence(
            text_token_tags,
            num_video_latents,
            latent_height,
            latent_width,
            num_audio_latents,
            patch_size,
        )

    visual_anchor = _unbatch_anchor(
        raw_batch,
        _REF_VISUAL_ANCHOR_KEY,
        MINIMAX_H3_REF2VA_VISUAL_ROW_WIDTH,
        self.device,
    )
    audio_anchor = _unbatch_anchor(
        raw_batch,
        _REF_AUDIO_ANCHOR_KEY,
        MINIMAX_H3_REF2VA_AUDIO_ROW_WIDTH,
        self.device,
    )
    if visual_anchor.shape[0] != layout.num_condition_video_rows:
        raise ValueError("ref_visual_anchor row count does not match the ordered reference layout: "
                         f"{visual_anchor.shape[0]} != {layout.num_condition_video_rows}")
    if audio_anchor.shape[0] != layout.num_condition_audio_rows:
        raise ValueError("ref_audio_anchor row count does not match the ordered reference layout: "
                         f"{audio_anchor.shape[0]} != {layout.num_condition_audio_rows}")

    batch.minimax_h3_layout = layout
    batch.input_kwargs = dict(batch.input_kwargs or {})
    batch.input_kwargs[_REF_VISUAL_ANCHOR_KEY] = visual_anchor
    batch.input_kwargs[_REF_AUDIO_ANCHOR_KEY] = audio_anchor
    return batch

Functions: