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 Euler-eta0 update 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,
) -> 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.