Skip to content

vllm_omni.diffusion.models.minimax_h3.denoise_loop

MiniMax H3 cfg-distilled full denoise loop.

Per step, the positive presentation is forwarded exactly once. Video and audio target rows chain through the selected solver while visual and audio condition rows stay pinned to their noised step-0 anchors.

MINIMAX_H3_AUDIO_REF_COND_TIMESTEP module-attribute

MINIMAX_H3_AUDIO_REF_COND_TIMESTEP = 1.0

MINIMAX_H3_AUDIO_ROW_WIDTH module-attribute

MINIMAX_H3_AUDIO_ROW_WIDTH = 32

MINIMAX_H3_IMGVID_COND_TIMESTEP module-attribute

MINIMAX_H3_IMGVID_COND_TIMESTEP = 0.999

MINIMAX_H3_VIDEO_ROW_WIDTH module-attribute

MINIMAX_H3_VIDEO_ROW_WIDTH = 96

MiniMaxH3DenoiseBranch

Static per-branch state: packed layout + fixed forward kwargs.

packed is a minimax_h3_packed_sequence(...) result (or equivalent layout dict); text_embeddings is the branch's [text_len, 5120] hidden states; token_tags must already carry any fl2va vision-span overrides.

audio_pos instance-attribute

audio_pos = packed['audio_pos'].view(-1).to(torch.long)

audio_pos_dev instance-attribute

audio_pos_dev = self.audio_pos.to(device)

audio_update_mask instance-attribute

audio_update_mask = (
    packed["audio_update_mask"].view(-1).to(torch.bool)
)

audio_update_mask_dev instance-attribute

audio_update_mask_dev = self.audio_update_mask.to(device)

audio_x_base instance-attribute

audio_x_base = torch.zeros(
    1,
    seq_len,
    MINIMAX_H3_AUDIO_ROW_WIDTH,
    dtype=torch.float32,
    device=device,
)

device instance-attribute

device = device

img_pos instance-attribute

img_pos = packed['img_pos'].view(-1).to(torch.long)

img_pos_dev instance-attribute

img_pos_dev = self.img_pos.to(device)

img_position_ids_dev instance-attribute

img_position_ids_dev = packed["img_position_ids"].to(
    device
)

locked_audio_rows instance-attribute

locked_audio_rows: Tensor | None = None

seq_len instance-attribute

seq_len = seq_len

static_kwargs instance-attribute

static_kwargs: dict[str, Any] = {
    "img_position_ids": self.img_position_ids_dev[None],
    "update_mask": self.update_mask_dev,
    "token_tags": self.token_tags_dev,
    "skip_mask_out_condition": False,
    "prompt_embeds": self.text_embeddings_dev,
    "img_pos_info": {"position_ids": self.img_pos_dev},
    "audio_pos_info": {"position_ids": self.audio_pos_dev},
    "text_pos_info": {"position_ids": self.text_pos_dev},
    "img_pos_for_infer_output_info": {
        "position_ids": self.img_pos_dev
    },
    "packed_seq_params": {
        "cu_seqlens_q": cu.to(device),
        "max_seqlen_q": self.used_len,
        "num_requests": 1,
    },
    "refiner_packed_seq_params": {
        "cu_seqlens_q": torch.tensor(
            [0, text_len, text_len],
            dtype=torch.int32,
            device=device,
        ),
        "max_seqlen_q": text_len,
    },
}

text_embeddings_dev instance-attribute

text_embeddings_dev = text_embeddings.to(device)

text_len instance-attribute

text_len = text_len

text_pos_dev instance-attribute

text_pos_dev = (
    packed["text_pos"].view(-1).to(torch.long).to(device)
)

token_tags_dev instance-attribute

token_tags_dev = (
    token_tags.view(-1).to(torch.long).to(device)
)

update_mask instance-attribute

update_mask = packed["update_mask"].view(-1).to(torch.bool)

update_mask_dev instance-attribute

update_mask_dev = self.update_mask.to(device)

used_len instance-attribute

used_len = int(cu[1])

x_base instance-attribute

x_base = torch.zeros(
    1,
    seq_len,
    MINIMAX_H3_VIDEO_ROW_WIDTH,
    dtype=torch.float32,
    device=device,
)

fill_timesteps

fill_timesteps(
    timesteps: Tensor,
    *,
    t_video: float,
    t_audio: float,
    imgvid_cond_timestep: float,
    audio_ref_cond_timestep: float,
    video_target_timesteps: Tensor | None = None,
    audio_target_timesteps: Tensor | None = None,
) -> None

Fill one request's packed row timesteps in an existing tensor.

forward_kwargs

forward_kwargs(
    *,
    video_rows: Tensor,
    audio_rows: Tensor,
    t_video: float,
    t_audio: float,
    imgvid_cond_timestep: float,
    audio_ref_cond_timestep: float,
    video_target_timesteps: Tensor | None = None,
    audio_target_timesteps: Tensor | None = None,
) -> dict[str, Any]

prepare_rope_table

prepare_rope_table(model: Any) -> None

Materialize the branch-local DiT RoPE table once per denoise run.

img_position_ids is immutable for this branch while latents and timesteps change every scheduler step. Keeping the table in static_kwargs makes every model call reuse the exact same BF16 tensor without extending its lifetime beyond this request branch.

minimax_h3_denoise_loop

minimax_h3_denoise_loop(
    *,
    model: Any,
    positive: MiniMaxH3DenoiseBranch,
    initial_video_rows: Tensor,
    initial_audio_rows: Tensor,
    keyframe_cond_rows: Tensor | None,
    audio_ref_rows: Tensor | None = None,
    video_edit: MiniMaxH3LatentEdit | None = None,
    audio_edit: MiniMaxH3LatentEdit | None = None,
    sigmas_video: list[float],
    sigmas_audio: list[float],
    device: device,
    imgvid_cond_noise_aug_for_inference: float = MINIMAX_H3_IMGVID_COND_TIMESTEP,
    audio_cond_noise_aug_for_inference: float = MINIMAX_H3_AUDIO_REF_COND_TIMESTEP,
    on_step: Callable[[int, Tensor, Tensor], None]
    | None = None,
    step_profiler: Callable[[int], AbstractContextManager]
    | None = None,
    sampler: str = "euler",
) -> tuple[Tensor, Tensor]

Run the full denoise loop; returns final (video_rows, audio_rows).

initial_video_rows covers all image rows of the positive layout. For a conditional task, pass keyframe_cond_rows and/or audio_ref_rows to pin those rows across every step. The model's raw positive velocity is the update signal; MiniMax H3 only supports cfg-distilled checkpoints.

minimax_h3_prepare_denoise_rows

minimax_h3_prepare_denoise_rows(
    *,
    positive: MiniMaxH3DenoiseBranch,
    initial_video_rows: Tensor,
    initial_audio_rows: Tensor,
    keyframe_cond_rows: Tensor | None,
    audio_ref_rows: Tensor | None,
    device: device,
) -> tuple[Tensor, Tensor, Tensor | None, Tensor | None]

Validate the initial rows against the layout and pin the condition rows.

Returns (video_rows, audio_rows, cond_anchor, audio_anchor) with the anchors already written into their rows, which is the state both the request-mode loop and step-mode prepare_encode() start from.

minimax_h3_publish_denoise_progress

minimax_h3_publish_denoise_progress(
    step: int | None,
    sigma_video: float | None,
    total_steps: int | None = None,
) -> None

Publish denoise progress for step-gated attention features.

Both execution modes must publish the same trio: the step index drives the dense warmup of RAINFUSION_ATTN, the normalized descending timestep drives the TRTLLM_ATTN skip gate (which stays dense while it is unset), and the total step count enables the end_step tail fallback.