Compare commits
39
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0a861031c7 | ||
|
|
f08c5ee8af | ||
|
|
755f4a7967 | ||
|
|
d543a67b10 | ||
|
|
741aa8d289 | ||
|
|
ad9cd63122 | ||
|
|
68e6ffca9e | ||
|
|
64cdcf6be4 | ||
|
|
aa0d98a6b8 | ||
|
|
93b03bc14d | ||
|
|
160f0c9ccf | ||
|
|
e5d1110a0f | ||
|
|
622217ff2a | ||
|
|
cbab605eff | ||
|
|
ac98869aa1 | ||
|
|
9713ea1275 | ||
|
|
37aa382cce | ||
|
|
3f00983287 | ||
|
|
2dc57f4070 | ||
|
|
9df19be719 | ||
|
|
089eea3970 | ||
|
|
dd8447ecc5 | ||
|
|
dca423fd31 | ||
|
|
1b43af8e8e | ||
|
|
628591b620 | ||
|
|
0980ca563f | ||
|
|
b158388733 | ||
|
|
ac56806aff | ||
|
|
aa95a4c18e | ||
|
|
942f7db3db | ||
|
|
aadb23f409 | ||
|
|
f56f567042 | ||
|
|
15a164a052 | ||
|
|
aaaa7a14a3 | ||
|
|
74b409d7cf | ||
|
|
528cef02c4 | ||
|
|
3c3da4d057 | ||
|
|
0399713e7b | ||
|
|
0af2e9e8ef |
@@ -76,6 +76,10 @@ surfaces:
|
||||
lora_target_modules: "Legacy LoRA configuration surface pending dedicated component API."
|
||||
output_type: "Legacy output formatting surface pending GenerationResult cleanup."
|
||||
VSA_sparsity: "Model-specific inference optimization not yet represented in the typed public schema."
|
||||
VSA_tile_size: "VSA-H3 tile geometry request; model-specific optimization not yet represented in the typed public schema."
|
||||
vae_parallel_decode: "MiniMax-H3 sequence-parallel VAE decode opt-in; model-specific optimization not yet represented in the typed public schema."
|
||||
vae_parallel_encode: "MiniMax-H3 sequence-parallel reference VAE encode opt-in; model-specific optimization not yet represented in the typed public schema."
|
||||
vae_parallel_decode_strategy: "Chunk-transport collective for vae_parallel_decode; model-specific optimization not yet represented in the typed public schema."
|
||||
attention_backend: "Process-wide default attention-backend request applied per component at load time; kernel-selection knob not yet represented in the typed public schema."
|
||||
moba_config_path: "Model-specific MoBA optimization surface not yet represented in the typed public schema."
|
||||
master_port: "Executor/bootstrap compatibility field; not part of the canonical inference schema."
|
||||
|
||||
@@ -33,6 +33,14 @@ For the typed config/request path added during the inference API refactor:
|
||||
python examples/inference/basic/basic_dmd_new_api.py
|
||||
```
|
||||
|
||||
For the few-step (4-step, DMD2-distilled) MiniMax-H3 preview, generating synchronized video and audio, optionally with block-sparse VSA attention:
|
||||
```
|
||||
python examples/inference/basic/basic_fasth3.py --prompt "your prompt" [--vsa-sparsity 0.9]
|
||||
```
|
||||
The default checkpoint `FastVideo/FastVideo-Minimax-FastH3-Preview-v0.1` is private on the Hub while its license review completes; until it flips public, pass `--model-path` with a local snapshot of the release.
|
||||
|
||||
On Blackwell (sm_100) GPUs with a `fastvideo-kernel` build that carries the sm_100a block-sparse extension, `--vsa-kernel sm100a` routes the tile-64 attention forwards through the CUDA kernel instead of Triton (it sets `FASTVIDEO_VSA_SM100A=1` before the pipeline boots); if the extension or the arch is missing, the run warns once and falls back to Triton.
|
||||
|
||||
## Basic Walkthrough
|
||||
|
||||
All you need to generate videos using multi-gpus from state-of-the-art diffusion pipelines is the following few lines!
|
||||
|
||||
@@ -0,0 +1,180 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Few-step video+audio generation with the DMD2-distilled MiniMax H3 preview.
|
||||
|
||||
FastVideo/FastVideo-Minimax-FastH3-Preview-v0.1 is a 4-step distillation of
|
||||
MiniMaxAI/MiniMax-H3 (data-free DMD2): it walks a 4-step grid on the release
|
||||
sampler's shift-12 schedule instead of the base model's 50 steps, generating
|
||||
synchronized video and audio in one pipeline call.
|
||||
|
||||
The student was trained with block-sparse video attention (VSA, 64-token
|
||||
tiles) and its checkpoint carries the trained sparse-gate parameters
|
||||
(``attn.to_gate_compress``), so this script always runs the VSA-H3 attention
|
||||
backend. At the default ``--vsa-sparsity 0.0`` the attention math is exactly
|
||||
dense (every tile is selected); raise the sparsity for additional speedup.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.api import (
|
||||
CompileConfig,
|
||||
EngineConfig,
|
||||
GenerationRequest,
|
||||
GeneratorConfig,
|
||||
OffloadConfig,
|
||||
OutputConfig,
|
||||
ParallelismConfig,
|
||||
PipelineSelection,
|
||||
SamplingConfig,
|
||||
)
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--model-path", default="FastVideo/FastVideo-Minimax-FastH3-Preview-v0.1")
|
||||
# The HF repo is private while the MiniMax H3 Community License review
|
||||
# completes; until it flips public, pass --model-path with a local
|
||||
# snapshot of the release instead (e.g. the team export at
|
||||
# /mnt/lustre/vlm-wlsaidhi/fastvideo/exports/FastVideo-Minimax-FastH3-Preview-v0.1).
|
||||
parser.add_argument("--prompt", required=True)
|
||||
parser.add_argument("--output", default="outputs/fasth3")
|
||||
parser.add_argument("--height", type=int, default=768)
|
||||
parser.add_argument("--width", type=int, default=1344)
|
||||
parser.add_argument("--num-frames", type=int, default=124)
|
||||
# num_inference_steps counts sigma-GRID POINTS, matching the base model's
|
||||
# convention ("50 steps" = a 50-point grid = 49 transformer forwards). The
|
||||
# student's distilled 4-step grid is 4 FORWARDS, i.e. a 5-point grid
|
||||
# (t = 1000, 750, 500, 250 -> 0 on the shift-12 schedule) — so the correct
|
||||
# default here is 5. Other grids are off-distribution.
|
||||
parser.add_argument("--steps",
|
||||
type=int,
|
||||
default=5,
|
||||
help="num_inference_steps = sigma-grid points; N points run N-1 denoising "
|
||||
"forwards. 5 (default) is the distilled 4-forward grid")
|
||||
parser.add_argument("--seed", type=int, default=0)
|
||||
parser.add_argument("--num-gpus", type=int, default=4)
|
||||
parser.add_argument("--vsa-sparsity",
|
||||
type=float,
|
||||
default=0.0,
|
||||
help="Run-level VSA sparsity in [0, 1). 0.0 (default) selects every tile, which is "
|
||||
"exactly dense attention; the student was trained at 0.9")
|
||||
# 64 is the trained contract: the student was TRAINED with 64-token
|
||||
# (4,4,4) tiles, and its to_gate_compress gates were learned against
|
||||
# pooling at that granularity — keep 64 unless you are ablating.
|
||||
parser.add_argument("--vsa-tile-size",
|
||||
type=int,
|
||||
choices=(64, 256),
|
||||
default=64,
|
||||
help="VSA-H3 tile size in tokens; 64 (default) is what the student was trained "
|
||||
"with and runs the native Triton block-sparse path, 256 is the FA4-CuTe-capable "
|
||||
"geometry for ablations")
|
||||
parser.add_argument("--vsa-kernel",
|
||||
choices=("triton", "sm100a"),
|
||||
default="triton",
|
||||
help="Block-sparse kernel for the tile-64 attention forward: triton (default, "
|
||||
"portable fwd+bwd) or sm100a — the opt-in Blackwell CUDA forward "
|
||||
"(fastvideo_kernel.block_sparse_attn_sm100a). sm100a needs an sm_100 GPU and a "
|
||||
"fastvideo-kernel build that carries the extension; if a precondition fails at "
|
||||
"run time the attention layer logs one warning and falls back to Triton. Only "
|
||||
"meaningful with --vsa-tile-size 64")
|
||||
parser.add_argument("--torch-compile", action="store_true", help="torch.compile the DiT transformer path")
|
||||
parser.add_argument("--compile-mode",
|
||||
default=None,
|
||||
help='torch.compile mode, e.g. "reduce-overhead" for CUDA graphs')
|
||||
parser.add_argument("--repeats",
|
||||
type=int,
|
||||
default=1,
|
||||
help="generate N times; with --torch-compile the first run pays "
|
||||
"compilation, so steady-state is the last repeat")
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
output_dir = Path(args.output)
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
if args.vsa_kernel == "sm100a":
|
||||
# The attention backend reads FASTVIDEO_VSA_SM100A per forward; set it
|
||||
# before the pipeline boots so spawned GPU workers inherit it. The
|
||||
# kernel is forward-only and inference runs under no-grad, so every
|
||||
# denoising forward qualifies for the CUDA route.
|
||||
os.environ["FASTVIDEO_VSA_SM100A"] = "1"
|
||||
|
||||
# Boot-time run configuration, folded into FastVideoArgs (the same route
|
||||
# examples/inference/basic/basic_minimax_h3_t2v.py uses for sparsity):
|
||||
# - attention_backend: the checkpoint carries trained to_gate_compress
|
||||
# gates, which only exist under the VSA-H3 backend — a dense-backend
|
||||
# load would reject them as unexpected weights. Layers that do not
|
||||
# support VSA-H3 (e.g. the token refiner) fall back to flash attention.
|
||||
# - VSA_tile_size: forwarded even at sparsity 0.0 because the gate-compress
|
||||
# branch pools per tile, and the gates were trained at 64 tokens/tile.
|
||||
experimental: dict[str, object] = {
|
||||
"attention_backend": "VIDEO_SPARSE_ATTN_H3",
|
||||
"VSA_tile_size": args.vsa_tile_size,
|
||||
}
|
||||
if args.vsa_sparsity > 0.0:
|
||||
experimental["VSA_sparsity"] = args.vsa_sparsity
|
||||
|
||||
generator = VideoGenerator.from_config(
|
||||
GeneratorConfig(
|
||||
model_path=args.model_path,
|
||||
pipeline=PipelineSelection(experimental=experimental),
|
||||
engine=EngineConfig(
|
||||
num_gpus=args.num_gpus,
|
||||
use_fsdp_inference=args.num_gpus > 1,
|
||||
parallelism=ParallelismConfig(tp_size=1, sp_size=args.num_gpus),
|
||||
offload=OffloadConfig(
|
||||
dit=False,
|
||||
dit_layerwise=False,
|
||||
text_encoder=True,
|
||||
vae=True,
|
||||
pin_cpu_memory=False,
|
||||
),
|
||||
compile=CompileConfig(
|
||||
enabled=args.torch_compile,
|
||||
mode=args.compile_mode,
|
||||
),
|
||||
),
|
||||
))
|
||||
try:
|
||||
request = GenerationRequest(
|
||||
prompt=args.prompt,
|
||||
negative_prompt="",
|
||||
sampling=SamplingConfig(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
fps=24,
|
||||
num_inference_steps=args.steps,
|
||||
# the base model is guidance-distilled; the student inherits it
|
||||
guidance_scale=1.0,
|
||||
batch_cfg=False,
|
||||
seed=args.seed,
|
||||
),
|
||||
output=OutputConfig(
|
||||
output_path=str(output_dir / "fasth3.mp4"),
|
||||
save_video=True,
|
||||
return_frames=False,
|
||||
),
|
||||
)
|
||||
result = generator.generate(request)
|
||||
print(f"Output written to: {result.video_path}")
|
||||
if result.generation_time is not None:
|
||||
# machine-readable: benchmark harnesses parse this line to separate
|
||||
# generation from model-load time (last occurrence = steady state)
|
||||
print(f"Generation time: {result.generation_time:.2f}s")
|
||||
for _ in range(args.repeats - 1):
|
||||
result = generator.generate(request)
|
||||
if result.generation_time is not None:
|
||||
print(f"Generation time: {result.generation_time:.2f}s")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -19,7 +19,12 @@ from fastvideo.attention.backends.abstract import (
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
logger.info("Using FlashAttention-%s backend", fa_version)
|
||||
# Every worker records the loaded FlashAttention implementation so a
|
||||
# distributed profiling log contains one backend receipt per rank.
|
||||
logger.info("Worker %s Using FlashAttention-%s backend",
|
||||
os.environ.get("RANK", "0"),
|
||||
fa_version,
|
||||
local_main_process_only=False)
|
||||
|
||||
# FP4 FA4 support: quantize Q/K to NVFP4 E2M1 for block-scaled MMA on Blackwell.
|
||||
# Requires: flash-attention-fp4, flashinfer, cutlass-dsl. Enable via nvfp4_fa4=True kwarg.
|
||||
|
||||
@@ -5,13 +5,17 @@ H3 runs one joint bidirectional attention over
|
||||
``[text | condition keyframes | audio | generated video]``, so this
|
||||
backend differs from the Wan-tuned ``video_sparse_attn``:
|
||||
|
||||
- Tiles are ``[segment-pure prefix chunks] + [3D (4,8,8) video tiles]``;
|
||||
prefix tiles never straddle segment boundaries.
|
||||
- Tiles are ``[segment-pure prefix chunks] + [3D video tiles]``; prefix
|
||||
tiles never straddle segment boundaries. The tile size is selectable at
|
||||
metadata build time: 256 tokens ``(4,8,8)`` (default) or 64 tokens
|
||||
``(4,4,4)`` (see ``VSA_H3_TILE_SHAPES``).
|
||||
- Selection is pure Python on pooled tile scores; the block-sparse kernel
|
||||
consumes an explicit bool mask, so no kernel changes are needed.
|
||||
- The compression branch is gated by ``to_gate_compress``, which the H3
|
||||
checkpoint does not carry: the loader zero-initializes it, so untrained
|
||||
inference is exactly pure sparse and finetuning can learn the gate.
|
||||
- The compression branch is gated by ``to_gate_compress``, which the base
|
||||
H3 checkpoint does not carry: the loader zero-initializes it, so
|
||||
untrained inference is exactly pure sparse and finetuning can learn the
|
||||
gate. VSA-distilled students (e.g. FastVideo-Minimax-H3-Preview) ship
|
||||
trained gates, which load and activate the branch.
|
||||
- Non-video *queries* are always dense. Non-video *keys* are either
|
||||
always-selected for every query ("exempt", default) or compete in
|
||||
top-k under a FLOP-matched budget ("compete") — the ablation axis,
|
||||
@@ -20,22 +24,46 @@ backend differs from the Wan-tuned ``video_sparse_attn``:
|
||||
(``vsa_dense_first_n_steps``, ``vsa_dense_layers``) let mixed schedules
|
||||
run the diffuse steps/layers dense while pushing the rest harder.
|
||||
|
||||
Targets sm10.x through the FA4 CuTe 256-tile path
|
||||
At tile 256 this targets sm10.x through the FA4 CuTe 256-tile path
|
||||
(``FASTVIDEO_VSA_CUTEDSL=1``); the Triton 256→64 expansion is the
|
||||
fallback and keeps identical mask semantics.
|
||||
fallback and keeps identical mask semantics. At tile 64 the block map is
|
||||
already at the kernels' native 64-token granularity, so both forward and
|
||||
backward run the Triton block-sparse kernels directly (no expansion,
|
||||
``FASTVIDEO_VSA_CUTEDSL`` does not apply). A third, opt-in route exists
|
||||
for the tile-64 FORWARD only: ``FASTVIDEO_VSA_SM100A=1`` sends no-grad
|
||||
forwards through the sm_100a CUDA block-sparse kernel
|
||||
(``fastvideo_kernel.block_sparse_attn_sm100a``, upstream PR #1719 plus
|
||||
our per-q-tile ``q2k_num`` fix) when the extension is built, the device
|
||||
is sm_100, and the geometry qualifies; grad-tracking forwards and every
|
||||
backward stay on Triton unchanged. If the env is set but a precondition
|
||||
fails, the route logs one warning and falls back.
|
||||
"""
|
||||
|
||||
import functools
|
||||
import math
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
try:
|
||||
from fastvideo_kernel.block_sparse_attn import block_sparse_attn as block_sparse_attn_64_bhsd
|
||||
from fastvideo_kernel.block_sparse_attn_256 import block_sparse_attn_256_bshd
|
||||
from fastvideo_kernel.triton_kernels.index import map_to_index
|
||||
except ImportError:
|
||||
block_sparse_attn_64_bhsd = None
|
||||
block_sparse_attn_256_bshd = None
|
||||
map_to_index = None
|
||||
|
||||
try:
|
||||
# Optional: only present in fastvideo_kernel builds that carry the sm_100a
|
||||
# CUDA block-sparse forward (upstream PR #1719). The module itself imports
|
||||
# fine without the compiled symbols (`_HAS_VSA_SM100A` is then False and
|
||||
# `is_supported` says no), so this only guards *module* availability.
|
||||
from fastvideo_kernel import block_sparse_attn_sm100a as _sm100a
|
||||
except ImportError:
|
||||
_sm100a = None
|
||||
|
||||
from fastvideo.attention.backends.abstract import (AttentionBackend, AttentionImpl, AttentionMetadata,
|
||||
AttentionMetadataBuilder, layer_idx_from_prefix)
|
||||
@@ -43,51 +71,115 @@ from fastvideo.attention.backends.video_sparse_attn import (compute_topk, constr
|
||||
get_non_pad_index, get_tile_partition_indices,
|
||||
scatter_into_tile_buf)
|
||||
from fastvideo.attention.backends.video_sparse_attn_h3_probe import probe_enabled, record_probe
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
VSA_H3_TILE_SIZE = (4, 8, 8) # 256 elements -> FA4 CuTe fastpath on sm10.x
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Opt-in switch for the sm_100a CUDA forward on the tile-64 no-grad path.
|
||||
VSA_SM100A_ENV = "FASTVIDEO_VSA_SM100A"
|
||||
|
||||
VSA_H3_TILE_SIZE = (4, 8, 8) # 256 elements -> FA4 CuTe fastpath on sm10.x (default)
|
||||
_TILE_ELEMS = math.prod(VSA_H3_TILE_SIZE)
|
||||
# Selectable tile geometries, keyed by element count (= the build-time
|
||||
# ``tile_size``). 64 runs the native 64-token Triton block-sparse kernels for
|
||||
# forward AND backward — the block map is already at kernel granularity, so no
|
||||
# 256->64 mask expansion is involved and FASTVIDEO_VSA_CUTEDSL does not apply.
|
||||
VSA_H3_TILE_SHAPES: dict[int, tuple[int, int, int]] = {
|
||||
_TILE_ELEMS: VSA_H3_TILE_SIZE,
|
||||
64: (4, 4, 4),
|
||||
}
|
||||
|
||||
|
||||
def token_tile_and_valid(variable_block_sizes: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
def token_tile_and_valid(variable_block_sizes: torch.Tensor,
|
||||
tile_elems: int = _TILE_ELEMS) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Per padded-token tile id and pad-validity mask.
|
||||
|
||||
The single encoding of the padding contract, shared by the probe and the
|
||||
test oracle so they cannot drift from the backend's tile geometry.
|
||||
``tile_elems`` must match the metadata the sizes came from
|
||||
(``MiniMaxH3VSAMetadata.tile_elems``).
|
||||
"""
|
||||
device = variable_block_sizes.device
|
||||
token_tile = torch.arange(variable_block_sizes.numel(), device=device).repeat_interleave(_TILE_ELEMS)
|
||||
token_valid = (torch.arange(_TILE_ELEMS, device=device)[None, :] < variable_block_sizes[:, None]).reshape(-1)
|
||||
token_tile = torch.arange(variable_block_sizes.numel(), device=device).repeat_interleave(tile_elems)
|
||||
token_valid = (torch.arange(tile_elems, device=device)[None, :] < variable_block_sizes[:, None]).reshape(-1)
|
||||
return token_tile, token_valid
|
||||
|
||||
|
||||
def _validate_h3_tile_geometry(
|
||||
prefix_segments: tuple[int, ...],
|
||||
dit_seq_shape: tuple[int, int, int],
|
||||
variable_block_sizes: torch.Tensor,
|
||||
untile_combined_index: torch.Tensor,
|
||||
tile_elems: int = _TILE_ELEMS,
|
||||
) -> None:
|
||||
"""Fail synchronously on out-of-bounds tile geometry.
|
||||
|
||||
Invariants the block-sparse kernel trusts without checking:
|
||||
every tile's valid size is in (0, tile_elems]; the sizes sum to the
|
||||
packed sequence length; and ``untile_combined_index`` maps each packed
|
||||
row to exactly one non-pad slot of the padded tile buffer. A violation
|
||||
would surface only as an async device fault at some later kernel or
|
||||
collective (e.g. an FSDP all-gather), which is unattributable — so raise
|
||||
here, once per cached geometry, with the numbers in hand.
|
||||
"""
|
||||
total = sum(prefix_segments) + math.prod(dit_seq_shape)
|
||||
n_pad = variable_block_sizes.numel() * tile_elems
|
||||
sizes_min = int(variable_block_sizes.min())
|
||||
sizes_max = int(variable_block_sizes.max())
|
||||
sizes_sum = int(variable_block_sizes.sum())
|
||||
if sizes_min < 1 or sizes_max > tile_elems or sizes_sum != total:
|
||||
raise ValueError(f"VSA-H3 tile sizes out of bounds for prefix={prefix_segments}, video={dit_seq_shape}, "
|
||||
f"tile_elems={tile_elems}: min={sizes_min}, max={sizes_max}, sum={sizes_sum}, "
|
||||
f"expected sum={total}.")
|
||||
if untile_combined_index.numel() != total:
|
||||
raise ValueError(f"VSA-H3 untile index has {untile_combined_index.numel()} entries for a packed "
|
||||
f"sequence of {total} rows (prefix={prefix_segments}, video={dit_seq_shape}).")
|
||||
idx_min = int(untile_combined_index.min())
|
||||
idx_max = int(untile_combined_index.max())
|
||||
if idx_min < 0 or idx_max >= n_pad:
|
||||
# Range first: the pad-slot gather below would itself index out of
|
||||
# bounds (the very async fault this guard exists to preempt).
|
||||
raise ValueError(f"VSA-H3 untile index is not an injective map into non-pad slots: range "
|
||||
f"[{idx_min}, {idx_max}] vs padded length {n_pad} "
|
||||
f"(prefix={prefix_segments}, video={dit_seq_shape}).")
|
||||
in_tile_offset = untile_combined_index % tile_elems
|
||||
maps_into_pad = bool((in_tile_offset >= variable_block_sizes[untile_combined_index // tile_elems]).any())
|
||||
if maps_into_pad or int(torch.unique(untile_combined_index).numel()) != total:
|
||||
raise ValueError(f"VSA-H3 untile index is not an injective map into non-pad slots: "
|
||||
f"pad-slot hit={maps_into_pad} "
|
||||
f"(prefix={prefix_segments}, video={dit_seq_shape}).")
|
||||
|
||||
|
||||
@functools.lru_cache(maxsize=10)
|
||||
def _h3_tile_geometry(
|
||||
prefix_segments: tuple[int, ...],
|
||||
dit_seq_shape: tuple[int, int, int],
|
||||
device: torch.device,
|
||||
tile_shape: tuple[int, int, int] = VSA_H3_TILE_SIZE,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, int, int]:
|
||||
"""Tile the packed sequence: segment-pure prefix chunks, then video tiles.
|
||||
|
||||
Returns (tile_partition_indices, variable_block_sizes,
|
||||
untile_combined_index, num_prefix_tiles, num_video_tiles).
|
||||
"""
|
||||
tile_elems = math.prod(tile_shape)
|
||||
prefix_len = sum(prefix_segments)
|
||||
|
||||
prefix_sizes: list[int] = []
|
||||
for segment in prefix_segments:
|
||||
full, rem = divmod(segment, _TILE_ELEMS)
|
||||
prefix_sizes.extend([_TILE_ELEMS] * full)
|
||||
full, rem = divmod(segment, tile_elems)
|
||||
prefix_sizes.extend([tile_elems] * full)
|
||||
if rem:
|
||||
prefix_sizes.append(rem)
|
||||
num_prefix_tiles = len(prefix_sizes)
|
||||
|
||||
ts_t, ts_h, ts_w = VSA_H3_TILE_SIZE
|
||||
ts_t, ts_h, ts_w = tile_shape
|
||||
t, h, w = dit_seq_shape
|
||||
num_tiles = (math.ceil(t / ts_t), math.ceil(h / ts_h), math.ceil(w / ts_w))
|
||||
video_sizes = construct_variable_block_sizes(dit_seq_shape, num_tiles, device, VSA_H3_TILE_SIZE)
|
||||
video_sizes = construct_variable_block_sizes(dit_seq_shape, num_tiles, device, tile_shape)
|
||||
num_video_tiles = int(video_sizes.numel())
|
||||
|
||||
video_indices = get_tile_partition_indices(dit_seq_shape, VSA_H3_TILE_SIZE, device) + prefix_len
|
||||
video_indices = get_tile_partition_indices(dit_seq_shape, tile_shape, device) + prefix_len
|
||||
tile_partition_indices = torch.cat([
|
||||
torch.arange(prefix_len, device=device, dtype=torch.long),
|
||||
video_indices,
|
||||
@@ -100,9 +192,11 @@ def _h3_tile_geometry(
|
||||
|
||||
# get_non_pad_index is lru-cached on tensor identity; variable_block_sizes
|
||||
# is itself cached by this function, so the identity stays stable.
|
||||
non_pad_index = get_non_pad_index(variable_block_sizes, _TILE_ELEMS)
|
||||
non_pad_index = get_non_pad_index(variable_block_sizes, tile_elems)
|
||||
|
||||
untile_combined_index = non_pad_index[torch.argsort(tile_partition_indices)]
|
||||
# One-time (lru-cached) synchronous bounds check; see _validate_h3_tile_geometry.
|
||||
_validate_h3_tile_geometry(prefix_segments, dit_seq_shape, variable_block_sizes, untile_combined_index, tile_elems)
|
||||
return (tile_partition_indices, variable_block_sizes, untile_combined_index, num_prefix_tiles, num_video_tiles)
|
||||
|
||||
|
||||
@@ -139,6 +233,9 @@ class MiniMaxH3VSAMetadata(AttentionMetadata):
|
||||
exempt: bool
|
||||
variable_block_sizes: torch.Tensor
|
||||
untile_combined_index: torch.Tensor
|
||||
# tokens per tile (256 or 64); selects the tile geometry AND the kernel
|
||||
# route in forward() (256 -> VSA-256 CuTe/Triton, 64 -> native Triton)
|
||||
tile_elems: int = _TILE_ELEMS
|
||||
# layers forced dense regardless of sparsity (probe-guided opt-outs)
|
||||
dense_layers: tuple[int, ...] = ()
|
||||
# Single-slot holder for the padded tile buffer, owned by the BUILDER so
|
||||
@@ -158,24 +255,28 @@ class MiniMaxH3VSAMetadataBuilder(AttentionMetadataBuilder):
|
||||
pass
|
||||
|
||||
def build( # type: ignore
|
||||
self,
|
||||
current_timestep: int,
|
||||
raw_latent_shape: tuple[int, int, int],
|
||||
patch_size: tuple[int, int, int],
|
||||
VSA_sparsity: float,
|
||||
prefix_segments: tuple[int, ...],
|
||||
device: torch.device,
|
||||
exempt: bool = True,
|
||||
dense_layers: tuple[int, ...] = (),
|
||||
**kwargs: dict[str, Any],
|
||||
self,
|
||||
current_timestep: int,
|
||||
raw_latent_shape: tuple[int, int, int],
|
||||
patch_size: tuple[int, int, int],
|
||||
VSA_sparsity: float,
|
||||
prefix_segments: tuple[int, ...],
|
||||
device: torch.device,
|
||||
exempt: bool = True,
|
||||
dense_layers: tuple[int, ...] = (),
|
||||
tile_size: int = _TILE_ELEMS,
|
||||
**kwargs: dict[str, Any],
|
||||
) -> MiniMaxH3VSAMetadata:
|
||||
tile_shape = VSA_H3_TILE_SHAPES.get(int(tile_size))
|
||||
if tile_shape is None:
|
||||
raise ValueError(f"VSA-H3 tile_size must be one of {sorted(VSA_H3_TILE_SHAPES)}, got {tile_size!r}")
|
||||
dit_seq_shape = (raw_latent_shape[0] // patch_size[0], raw_latent_shape[1] // patch_size[1],
|
||||
raw_latent_shape[2] // patch_size[2])
|
||||
prefix_segments = tuple(int(s) for s in prefix_segments if s > 0)
|
||||
total_seq_length = sum(prefix_segments) + math.prod(dit_seq_shape)
|
||||
|
||||
(_tile_partition_indices, variable_block_sizes, untile_combined_index, num_prefix_tiles,
|
||||
num_video_tiles) = _h3_tile_geometry(prefix_segments, dit_seq_shape, device)
|
||||
num_video_tiles) = _h3_tile_geometry(prefix_segments, dit_seq_shape, device, tile_shape)
|
||||
|
||||
return MiniMaxH3VSAMetadata(
|
||||
current_timestep=current_timestep,
|
||||
@@ -186,13 +287,14 @@ class MiniMaxH3VSAMetadataBuilder(AttentionMetadataBuilder):
|
||||
exempt=exempt,
|
||||
variable_block_sizes=variable_block_sizes,
|
||||
untile_combined_index=untile_combined_index,
|
||||
tile_elems=int(tile_size),
|
||||
dense_layers=tuple(int(layer) for layer in dense_layers),
|
||||
tile_buf_holder=self._tile_buf_holder,
|
||||
)
|
||||
|
||||
|
||||
def _pool_tiles(x: torch.Tensor, variable_block_sizes: torch.Tensor) -> torch.Tensor:
|
||||
"""fp32 mean over each 256-token tile. x: [B, S_pad, H, D] -> [B, H, n_tiles, D].
|
||||
def _pool_tiles(x: torch.Tensor, variable_block_sizes: torch.Tensor, tile_elems: int = _TILE_ELEMS) -> torch.Tensor:
|
||||
"""fp32 mean over each tile_elems-token tile. x: [B, S_pad, H, D] -> [B, H, n_tiles, D].
|
||||
|
||||
Pad positions in the tile buffer are guaranteed zero (zeros-init, never
|
||||
written), so a plain sum with fp32 accumulation needs no validity mask
|
||||
@@ -200,8 +302,8 @@ def _pool_tiles(x: torch.Tensor, variable_block_sizes: torch.Tensor) -> torch.Te
|
||||
the masked mean exactly.
|
||||
"""
|
||||
batch, seq_len, heads, dim = x.shape
|
||||
n_tiles = seq_len // _TILE_ELEMS
|
||||
pooled = x.view(batch, n_tiles, _TILE_ELEMS, heads, dim).sum(dim=2, dtype=torch.float32)
|
||||
n_tiles = seq_len // tile_elems
|
||||
pooled = x.view(batch, n_tiles, tile_elems, heads, dim).sum(dim=2, dtype=torch.float32)
|
||||
pooled = pooled / variable_block_sizes.view(1, -1, 1, 1)
|
||||
return pooled.permute(0, 2, 1, 3)
|
||||
|
||||
@@ -232,6 +334,24 @@ def _build_block_mask(
|
||||
return mask
|
||||
|
||||
|
||||
def _sm100a_unavailable_reason(sm100a_mod: Any, query_bhsd: torch.Tensor, variable_block_sizes: torch.Tensor,
|
||||
grad_mode: bool) -> str | None:
|
||||
"""Why the opt-in sm_100a forward route cannot run here, or None if it can.
|
||||
|
||||
Pure decision logic, split out so the routing is unit-testable without a
|
||||
GPU or the compiled extension (tests substitute ``sm100a_mod``). Order
|
||||
matters only for the message: the cheapest, most actionable reason first.
|
||||
"""
|
||||
if sm100a_mod is None:
|
||||
return "fastvideo_kernel.block_sparse_attn_sm100a is not installed"
|
||||
if grad_mode:
|
||||
return "inputs require grad and the sm_100a kernel is forward-only; grad paths keep Triton"
|
||||
if not sm100a_mod.is_supported(query_bhsd, variable_block_sizes):
|
||||
return ("block_sparse_attn_sm100a.is_supported returned False (needs an sm_100 device, a built "
|
||||
"extension, bf16, head_dim 128, an even tile count, and integer tile sizes)")
|
||||
return None
|
||||
|
||||
|
||||
class MiniMaxH3VSAImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
@@ -259,7 +379,7 @@ class MiniMaxH3VSAImpl(AttentionImpl):
|
||||
f"got {x.shape[1]}. A non-packed sequence (e.g. the token refiner) is "
|
||||
"routed to the VSA-H3 backend; exclude it from the supported backends.")
|
||||
n_tiles = attn_metadata.variable_block_sizes.numel()
|
||||
target_shape = (x.shape[0], n_tiles * _TILE_ELEMS, x.shape[-2], x.shape[-1])
|
||||
target_shape = (x.shape[0], n_tiles * attn_metadata.tile_elems, x.shape[-2], x.shape[-1])
|
||||
|
||||
# single scatter: untile_combined_index maps original row i to its
|
||||
# padded slot, so this is exactly the inverse of postprocess_output
|
||||
@@ -281,7 +401,11 @@ class MiniMaxH3VSAImpl(AttentionImpl):
|
||||
gate_compress: torch.Tensor | None,
|
||||
attn_metadata: MiniMaxH3VSAMetadata,
|
||||
) -> torch.Tensor:
|
||||
if block_sparse_attn_256_bshd is None:
|
||||
tile_elems = attn_metadata.tile_elems
|
||||
if tile_elems == 64:
|
||||
if block_sparse_attn_64_bhsd is None:
|
||||
raise NotImplementedError("fastvideo_kernel.block_sparse_attn is not installed")
|
||||
elif block_sparse_attn_256_bshd is None:
|
||||
raise NotImplementedError("fastvideo_kernel.block_sparse_attn_256 is not installed")
|
||||
|
||||
# probe-guided per-layer opt-out: diffuse layers run dense (all-True
|
||||
@@ -291,8 +415,8 @@ class MiniMaxH3VSAImpl(AttentionImpl):
|
||||
|
||||
scores = None
|
||||
if layer_sparsity > 0.0 or gate_compress is not None or probe_dir is not None:
|
||||
q_pooled = _pool_tiles(query, attn_metadata.variable_block_sizes)
|
||||
k_pooled = _pool_tiles(key, attn_metadata.variable_block_sizes)
|
||||
q_pooled = _pool_tiles(query, attn_metadata.variable_block_sizes, tile_elems)
|
||||
k_pooled = _pool_tiles(key, attn_metadata.variable_block_sizes, tile_elems)
|
||||
scores = torch.matmul(q_pooled, k_pooled.transpose(-2, -1)) / (query.shape[-1]**0.5)
|
||||
if probe_dir is not None:
|
||||
record_probe(probe_dir, self.layer_idx, query, key, scores, attn_metadata)
|
||||
@@ -309,14 +433,66 @@ class MiniMaxH3VSAImpl(AttentionImpl):
|
||||
attn_metadata.exempt,
|
||||
)
|
||||
|
||||
out, _ = block_sparse_attn_256_bshd(query, key, value, mask, attn_metadata.variable_block_sizes)
|
||||
if tile_elems == 64:
|
||||
# Native 64-token path: the block map is already at the kernels'
|
||||
# granularity. Both 64-token entries take BHSD ([B, H, S_pad, D]);
|
||||
# mirror block_sparse_attn_256_bshd's Triton branch and transpose
|
||||
# around the call.
|
||||
q_bhsd = query.transpose(1, 2).contiguous()
|
||||
k_bhsd = key.transpose(1, 2).contiguous()
|
||||
v_bhsd = value.transpose(1, 2).contiguous()
|
||||
|
||||
# Opt-in sm_100a CUDA forward (upstream PR #1719 + per-q-tile
|
||||
# q2k_num fix). Forward-only: grad-tracking calls stay on Triton
|
||||
# so autograd keeps the Triton fwd+bwd pairing untouched. The
|
||||
# kernel does return an LSE in Triton's M format, so a future
|
||||
# fwd/bwd pairing is possible, but it is not built here.
|
||||
use_sm100a = False
|
||||
if os.environ.get(VSA_SM100A_ENV, "0") == "1":
|
||||
grad_mode = torch.is_grad_enabled() and (query.requires_grad or key.requires_grad
|
||||
or value.requires_grad)
|
||||
reason = _sm100a_unavailable_reason(_sm100a, q_bhsd, attn_metadata.variable_block_sizes, grad_mode)
|
||||
if reason is None and map_to_index is None:
|
||||
reason = "fastvideo_kernel.triton_kernels.index (map_to_index) is not importable"
|
||||
if reason is None:
|
||||
use_sm100a = True
|
||||
elif not torch.compiler.is_compiling():
|
||||
logger.warning_once(f"{VSA_SM100A_ENV}=1 but falling back to the Triton-64 kernels: {reason}")
|
||||
|
||||
if use_sm100a:
|
||||
# The sm_100a entry is index-native; compact the bool map the
|
||||
# same way the Triton bool entry does internally. Per-row
|
||||
# counts are NON-uniform here (prefix query tiles are dense,
|
||||
# video tiles run prefix+top-k) -- legal for the fixed kernel,
|
||||
# silently wrong on the pre-fix upstream one.
|
||||
q2k_idx, q2k_num = map_to_index(mask)
|
||||
out_bhsd, _ = _sm100a.block_sparse_attn_sm100a(
|
||||
q_bhsd,
|
||||
k_bhsd,
|
||||
v_bhsd,
|
||||
q2k_idx,
|
||||
q2k_num,
|
||||
attn_metadata.variable_block_sizes.to(torch.int32),
|
||||
need_lse=False,
|
||||
)
|
||||
else:
|
||||
out_bhsd, _ = block_sparse_attn_64_bhsd(
|
||||
q_bhsd,
|
||||
k_bhsd,
|
||||
v_bhsd,
|
||||
mask,
|
||||
attn_metadata.variable_block_sizes,
|
||||
)
|
||||
out = out_bhsd.transpose(1, 2).contiguous()
|
||||
else:
|
||||
out, _ = block_sparse_attn_256_bshd(query, key, value, mask, attn_metadata.variable_block_sizes)
|
||||
|
||||
if gate_compress is not None:
|
||||
# Wan-style compression branch: dense attention over pooled tiles,
|
||||
# broadcast to each tile's rows, scaled by the learned gate
|
||||
# (zero-initialized for H3 => branch contributes nothing until
|
||||
# finetuned; the model layer skips it entirely for all-zero gates).
|
||||
v_pooled = _pool_tiles(value, attn_metadata.variable_block_sizes)
|
||||
v_pooled = _pool_tiles(value, attn_metadata.variable_block_sizes, tile_elems)
|
||||
out_c = torch.matmul(torch.softmax(scores, dim=-1), v_pooled) # [B, H, n_tiles, D]
|
||||
out_c = out_c.permute(0, 2, 1, 3).to(out.dtype) # [B, n_tiles, H, D]
|
||||
batch, seq_len, heads, dim = out.shape
|
||||
@@ -325,7 +501,7 @@ class MiniMaxH3VSAImpl(AttentionImpl):
|
||||
# autograd node saved for its backward, so an in-place add here
|
||||
# bumps its version counter and backward dies with "one of the
|
||||
# variables needed for gradient computation has been modified".
|
||||
out_tiled = out.view(batch, n_tiles, _TILE_ELEMS, heads, dim)
|
||||
gate_tiled = gate_compress.view(batch, n_tiles, _TILE_ELEMS, heads, dim)
|
||||
out_tiled = out.view(batch, n_tiles, tile_elems, heads, dim)
|
||||
gate_tiled = gate_compress.view(batch, n_tiles, tile_elems, heads, dim)
|
||||
out = (out_tiled + out_c.unsqueeze(2) * gate_tiled).view(batch, seq_len, heads, dim)
|
||||
return out
|
||||
|
||||
@@ -60,7 +60,7 @@ def record_probe(
|
||||
gen = torch.Generator(device="cpu").manual_seed(step * 1000 + layer)
|
||||
# sample among video rows in the PADDED/tiled domain that are non-pad
|
||||
from fastvideo.attention.backends.video_sparse_attn_h3 import token_tile_and_valid
|
||||
token_tile, token_valid = token_tile_and_valid(attn_metadata.variable_block_sizes)
|
||||
token_tile, token_valid = token_tile_and_valid(attn_metadata.variable_block_sizes, attn_metadata.tile_elems)
|
||||
video_rows = torch.nonzero((token_tile >= P) & token_valid, as_tuple=False).flatten()
|
||||
idx = video_rows[torch.randint(0, video_rows.numel(), (_TRUE_ROWS, ), generator=gen).to(query.device)]
|
||||
|
||||
|
||||
@@ -62,14 +62,7 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
|
||||
hidden_size: int = 5120
|
||||
intermediate_size: int = 25600
|
||||
num_hidden_layers: int = 64
|
||||
# H3 conditions on one intermediate hidden state and reads nothing above it,
|
||||
# so the remaining layers are built, weight-loaded and then discarded: 14
|
||||
# layers, 13.7 GB in bf16. Building exactly this many leaves that hidden
|
||||
# state bit-identical, because the tuple records each layer's *input*, so
|
||||
# entry N is the output of layer N-1. Set to None to keep the full stack.
|
||||
# Must equal MINIMAX_H3_TEXT_ENCODER_LAYER in
|
||||
# fastvideo/pipelines/basic/minimax_h3/packing.py; a test pins them together
|
||||
# rather than importing across the models -> pipelines boundary.
|
||||
output_hidden_state_index: int = 50
|
||||
num_hidden_layers_override: int | None = 50
|
||||
num_attention_heads: int = 64
|
||||
num_key_value_heads: int = 8
|
||||
@@ -116,7 +109,7 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
|
||||
vision_initializer_range: float = 0.02
|
||||
vision_deepstack_visual_indexes: tuple[int, ...] = (8, 16, 24)
|
||||
|
||||
output_hidden_states: bool = True
|
||||
output_hidden_states: bool = False
|
||||
stacked_params_mapping: list[tuple[str, str, str | int]] = field(default_factory=list)
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [
|
||||
_is_language_transformer_layer,
|
||||
@@ -127,15 +120,16 @@ class MiniMaxH3Qwen3VLArchConfig(TextEncoderArchConfig):
|
||||
])
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
# Runs both at construction and after ``update_model_arch`` merges the
|
||||
# checkpoint's config.json, so it also guards config-file overrides. A
|
||||
# non-positive override would build no decoder layers at all, and a
|
||||
# negative one would additionally make the surplus-key filter drop
|
||||
# every ``language_model.layers.*`` checkpoint key, so the conditioner
|
||||
# would "load" with no transformer stack and only fail at generation.
|
||||
if self.num_hidden_layers_override is not None and self.num_hidden_layers_override < 1:
|
||||
raise ValueError("MiniMax H3 Qwen3-VL num_hidden_layers_override must be a positive layer count "
|
||||
f"or None for the full stack; got {self.num_hidden_layers_override}.")
|
||||
if self.output_hidden_state_index <= 0 or self.output_hidden_state_index > self.num_hidden_layers:
|
||||
raise ValueError("MiniMax H3 Qwen3-VL output_hidden_state_index must be in "
|
||||
f"[1, {self.num_hidden_layers}], got {self.output_hidden_state_index}.")
|
||||
if self.num_hidden_layers_override is not None:
|
||||
if self.num_hidden_layers_override <= 0:
|
||||
raise ValueError("MiniMax H3 Qwen3-VL num_hidden_layers_override must be positive or None.")
|
||||
if self.num_hidden_layers_override < self.output_hidden_state_index:
|
||||
raise ValueError("MiniMax H3 Qwen3-VL num_hidden_layers_override must build through "
|
||||
f"hidden_states[{self.output_hidden_state_index}], got "
|
||||
f"{self.num_hidden_layers_override}.")
|
||||
|
||||
rope_scaling = dict(self.rope_scaling or {})
|
||||
self.mrope_interleaved = bool(rope_scaling.get("mrope_interleaved", self.mrope_interleaved))
|
||||
|
||||
@@ -21,12 +21,17 @@ if TYPE_CHECKING:
|
||||
FASTVIDEO_TRACE_FUNCTION: int = 0
|
||||
FASTVIDEO_ATTENTION_BACKEND: str | None = None
|
||||
FASTVIDEO_FA4: bool = False
|
||||
FASTVIDEO_MINIMAX_H3_FUSIONS: str = ""
|
||||
FASTVIDEO_VAE_PARALLEL_DECODE: bool = False
|
||||
FASTVIDEO_VAE_PARALLEL_ENCODE: bool = False
|
||||
FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY: str | None = None
|
||||
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "spawn"
|
||||
FASTVIDEO_TARGET_DEVICE: str = "cuda"
|
||||
MAX_JOBS: str | None = None
|
||||
NVCC_THREADS: str | None = None
|
||||
CMAKE_BUILD_TYPE: str | None = None
|
||||
VERBOSE: bool = False
|
||||
FASTVIDEO_NVTX_PROFILE: bool = False
|
||||
FASTVIDEO_TORCH_PROFILER_DIR: str | None = None
|
||||
FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES: bool = False
|
||||
FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY: bool = False
|
||||
@@ -217,10 +222,34 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
"FASTVIDEO_FA4":
|
||||
lambda: os.getenv("FASTVIDEO_FA4", "0") != "0",
|
||||
|
||||
# If set (=1), MiniMax-H3 VAE decode (and, with the ENCODE variant,
|
||||
# reference-video encode) round-robins its temporal chunks across the
|
||||
# sequence-parallel ranks instead of running serially on the output rank.
|
||||
# Folded into FastVideoArgs.vae_parallel_decode / vae_parallel_encode at
|
||||
# construction (parse-once). The STRATEGY variant picks the chunk
|
||||
# transport collective: "gather" (default) or "all_gather".
|
||||
"FASTVIDEO_VAE_PARALLEL_DECODE":
|
||||
lambda: os.getenv("FASTVIDEO_VAE_PARALLEL_DECODE", "0") != "0",
|
||||
"FASTVIDEO_VAE_PARALLEL_ENCODE":
|
||||
lambda: os.getenv("FASTVIDEO_VAE_PARALLEL_ENCODE", "0") != "0",
|
||||
"FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY":
|
||||
lambda: os.getenv("FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY", None),
|
||||
|
||||
# Opt-in MiniMax-H3 inference-only Triton fusions adapted from the
|
||||
# NVlabs/Sana Sol-Engine implementation. Accepts `all`, `1`, or a
|
||||
# comma-separated subset of `modulate,qknorm_rope,swiglu`. An empty value
|
||||
# (the default), `0`, or `none` keeps the eager implementation.
|
||||
"FASTVIDEO_MINIMAX_H3_FUSIONS":
|
||||
lambda: os.getenv("FASTVIDEO_MINIMAX_H3_FUSIONS", ""),
|
||||
|
||||
# Use dedicated multiprocess context for workers.
|
||||
"FASTVIDEO_WORKER_MULTIPROC_METHOD":
|
||||
lambda: os.getenv("FASTVIDEO_WORKER_MULTIPROC_METHOD", "spawn"),
|
||||
|
||||
# Emit lightweight NVTX ranges for external profilers such as Nsight Systems.
|
||||
"FASTVIDEO_NVTX_PROFILE":
|
||||
lambda: os.getenv("FASTVIDEO_NVTX_PROFILE", "0") != "0",
|
||||
|
||||
# Enables torch profiler if set. Path to the directory where torch profiler
|
||||
# traces are saved. Note that it must be an absolute path.
|
||||
"FASTVIDEO_TORCH_PROFILER_DIR":
|
||||
|
||||
@@ -146,6 +146,19 @@ class FastVideoArgs:
|
||||
vae_cpu_offload: bool = True
|
||||
pin_cpu_memory: bool = True
|
||||
|
||||
# Sequence-parallel MiniMax-H3 VAE (opt-in, default off). With SP > 1 the
|
||||
# video VAE's temporal chunks (decode) and clips (reference encode) are
|
||||
# round-robined across the sequence-parallel ranks and reassembled
|
||||
# bit-exactly on the group's first rank instead of running serially on
|
||||
# one rank while the others idle. ``__post_init__`` folds the
|
||||
# FASTVIDEO_VAE_PARALLEL_DECODE / FASTVIDEO_VAE_PARALLEL_ENCODE env vars
|
||||
# into these fields (parse-once, like attention_backend), and
|
||||
# FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY overrides the chunk-transport
|
||||
# collective ("gather" or "all_gather").
|
||||
vae_parallel_decode: bool = False
|
||||
vae_parallel_encode: bool = False
|
||||
vae_parallel_decode_strategy: str | None = None
|
||||
|
||||
# Compilation
|
||||
# ``enable_torch_compile`` covers the DiT path (transformer,
|
||||
# transformer_2, and the LTX-2 stage-2 transformer_refine).
|
||||
@@ -169,6 +182,7 @@ class FastVideoArgs:
|
||||
|
||||
# VSA parameters
|
||||
VSA_sparsity: float = 0.0 # inference/validation sparsity
|
||||
VSA_tile_size: int = 256 # VSA-H3 tile size (256 or 64); 64 = native Triton path
|
||||
|
||||
# V-MoBA parameters
|
||||
moba_config_path: str | None = None
|
||||
@@ -286,8 +300,27 @@ class FastVideoArgs:
|
||||
env_backend = envs.FASTVIDEO_ATTENTION_BACKEND
|
||||
if env_backend is not None and backend_name_to_enum(env_backend) is not None:
|
||||
self.attention_backend = env_backend
|
||||
self._fold_vae_parallel_env()
|
||||
self.check_fastvideo_args()
|
||||
|
||||
def _fold_vae_parallel_env(self) -> None:
|
||||
"""Parse-once adapters for the sequence-parallel VAE env vars."""
|
||||
import fastvideo.envs as envs
|
||||
|
||||
# Mirrors fastvideo.models.vaes.minimax_h3_parallel.DECODE_GATHER_STRATEGIES /
|
||||
# DEFAULT_DECODE_GATHER_STRATEGY (kept literal here so constructing args
|
||||
# never imports model modules; a unit test pins the two in sync).
|
||||
strategies = ("gather", "all_gather")
|
||||
if not self.vae_parallel_decode and envs.FASTVIDEO_VAE_PARALLEL_DECODE:
|
||||
self.vae_parallel_decode = True
|
||||
if not self.vae_parallel_encode and envs.FASTVIDEO_VAE_PARALLEL_ENCODE:
|
||||
self.vae_parallel_encode = True
|
||||
if self.vae_parallel_decode_strategy is None:
|
||||
self.vae_parallel_decode_strategy = envs.FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY or "gather"
|
||||
if self.vae_parallel_decode_strategy not in strategies:
|
||||
raise ValueError(f"vae_parallel_decode_strategy must be one of {strategies}, "
|
||||
f"got {self.vae_parallel_decode_strategy!r}.")
|
||||
|
||||
def _apply_transformer_quant(self) -> None:
|
||||
"""Pin the typed ``transformer_quant`` instance onto ``dit_config``.
|
||||
|
||||
@@ -631,6 +664,18 @@ class FastVideoArgs:
|
||||
"Pin memory for CPU offload. Only added as a temp workaround if it throws \"CUDA error: invalid argument\". "
|
||||
"Should be enabled in almost all cases",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vae-parallel-decode",
|
||||
action=StoreBoolean,
|
||||
help="With sequence parallelism, round-robin MiniMax-H3 VAE decode chunks across the SP ranks "
|
||||
"and reassemble bit-exactly on the output rank (default: serial decode on the output rank)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vae-parallel-encode",
|
||||
action=StoreBoolean,
|
||||
help="With sequence parallelism, round-robin MiniMax-H3 reference-video VAE encode clips across "
|
||||
"the SP ranks; every rank keeps the identical full encoding (default: serial encode on every rank)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--disable-autocast",
|
||||
action=StoreBoolean,
|
||||
@@ -644,6 +689,12 @@ class FastVideoArgs:
|
||||
default=FastVideoArgs.VSA_sparsity,
|
||||
help="Validation sparsity for VSA",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--VSA-tile-size",
|
||||
type=int,
|
||||
default=FastVideoArgs.VSA_tile_size,
|
||||
help="VSA-H3 tile size in tokens (256 or 64); 64 runs the native Triton block-sparse path",
|
||||
)
|
||||
|
||||
# Master port for distributed training/inference
|
||||
parser.add_argument(
|
||||
|
||||
+4
-1
@@ -114,7 +114,10 @@ def _info(logger: Logger,
|
||||
is_local_main_process = local_rank == 0
|
||||
|
||||
if (main_process_only and is_main_process) or (local_main_process_only and is_local_main_process):
|
||||
logger.log(logging.INFO, msg, *args, stacklevel=2, **kwargs)
|
||||
# Honor an explicit stacklevel (info_once routes through here with
|
||||
# stacklevel already set) instead of passing the keyword twice.
|
||||
stacklevel = kwargs.pop("stacklevel", 2)
|
||||
logger.log(logging.INFO, msg, *args, stacklevel=stacklevel, **kwargs)
|
||||
|
||||
global _warned_local_main_process, _warned_main_process
|
||||
|
||||
|
||||
@@ -10,6 +10,7 @@ import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo import envs
|
||||
from fastvideo.attention import DistributedAttention
|
||||
from fastvideo.attention.layer import DistributedAttention_VSA
|
||||
from fastvideo.attention.selector import get_attn_backend
|
||||
@@ -23,12 +24,50 @@ from fastvideo.layers.linear import ReplicatedLinear
|
||||
from fastvideo.layers.mlp import MLP
|
||||
from fastvideo.layers.quantization import QuantizationConfig
|
||||
from fastvideo.layers.visual_embedding import Timesteps
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.dits.base import BaseDiT
|
||||
from fastvideo.models.dits.minimax_h3_fusions import (
|
||||
HAVE_TRITON,
|
||||
fused_qknorm_rope,
|
||||
fused_residual_gate_rmsnorm_modulate,
|
||||
fused_rmsnorm_modulate,
|
||||
minimax_h3_swiglu,
|
||||
)
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.profiler import nvtx_range
|
||||
from fastvideo.utils import get_compute_dtype
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
MINIMAX_H3_MODALITY_NUM = 3
|
||||
_CFG = MiniMaxH3Config()
|
||||
_MINIMAX_H3_FUSION_NAMES = frozenset({"modulate", "qknorm_rope", "swiglu"})
|
||||
|
||||
|
||||
def _enabled_minimax_h3_fusions(value: str | None = None) -> frozenset[str]:
|
||||
"""Parse the independently switchable inference fusion set."""
|
||||
raw = envs.FASTVIDEO_MINIMAX_H3_FUSIONS if value is None else value
|
||||
normalized = raw.strip().lower()
|
||||
if normalized in {"", "0", "none"}:
|
||||
return frozenset()
|
||||
if normalized in {"1", "all"}:
|
||||
return _MINIMAX_H3_FUSION_NAMES
|
||||
enabled = frozenset(item.strip() for item in normalized.split(",") if item.strip())
|
||||
unknown = enabled - _MINIMAX_H3_FUSION_NAMES
|
||||
if unknown:
|
||||
supported = ",".join(sorted(_MINIMAX_H3_FUSION_NAMES))
|
||||
raise ValueError(f"Unknown MiniMax H3 fusion(s) {sorted(unknown)}; expected a subset of {supported}.")
|
||||
return enabled
|
||||
|
||||
|
||||
def _can_run_minimax_h3_fusion(tensor: torch.Tensor) -> bool:
|
||||
"""Triton kernels are inference-only and stay outside Dynamo capture.
|
||||
|
||||
The ``HAVE_TRITON`` check makes the eager fallback exact: on a CUDA build
|
||||
whose Triton failed to import, an enabled fusion falls back instead of
|
||||
hitting the strict wrappers' hard RuntimeError mid-forward.
|
||||
"""
|
||||
return (HAVE_TRITON and tensor.is_cuda and not torch.is_grad_enabled() and not torch.compiler.is_compiling())
|
||||
|
||||
|
||||
class MiniMaxH3RotaryPosEmbed(nn.Module):
|
||||
@@ -62,6 +101,7 @@ class MiniMaxH3FeedForward(nn.Module):
|
||||
ffn_dim: int,
|
||||
quant_config: QuantizationConfig | None = None,
|
||||
prefix: str = "",
|
||||
fuse_swiglu: bool = False,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.fc_in = ReplicatedLinear(
|
||||
@@ -78,11 +118,15 @@ class MiniMaxH3FeedForward(nn.Module):
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.fc_out",
|
||||
)
|
||||
self.fuse_swiglu = fuse_swiglu
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states, _ = self.fc_in(hidden_states)
|
||||
hidden_states, gate = hidden_states.chunk(2, dim=-1)
|
||||
hidden_states = hidden_states * F.silu(gate)
|
||||
if self.fuse_swiglu and _can_run_minimax_h3_fusion(hidden_states):
|
||||
hidden_states = minimax_h3_swiglu(hidden_states)
|
||||
else:
|
||||
hidden_states, gate = hidden_states.chunk(2, dim=-1)
|
||||
hidden_states = hidden_states * F.silu(gate)
|
||||
hidden_states, _ = self.fc_out(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
@@ -99,6 +143,7 @@ class MiniMaxH3Attention(nn.Module):
|
||||
supported_attention_backends: tuple[AttentionBackendEnum, ...],
|
||||
quant_config: QuantizationConfig | None,
|
||||
prefix: str,
|
||||
fuse_qknorm_rope: bool = False,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.num_attention_heads = num_attention_heads
|
||||
@@ -134,6 +179,7 @@ class MiniMaxH3Attention(nn.Module):
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.to_out",
|
||||
)
|
||||
self.fuse_qknorm_rope = fuse_qknorm_rope
|
||||
# VSA carries a learned gate on its pooled-compression branch. The H3
|
||||
# checkpoint has no such weight, so the loader zero-initializes it
|
||||
# (ALLOWED_NEW_PARAM_PATTERNS) and the branch is exactly disabled
|
||||
@@ -211,11 +257,18 @@ class MiniMaxH3Attention(nn.Module):
|
||||
query = query.unflatten(-1, (self.num_attention_heads, self.attention_head_dim))
|
||||
key = key.unflatten(-1, (self.num_attention_heads, self.attention_head_dim))
|
||||
value = value.unflatten(-1, (self.num_attention_heads, self.attention_head_dim))
|
||||
query = self.norm_q(query)
|
||||
key = self.norm_k(key)
|
||||
if rotary_emb is not None:
|
||||
query = self._apply_rotary_emb(query, rotary_emb)
|
||||
key = self._apply_rotary_emb(key, rotary_emb)
|
||||
if (self.fuse_qknorm_rope and rotary_emb is not None and _can_run_minimax_h3_fusion(query)):
|
||||
cos, sin = rotary_emb
|
||||
cos = cos.to(query.dtype)
|
||||
sin = sin.to(query.dtype)
|
||||
query = fused_qknorm_rope(query, self.norm_q.weight, cos, sin, self.norm_q.eps)
|
||||
key = fused_qknorm_rope(key, self.norm_k.weight, cos, sin, self.norm_k.eps)
|
||||
else:
|
||||
query = self.norm_q(query)
|
||||
key = self.norm_k(key)
|
||||
if rotary_emb is not None:
|
||||
query = self._apply_rotary_emb(query, rotary_emb)
|
||||
key = self._apply_rotary_emb(key, rotary_emb)
|
||||
|
||||
# H3 rotates only 96/128 channels, which the generic `freqs_cis`
|
||||
# branch cannot express. Apply it above, then pass no RoPE here.
|
||||
@@ -397,6 +450,9 @@ class MiniMaxH3TransformerBlock(nn.Module):
|
||||
quant_config: QuantizationConfig | None,
|
||||
prefix: str,
|
||||
adaln_apply_silu: bool = True,
|
||||
fuse_modulate: bool = False,
|
||||
fuse_qknorm_rope: bool = False,
|
||||
fuse_swiglu: bool = False,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.norm1 = nn.RMSNorm(hidden_size, eps=norm_eps)
|
||||
@@ -408,6 +464,7 @@ class MiniMaxH3TransformerBlock(nn.Module):
|
||||
supported_attention_backends,
|
||||
quant_config,
|
||||
prefix=f"{prefix}.attn",
|
||||
fuse_qknorm_rope=fuse_qknorm_rope,
|
||||
)
|
||||
self.norm2 = nn.RMSNorm(hidden_size, eps=norm_eps)
|
||||
self.ff = MiniMaxH3FeedForward(
|
||||
@@ -415,6 +472,7 @@ class MiniMaxH3TransformerBlock(nn.Module):
|
||||
ffn_dim,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.ff",
|
||||
fuse_swiglu=fuse_swiglu,
|
||||
)
|
||||
self.adaln_proj = MiniMaxH3AdaLayerNormModulation(
|
||||
time_embed_dim,
|
||||
@@ -423,6 +481,7 @@ class MiniMaxH3TransformerBlock(nn.Module):
|
||||
prefix=f"{prefix}.adaln_proj",
|
||||
apply_silu=adaln_apply_silu,
|
||||
)
|
||||
self.fuse_modulate = fuse_modulate
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -435,19 +494,39 @@ class MiniMaxH3TransformerBlock(nn.Module):
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
|
||||
t.to(hidden_states.dtype) for t in self.adaln_proj(temb))
|
||||
|
||||
residual = hidden_states
|
||||
norm_hidden_states = self.norm1(hidden_states)
|
||||
norm_hidden_states = norm_hidden_states * (
|
||||
1.0 + scale_msa.index_select(0, adaln_indices)) + shift_msa.index_select(0, adaln_indices)
|
||||
use_modulate_fusion = self.fuse_modulate and _can_run_minimax_h3_fusion(hidden_states)
|
||||
if use_modulate_fusion:
|
||||
norm_hidden_states = fused_rmsnorm_modulate(
|
||||
hidden_states,
|
||||
self.norm1.weight,
|
||||
scale_msa,
|
||||
shift_msa,
|
||||
adaln_indices,
|
||||
self.norm1.eps,
|
||||
)
|
||||
else:
|
||||
norm_hidden_states = self.norm1(hidden_states)
|
||||
norm_hidden_states = norm_hidden_states * (
|
||||
1.0 + scale_msa.index_select(0, adaln_indices)) + shift_msa.index_select(0, adaln_indices)
|
||||
attention_output = self.attn(norm_hidden_states, rotary_emb, original_seq_len)
|
||||
hidden_states = residual + gate_msa.index_select(0, adaln_indices) * attention_output
|
||||
|
||||
residual = hidden_states
|
||||
norm_hidden_states = self.norm2(hidden_states)
|
||||
norm_hidden_states = norm_hidden_states * (
|
||||
1.0 + scale_mlp.index_select(0, adaln_indices)) + shift_mlp.index_select(0, adaln_indices)
|
||||
if use_modulate_fusion:
|
||||
hidden_states, norm_hidden_states = fused_residual_gate_rmsnorm_modulate(
|
||||
hidden_states,
|
||||
attention_output,
|
||||
gate_msa,
|
||||
self.norm2.weight,
|
||||
scale_mlp,
|
||||
shift_mlp,
|
||||
adaln_indices,
|
||||
self.norm2.eps,
|
||||
)
|
||||
else:
|
||||
hidden_states = hidden_states + gate_msa.index_select(0, adaln_indices) * attention_output
|
||||
norm_hidden_states = self.norm2(hidden_states)
|
||||
norm_hidden_states = norm_hidden_states * (
|
||||
1.0 + scale_mlp.index_select(0, adaln_indices)) + shift_mlp.index_select(0, adaln_indices)
|
||||
feed_forward_output = self.ff(norm_hidden_states)
|
||||
return residual + gate_mlp.index_select(0, adaln_indices) * feed_forward_output
|
||||
return hidden_states + gate_mlp.index_select(0, adaln_indices) * feed_forward_output
|
||||
|
||||
|
||||
class MiniMaxH3Transformer3DModel(BaseDiT):
|
||||
@@ -493,6 +572,17 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
|
||||
def __init__(self, config: MiniMaxH3Config, hf_config: dict[str, Any]) -> None:
|
||||
super().__init__(config, hf_config)
|
||||
arch = config.arch_config
|
||||
self.enabled_fusions = _enabled_minimax_h3_fusions()
|
||||
if self.enabled_fusions:
|
||||
if HAVE_TRITON:
|
||||
logger.info(
|
||||
"MiniMax H3 inference fusions enabled: %s (CUDA inference-only; grad-enabled and "
|
||||
"torch.compile-captured forwards fall back to eager).",
|
||||
",".join(sorted(self.enabled_fusions)))
|
||||
else:
|
||||
logger.warning(
|
||||
"FASTVIDEO_MINIMAX_H3_FUSIONS requested %s but Triton is unavailable; "
|
||||
"every forward stays on the eager path.", ",".join(sorted(self.enabled_fusions)))
|
||||
sp_world_size = get_sp_world_size() if model_parallel_is_initialized() else 1
|
||||
if arch.num_attention_heads % sp_world_size:
|
||||
raise ValueError(f"MiniMax H3 attention heads ({arch.num_attention_heads}) must be divisible by "
|
||||
@@ -590,6 +680,9 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
|
||||
config.quant_config,
|
||||
prefix=f"{config.prefix}.transformer_blocks.{index}",
|
||||
adaln_apply_silu=self.adaln_rank is None,
|
||||
fuse_modulate="modulate" in self.enabled_fusions,
|
||||
fuse_qknorm_rope="qknorm_rope" in self.enabled_fusions,
|
||||
fuse_swiglu="swiglu" in self.enabled_fusions,
|
||||
) for index in range(arch.num_layers)
|
||||
])
|
||||
self.norm_out = MiniMaxH3AdaLayerNormOut(
|
||||
@@ -616,6 +709,20 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
|
||||
)
|
||||
self.__post_init__()
|
||||
|
||||
def prepare_for_compile(self) -> None:
|
||||
"""Pipeline hook, called once right before torch.compile wraps the blocks.
|
||||
|
||||
Dynamo capture traces the eager branch of every fusion guard, so an
|
||||
enabled ``FASTVIDEO_MINIMAX_H3_FUSIONS`` set is silently inert inside
|
||||
compiled block forwards (H3 compiles per-block by default). Say so
|
||||
once instead of leaving the flag looking active.
|
||||
"""
|
||||
if self.enabled_fusions:
|
||||
logger.warning(
|
||||
"torch.compile is enabled for MiniMax H3, so the requested inference fusions (%s) are "
|
||||
"inert inside compiled block forwards; the compiled eager path runs instead.",
|
||||
",".join(sorted(self.enabled_fusions)))
|
||||
|
||||
def materialize_non_persistent_buffers(
|
||||
self,
|
||||
device: torch.device,
|
||||
@@ -734,14 +841,17 @@ class MiniMaxH3Transformer3DModel(BaseDiT):
|
||||
local_timestep_indices, _ = sequence_model_parallel_shard(local_timestep_indices, dim=0)
|
||||
rotary_emb = (rotary_cos, rotary_sin)
|
||||
|
||||
for block in self.transformer_blocks:
|
||||
packed_hidden_states = block(
|
||||
packed_hidden_states,
|
||||
temb,
|
||||
adaln_indices,
|
||||
rotary_emb,
|
||||
original_seq_len,
|
||||
)
|
||||
# The eager driver owns profiling markers while each block's compiled
|
||||
# forward owns the graph that the marker surrounds.
|
||||
for block_index, block in enumerate(self.transformer_blocks):
|
||||
with nvtx_range(f"minimax_h3.transformer_block.{block_index}"):
|
||||
packed_hidden_states = block(
|
||||
packed_hidden_states,
|
||||
temb,
|
||||
adaln_indices,
|
||||
rotary_emb,
|
||||
original_seq_len,
|
||||
)
|
||||
|
||||
packed_hidden_states = self.norm_out(
|
||||
packed_hidden_states,
|
||||
|
||||
@@ -0,0 +1,20 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Inference-only MiniMax H3 fusions adapted from NVlabs/Sana Sol-Engine.
|
||||
|
||||
Source: https://github.com/NVlabs/Sana/tree/sol-engine/models/minimax_h3/GB200
|
||||
"""
|
||||
|
||||
from .modulation import (
|
||||
fused_residual_gate_rmsnorm_modulate,
|
||||
fused_rmsnorm_modulate,
|
||||
)
|
||||
from .qknorm_rope import HAVE_TRITON, fused_qknorm_rope
|
||||
from .swiglu import minimax_h3_swiglu
|
||||
|
||||
__all__ = [
|
||||
"HAVE_TRITON",
|
||||
"fused_qknorm_rope",
|
||||
"fused_residual_gate_rmsnorm_modulate",
|
||||
"fused_rmsnorm_modulate",
|
||||
"minimax_h3_swiglu",
|
||||
]
|
||||
@@ -0,0 +1,302 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""MiniMax H3 RMSNorm and row-indexed modulation fusions."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
|
||||
try:
|
||||
import triton
|
||||
import triton.language as tl
|
||||
except ImportError as exc: # pragma: no cover - depends on the runtime image
|
||||
triton = None
|
||||
tl = None
|
||||
_TRITON_IMPORT_ERROR: ImportError | None = exc
|
||||
else:
|
||||
_TRITON_IMPORT_ERROR = None
|
||||
|
||||
|
||||
__all__ = [
|
||||
"fused_residual_gate_rmsnorm_modulate",
|
||||
"fused_rmsnorm_modulate",
|
||||
]
|
||||
|
||||
_SUPPORTED_DTYPES = (torch.float16, torch.bfloat16, torch.float32)
|
||||
_rmsnorm_modulate_kernel = None
|
||||
_residual_gate_rmsnorm_modulate_kernel = None
|
||||
|
||||
|
||||
if triton is not None:
|
||||
|
||||
@triton.jit
|
||||
def _rmsnorm_modulate_kernel(
|
||||
out_ptr,
|
||||
x_ptr,
|
||||
weight_ptr,
|
||||
scale_ptr,
|
||||
shift_ptr,
|
||||
index_ptr,
|
||||
n_cols,
|
||||
n_index,
|
||||
eps,
|
||||
stride_x_row,
|
||||
stride_scale_row,
|
||||
stride_shift_row,
|
||||
BLOCK: tl.constexpr,
|
||||
):
|
||||
row = tl.program_id(0).to(tl.int64)
|
||||
cols = tl.arange(0, BLOCK)
|
||||
mask = cols < n_cols
|
||||
x_offsets = row * stride_x_row + cols
|
||||
table_row = tl.load(index_ptr + row % n_index).to(tl.int64)
|
||||
|
||||
x = tl.load(x_ptr + x_offsets, mask=mask, other=0.0).to(tl.float32)
|
||||
variance = tl.sum(x * x, axis=0) / n_cols
|
||||
normed = x * tl.math.rsqrt(variance + eps)
|
||||
weight = tl.load(weight_ptr + cols, mask=mask, other=0.0).to(tl.float32)
|
||||
scale = tl.load(
|
||||
scale_ptr + table_row * stride_scale_row + cols,
|
||||
mask=mask,
|
||||
other=0.0,
|
||||
).to(tl.float32)
|
||||
shift = tl.load(
|
||||
shift_ptr + table_row * stride_shift_row + cols,
|
||||
mask=mask,
|
||||
other=0.0,
|
||||
).to(tl.float32)
|
||||
output = normed * weight * (1.0 + scale) + shift
|
||||
tl.store(out_ptr + x_offsets, output.to(out_ptr.dtype.element_ty), mask=mask)
|
||||
|
||||
@triton.jit
|
||||
def _residual_gate_rmsnorm_modulate_kernel(
|
||||
hidden_out_ptr,
|
||||
normed_out_ptr,
|
||||
residual_ptr,
|
||||
branch_ptr,
|
||||
gate_ptr,
|
||||
weight_ptr,
|
||||
scale_ptr,
|
||||
shift_ptr,
|
||||
index_ptr,
|
||||
n_cols,
|
||||
n_index,
|
||||
eps,
|
||||
stride_input_row,
|
||||
stride_gate_row,
|
||||
stride_scale_row,
|
||||
stride_shift_row,
|
||||
BLOCK: tl.constexpr,
|
||||
):
|
||||
row = tl.program_id(0).to(tl.int64)
|
||||
cols = tl.arange(0, BLOCK)
|
||||
mask = cols < n_cols
|
||||
input_offsets = row * stride_input_row + cols
|
||||
table_row = tl.load(index_ptr + row % n_index).to(tl.int64)
|
||||
|
||||
residual = tl.load(residual_ptr + input_offsets, mask=mask, other=0.0).to(tl.float32)
|
||||
branch = tl.load(branch_ptr + input_offsets, mask=mask, other=0.0).to(tl.float32)
|
||||
gate = tl.load(
|
||||
gate_ptr + table_row * stride_gate_row + cols,
|
||||
mask=mask,
|
||||
other=0.0,
|
||||
).to(tl.float32)
|
||||
hidden = residual + gate * branch
|
||||
tl.store(hidden_out_ptr + input_offsets, hidden.to(hidden_out_ptr.dtype.element_ty), mask=mask)
|
||||
|
||||
variance = tl.sum(hidden * hidden, axis=0) / n_cols
|
||||
normed = hidden * tl.math.rsqrt(variance + eps)
|
||||
weight = tl.load(weight_ptr + cols, mask=mask, other=0.0).to(tl.float32)
|
||||
scale = tl.load(
|
||||
scale_ptr + table_row * stride_scale_row + cols,
|
||||
mask=mask,
|
||||
other=0.0,
|
||||
).to(tl.float32)
|
||||
shift = tl.load(
|
||||
shift_ptr + table_row * stride_shift_row + cols,
|
||||
mask=mask,
|
||||
other=0.0,
|
||||
).to(tl.float32)
|
||||
output = normed * weight * (1.0 + scale) + shift
|
||||
tl.store(normed_out_ptr + input_offsets, output.to(normed_out_ptr.dtype.element_ty), mask=mask)
|
||||
|
||||
|
||||
def _validate_contract(
|
||||
x: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
tables: tuple[torch.Tensor, ...],
|
||||
index: torch.Tensor,
|
||||
eps: float,
|
||||
) -> None:
|
||||
if x.ndim < 2:
|
||||
raise ValueError(f"x must have shape (..., sequence_length, hidden_size), got {tuple(x.shape)}.")
|
||||
if x.numel() == 0 or x.shape[-1] == 0:
|
||||
raise ValueError("x must not be empty.")
|
||||
if x.dtype not in _SUPPORTED_DTYPES:
|
||||
raise TypeError(f"x must use float16, bfloat16, or float32, got {x.dtype}.")
|
||||
hidden_size = x.shape[-1]
|
||||
sequence_length = x.shape[-2]
|
||||
if weight.shape != (hidden_size, ):
|
||||
raise ValueError(f"weight must have shape ({hidden_size},), got {tuple(weight.shape)}.")
|
||||
if weight.dtype not in _SUPPORTED_DTYPES:
|
||||
raise TypeError(f"weight must use float16, bfloat16, or float32, got {weight.dtype}.")
|
||||
if index.ndim != 1 or index.numel() != sequence_length:
|
||||
raise ValueError(
|
||||
f"index must have shape ({sequence_length},) so it can wrap over batch rows, got {tuple(index.shape)}."
|
||||
)
|
||||
if index.dtype not in (torch.int32, torch.int64):
|
||||
raise TypeError(f"index must use int32 or int64, got {index.dtype}.")
|
||||
if not isinstance(eps, (float, int)) or isinstance(eps, bool) or not math.isfinite(eps) or eps <= 0:
|
||||
raise ValueError(f"eps must be a positive finite number, got {eps!r}.")
|
||||
|
||||
table_rows = tables[0].shape[0] if tables and tables[0].ndim == 2 else None
|
||||
for name, table in zip(("gate", "scale", "shift")[-len(tables):], tables, strict=True):
|
||||
if table.ndim != 2 or table.shape[1] != hidden_size:
|
||||
raise ValueError(f"{name} must have shape (table_rows, {hidden_size}), got {tuple(table.shape)}.")
|
||||
if table.shape[0] == 0 or table.shape[0] != table_rows:
|
||||
raise ValueError("all modulation tables must have the same non-zero row count.")
|
||||
if table.dtype not in _SUPPORTED_DTYPES:
|
||||
raise TypeError(f"{name} must use float16, bfloat16, or float32, got {table.dtype}.")
|
||||
|
||||
tensors = (x, weight, *tables, index)
|
||||
if any(tensor.device != x.device for tensor in tensors[1:]):
|
||||
raise ValueError("x, weight, modulation tables, and index must be on the same device.")
|
||||
|
||||
|
||||
def _validate_residual_branch(residual: torch.Tensor, branch: torch.Tensor) -> None:
|
||||
if branch.shape != residual.shape:
|
||||
raise ValueError(f"branch must match residual shape {tuple(residual.shape)}, got {tuple(branch.shape)}.")
|
||||
if branch.dtype != residual.dtype:
|
||||
raise TypeError(f"branch dtype must match residual dtype {residual.dtype}, got {branch.dtype}.")
|
||||
if branch.device != residual.device:
|
||||
raise ValueError("branch and residual must be on the same device.")
|
||||
|
||||
|
||||
def _require_triton_cuda(x: torch.Tensor) -> None:
|
||||
if triton is None:
|
||||
detail = f": {_TRITON_IMPORT_ERROR}" if _TRITON_IMPORT_ERROR is not None else ""
|
||||
raise RuntimeError(f"MiniMax H3 modulation fusion requires Triton{detail}.")
|
||||
if x.device.type != "cuda":
|
||||
raise RuntimeError(f"MiniMax H3 modulation fusion requires CUDA tensors, got device {x.device}.")
|
||||
|
||||
|
||||
def _require_forward_only(*tensors: torch.Tensor) -> None:
|
||||
if torch.is_grad_enabled() and any(tensor.requires_grad for tensor in tensors):
|
||||
raise RuntimeError("MiniMax H3 modulation fusion is forward-only and does not support autograd.")
|
||||
|
||||
|
||||
def _next_power_of_two(value: int) -> int:
|
||||
return 1 << (value - 1).bit_length()
|
||||
|
||||
|
||||
def _num_warps(block_size: int) -> int:
|
||||
if block_size >= 8192:
|
||||
return 16
|
||||
if block_size >= 2048:
|
||||
return 8
|
||||
return 4
|
||||
|
||||
|
||||
def _row_addressable(table: torch.Tensor) -> torch.Tensor:
|
||||
return table if table.stride(-1) == 1 else table.contiguous()
|
||||
|
||||
|
||||
def fused_rmsnorm_modulate(
|
||||
x: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
scale: torch.Tensor,
|
||||
shift: torch.Tensor,
|
||||
index: torch.Tensor,
|
||||
eps: float,
|
||||
) -> torch.Tensor:
|
||||
"""Run RMSNorm and row-indexed modulation in one strict Triton kernel.
|
||||
|
||||
``index`` values must lie in ``[0, table_rows)``. Unlike eager
|
||||
``index_select``, the kernel does not raise on out-of-range values (a
|
||||
device-side bounds check would synchronize); callers are safe by
|
||||
construction (``timestep_indices * 3 + token_tags``, SP pads with 0).
|
||||
"""
|
||||
_validate_contract(x, weight, (scale, shift), index, eps)
|
||||
_require_forward_only(x, weight, scale, shift)
|
||||
_require_triton_cuda(x)
|
||||
|
||||
hidden_size = x.shape[-1]
|
||||
flat_x = x.reshape(-1, hidden_size).contiguous()
|
||||
weight = weight.contiguous()
|
||||
scale = _row_addressable(scale)
|
||||
shift = _row_addressable(shift)
|
||||
index = index.contiguous()
|
||||
output = torch.empty_like(flat_x)
|
||||
block_size = _next_power_of_two(hidden_size)
|
||||
_rmsnorm_modulate_kernel[(flat_x.shape[0], )](
|
||||
output,
|
||||
flat_x,
|
||||
weight,
|
||||
scale,
|
||||
shift,
|
||||
index,
|
||||
hidden_size,
|
||||
index.numel(),
|
||||
eps,
|
||||
flat_x.stride(0),
|
||||
scale.stride(0),
|
||||
shift.stride(0),
|
||||
BLOCK=block_size,
|
||||
num_warps=_num_warps(block_size),
|
||||
)
|
||||
return output.view_as(x)
|
||||
|
||||
|
||||
def fused_residual_gate_rmsnorm_modulate(
|
||||
residual: torch.Tensor,
|
||||
branch: torch.Tensor,
|
||||
gate: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
scale: torch.Tensor,
|
||||
shift: torch.Tensor,
|
||||
index: torch.Tensor,
|
||||
eps: float,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Fuse residual update, row-indexed gate, RMSNorm, and modulation.
|
||||
|
||||
``index`` values must lie in ``[0, table_rows)``; see
|
||||
:func:`fused_rmsnorm_modulate` for why the wrapper does not check them.
|
||||
"""
|
||||
_validate_residual_branch(residual, branch)
|
||||
_validate_contract(residual, weight, (gate, scale, shift), index, eps)
|
||||
_require_forward_only(residual, branch, gate, weight, scale, shift)
|
||||
_require_triton_cuda(residual)
|
||||
|
||||
hidden_size = residual.shape[-1]
|
||||
flat_residual = residual.reshape(-1, hidden_size).contiguous()
|
||||
flat_branch = branch.reshape(-1, hidden_size).contiguous()
|
||||
weight = weight.contiguous()
|
||||
gate = _row_addressable(gate)
|
||||
scale = _row_addressable(scale)
|
||||
shift = _row_addressable(shift)
|
||||
index = index.contiguous()
|
||||
hidden = torch.empty_like(flat_residual)
|
||||
modulated = torch.empty_like(flat_residual)
|
||||
block_size = _next_power_of_two(hidden_size)
|
||||
_residual_gate_rmsnorm_modulate_kernel[(flat_residual.shape[0], )](
|
||||
hidden,
|
||||
modulated,
|
||||
flat_residual,
|
||||
flat_branch,
|
||||
gate,
|
||||
weight,
|
||||
scale,
|
||||
shift,
|
||||
index,
|
||||
hidden_size,
|
||||
index.numel(),
|
||||
eps,
|
||||
flat_residual.stride(0),
|
||||
gate.stride(0),
|
||||
scale.stride(0),
|
||||
shift.stride(0),
|
||||
BLOCK=block_size,
|
||||
num_warps=_num_warps(block_size),
|
||||
)
|
||||
return hidden.view_as(residual), modulated.view_as(residual)
|
||||
@@ -0,0 +1,174 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Fused per-head RMSNorm and partial rotary embedding for MiniMax H3."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
import torch
|
||||
|
||||
try:
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
HAVE_TRITON = True
|
||||
except ImportError: # pragma: no cover - exercised only in environments without Triton
|
||||
triton = None
|
||||
tl = None
|
||||
HAVE_TRITON = False
|
||||
|
||||
|
||||
if HAVE_TRITON:
|
||||
|
||||
@triton.jit
|
||||
def _qknorm_partial_rope_kernel(
|
||||
out_ptr,
|
||||
x_ptr,
|
||||
weight_ptr,
|
||||
cos_ptr,
|
||||
sin_ptr,
|
||||
head_dim,
|
||||
rotary_dim,
|
||||
half_rotary_dim,
|
||||
num_heads,
|
||||
seq_len,
|
||||
eps,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
# int64, like the sibling kernels: with int32 program ids,
|
||||
# ``row * head_dim`` wraps once the flattened input reaches 2**31
|
||||
# elements (H3's 56 heads x 128 head_dim crosses that at
|
||||
# batch*seq >= 299_593 tokens per rank) and the loads/stores below
|
||||
# become out-of-bounds. ``seq_index`` inherits int64 from ``row``.
|
||||
row = tl.program_id(0).to(tl.int64)
|
||||
seq_index = (row // num_heads) % seq_len
|
||||
cols = tl.arange(0, BLOCK_SIZE)
|
||||
head_mask = cols < head_dim
|
||||
row_offset = row * head_dim
|
||||
|
||||
x = tl.load(x_ptr + row_offset + cols, mask=head_mask, other=0.0).to(tl.float32)
|
||||
variance = tl.sum(x * x, axis=0) / head_dim
|
||||
inv_rms = tl.math.rsqrt(variance + eps)
|
||||
weight = tl.load(weight_ptr + cols, mask=head_mask, other=0.0).to(tl.float32)
|
||||
normalized = x * inv_rms * weight
|
||||
|
||||
rotary_mask = cols < rotary_dim
|
||||
first_half = cols < half_rotary_dim
|
||||
partner_col = tl.where(first_half, cols + half_rotary_dim, cols - half_rotary_dim)
|
||||
partner_x = tl.load(
|
||||
x_ptr + row_offset + partner_col,
|
||||
mask=rotary_mask,
|
||||
other=0.0,
|
||||
).to(tl.float32)
|
||||
partner_weight = tl.load(weight_ptr + partner_col, mask=rotary_mask, other=0.0).to(tl.float32)
|
||||
partner_normalized = partner_x * inv_rms * partner_weight
|
||||
rotated = tl.where(first_half, -partner_normalized, partner_normalized)
|
||||
|
||||
table_offset = seq_index * rotary_dim + cols
|
||||
cos = tl.load(cos_ptr + table_offset, mask=rotary_mask, other=1.0).to(tl.float32)
|
||||
sin = tl.load(sin_ptr + table_offset, mask=rotary_mask, other=0.0).to(tl.float32)
|
||||
rotary_output = normalized * cos + rotated * sin
|
||||
output = tl.where(rotary_mask, rotary_output, normalized)
|
||||
tl.store(out_ptr + row_offset + cols, output.to(out_ptr.dtype.element_ty), mask=head_mask)
|
||||
|
||||
|
||||
_SUPPORTED_DTYPES = (torch.float16, torch.bfloat16, torch.float32)
|
||||
|
||||
|
||||
def _validate_inputs(
|
||||
x: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
cos: torch.Tensor,
|
||||
sin: torch.Tensor,
|
||||
eps: float,
|
||||
) -> tuple[int, int, int, int, int]:
|
||||
for name, tensor in (("x", x), ("weight", weight), ("cos", cos), ("sin", sin)):
|
||||
if not isinstance(tensor, torch.Tensor):
|
||||
raise TypeError(f"{name} must be a torch.Tensor, got {type(tensor).__name__}")
|
||||
|
||||
if x.ndim != 4:
|
||||
raise ValueError(f"x must have shape (batch, seq, heads, head_dim), got {tuple(x.shape)}")
|
||||
batch, seq_len, num_heads, head_dim = x.shape
|
||||
if min(batch, seq_len, num_heads, head_dim) <= 0:
|
||||
raise ValueError(f"x dimensions must all be positive, got {tuple(x.shape)}")
|
||||
if weight.shape != (head_dim, ):
|
||||
raise ValueError(f"weight must have shape ({head_dim},), got {tuple(weight.shape)}")
|
||||
if cos.ndim != 2:
|
||||
raise ValueError(f"cos must have shape (seq, rotary_dim), got {tuple(cos.shape)}")
|
||||
if sin.shape != cos.shape:
|
||||
raise ValueError(f"sin must match cos shape {tuple(cos.shape)}, got {tuple(sin.shape)}")
|
||||
if cos.shape[0] != seq_len:
|
||||
raise ValueError(f"cos/sin sequence length must be {seq_len}, got {cos.shape[0]}")
|
||||
|
||||
rotary_dim = cos.shape[1]
|
||||
if rotary_dim <= 0:
|
||||
raise ValueError(f"rotary_dim must be positive, got {rotary_dim}")
|
||||
if rotary_dim > head_dim:
|
||||
raise ValueError(f"rotary_dim must not exceed head_dim, got rotary_dim={rotary_dim}, head_dim={head_dim}")
|
||||
if rotary_dim % 2:
|
||||
raise ValueError(f"rotary_dim must be even, got {rotary_dim}")
|
||||
|
||||
if x.dtype not in _SUPPORTED_DTYPES:
|
||||
raise TypeError(f"x dtype must be float16, bfloat16, or float32, got {x.dtype}")
|
||||
for name, tensor in (("weight", weight), ("cos", cos), ("sin", sin)):
|
||||
if tensor.dtype != x.dtype:
|
||||
raise TypeError(f"{name} dtype must match x dtype {x.dtype}, got {tensor.dtype}")
|
||||
if tensor.device != x.device:
|
||||
raise ValueError(f"{name} device must match x device {x.device}, got {tensor.device}")
|
||||
|
||||
if not isinstance(eps, (float, int)) or not math.isfinite(float(eps)) or eps <= 0:
|
||||
raise ValueError(f"eps must be a positive finite number, got {eps!r}")
|
||||
return batch, seq_len, num_heads, head_dim, rotary_dim
|
||||
|
||||
|
||||
def fused_qknorm_rope(
|
||||
x: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
cos: torch.Tensor,
|
||||
sin: torch.Tensor,
|
||||
eps: float,
|
||||
) -> torch.Tensor:
|
||||
"""Run per-head RMSNorm and partial RoPE in one Sol-Engine-style kernel.
|
||||
|
||||
RMSNorm reduction and RoPE arithmetic stay in FP32 registers until the
|
||||
final store. Triton's reduction order and the absence of eager's BF16
|
||||
intermediate materializations can produce small, expected rounding drift.
|
||||
|
||||
Row offsets are computed in int64, so inputs beyond 2**31 total elements
|
||||
(about 300k tokens per rank at H3's 56 heads x 128 head_dim) address
|
||||
correctly.
|
||||
"""
|
||||
batch, seq_len, num_heads, head_dim, rotary_dim = _validate_inputs(x, weight, cos, sin, eps)
|
||||
if not weight.is_contiguous():
|
||||
raise ValueError("weight must be contiguous")
|
||||
if not cos.is_contiguous() or not sin.is_contiguous():
|
||||
raise ValueError("cos and sin must be contiguous (seq, rotary_dim) tables")
|
||||
if not x.is_cuda:
|
||||
raise RuntimeError("fused_qknorm_rope requires CUDA tensors")
|
||||
if not HAVE_TRITON:
|
||||
raise RuntimeError("fused_qknorm_rope requires Triton")
|
||||
if torch.is_grad_enabled() and any(tensor.requires_grad for tensor in (x, weight, cos, sin)):
|
||||
raise RuntimeError("fused_qknorm_rope is inference-only and does not implement autograd")
|
||||
|
||||
flat_x = x.reshape(-1, head_dim).contiguous()
|
||||
flat_out = torch.empty_like(flat_x)
|
||||
block_size = 1 << (head_dim - 1).bit_length()
|
||||
_qknorm_partial_rope_kernel[(flat_x.shape[0], )](
|
||||
flat_out,
|
||||
flat_x,
|
||||
weight,
|
||||
cos,
|
||||
sin,
|
||||
head_dim,
|
||||
rotary_dim,
|
||||
rotary_dim // 2,
|
||||
num_heads,
|
||||
seq_len,
|
||||
eps,
|
||||
BLOCK_SIZE=block_size,
|
||||
num_warps=4,
|
||||
)
|
||||
return flat_out.view(batch, seq_len, num_heads, head_dim)
|
||||
|
||||
|
||||
__all__ = ["HAVE_TRITON", "fused_qknorm_rope"]
|
||||
@@ -0,0 +1,104 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""MiniMax H3's value-first packed SwiGLU fusion."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
try:
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
HAVE_TRITON = True
|
||||
except ImportError: # pragma: no cover - exercised only in environments without Triton
|
||||
triton = None
|
||||
tl = None
|
||||
HAVE_TRITON = False
|
||||
|
||||
|
||||
def _validate_input(x: torch.Tensor) -> int:
|
||||
if x.ndim == 0:
|
||||
raise ValueError("MiniMax H3 SwiGLU expects at least one dimension")
|
||||
|
||||
packed_width = x.shape[-1]
|
||||
if packed_width == 0 or packed_width % 2 != 0:
|
||||
raise ValueError(
|
||||
"MiniMax H3 SwiGLU expects a positive even last dimension containing packed (value, gate) halves, "
|
||||
f"got {packed_width}"
|
||||
)
|
||||
if not x.is_floating_point():
|
||||
raise TypeError(f"MiniMax H3 SwiGLU expects a floating-point tensor, got {x.dtype}")
|
||||
return packed_width // 2
|
||||
|
||||
|
||||
if HAVE_TRITON:
|
||||
|
||||
@triton.jit
|
||||
def _minimax_h3_swiglu_kernel(
|
||||
out_ptr,
|
||||
x_ptr,
|
||||
ffn_dim,
|
||||
stride_in_row,
|
||||
stride_out_row,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
row = tl.program_id(0).to(tl.int64)
|
||||
cols = tl.arange(0, BLOCK_SIZE)
|
||||
mask = cols < ffn_dim
|
||||
|
||||
value = tl.load(x_ptr + row * stride_in_row + cols, mask=mask, other=0.0).to(tl.float32)
|
||||
gate = tl.load(x_ptr + row * stride_in_row + ffn_dim + cols, mask=mask, other=0.0).to(tl.float32)
|
||||
# Match Sol-Engine: keep the complete SwiGLU expression in FP32 and
|
||||
# convert only the final output store.
|
||||
out = value * (gate * tl.sigmoid(gate))
|
||||
tl.store(out_ptr + row * stride_out_row + cols, out.to(out_ptr.dtype.element_ty), mask=mask)
|
||||
|
||||
else:
|
||||
_minimax_h3_swiglu_kernel = None
|
||||
|
||||
|
||||
def _num_warps(block_size: int) -> int:
|
||||
if block_size >= 8192:
|
||||
return 16
|
||||
if block_size >= 2048:
|
||||
return 8
|
||||
return 4
|
||||
|
||||
|
||||
def minimax_h3_swiglu(x: torch.Tensor) -> torch.Tensor:
|
||||
"""Run the forward-only Triton fusion over an H3 ``(..., 2 * ffn_dim)`` input.
|
||||
|
||||
This is intentionally a strict kernel wrapper: callers own fallback policy and
|
||||
must only invoke it for a supported CUDA inference path.
|
||||
"""
|
||||
ffn_dim = _validate_input(x)
|
||||
if not x.is_cuda:
|
||||
raise ValueError("MiniMax H3 fused SwiGLU requires a CUDA tensor")
|
||||
if x.dtype not in (torch.float16, torch.bfloat16, torch.float32):
|
||||
raise TypeError(f"MiniMax H3 fused SwiGLU supports float16, bfloat16, and float32, got {x.dtype}")
|
||||
if torch.is_grad_enabled() and x.requires_grad:
|
||||
raise RuntimeError("MiniMax H3 fused SwiGLU is forward-only and does not implement autograd")
|
||||
if _minimax_h3_swiglu_kernel is None:
|
||||
raise RuntimeError("MiniMax H3 fused SwiGLU requires Triton")
|
||||
|
||||
packed_width = x.shape[-1]
|
||||
flat = x.reshape(-1, packed_width).contiguous()
|
||||
output_shape = (*x.shape[:-1], ffn_dim)
|
||||
if flat.shape[0] == 0:
|
||||
return torch.empty(output_shape, dtype=x.dtype, device=x.device)
|
||||
|
||||
out = torch.empty((flat.shape[0], ffn_dim), dtype=x.dtype, device=x.device)
|
||||
block_size = triton.next_power_of_2(ffn_dim)
|
||||
_minimax_h3_swiglu_kernel[(flat.shape[0],)](
|
||||
out,
|
||||
flat,
|
||||
ffn_dim,
|
||||
flat.stride(0),
|
||||
out.stride(0),
|
||||
BLOCK_SIZE=block_size,
|
||||
num_warps=_num_warps(block_size),
|
||||
)
|
||||
return out.view(output_shape)
|
||||
|
||||
|
||||
__all__ = ["HAVE_TRITON", "minimax_h3_swiglu"]
|
||||
@@ -1,6 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import field
|
||||
from typing import Any, Generic, TypeVar
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
@@ -8,11 +9,16 @@ from torch import nn
|
||||
from fastvideo.configs.models.encoders import (BaseEncoderOutput, ImageEncoderConfig, TextEncoderConfig)
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
TextEncoderOutputT = TypeVar("TextEncoderOutputT")
|
||||
|
||||
|
||||
class TextEncoder(nn.Module, ABC, Generic[TextEncoderOutputT]):
|
||||
"""Base for native encoders with a model-specific forward output contract."""
|
||||
|
||||
class TextEncoder(nn.Module, ABC):
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [])
|
||||
_stacked_params_mapping: list[tuple[str, str, str]] = field(default_factory=list)
|
||||
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = TextEncoderConfig()._supported_attention_backends
|
||||
supported_checkpoint_quantization_methods: frozenset[str] = frozenset()
|
||||
|
||||
def __init__(self, config: TextEncoderConfig) -> None:
|
||||
super().__init__()
|
||||
@@ -23,13 +29,7 @@ class TextEncoder(nn.Module, ABC):
|
||||
raise ValueError(f"Subclass {self.__class__.__name__} must define _supported_attention_backends")
|
||||
|
||||
@abstractmethod
|
||||
def forward(self,
|
||||
input_ids: torch.Tensor | None,
|
||||
position_ids: torch.Tensor | None = None,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
inputs_embeds: torch.Tensor | None = None,
|
||||
output_hidden_states: bool | None = None,
|
||||
**kwargs) -> BaseEncoderOutput:
|
||||
def forward(self, *args: Any, **kwargs: Any) -> TextEncoderOutputT:
|
||||
pass
|
||||
|
||||
@property
|
||||
|
||||
@@ -0,0 +1,453 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Serialized block-FP8 execution for the MiniMax-H3 Qwen3-VL encoder."""
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
from torch.nn.parameter import Parameter
|
||||
|
||||
try:
|
||||
import triton
|
||||
import triton.language as tl
|
||||
except ImportError:
|
||||
triton = None
|
||||
tl = None
|
||||
|
||||
from fastvideo.distributed import get_tp_world_size
|
||||
from fastvideo.layers.linear import LinearBase, LinearMethodBase
|
||||
from fastvideo.layers.quantization.base_config import QuantizationConfig
|
||||
from fastvideo.layers.quantization.fp8_config import FP8_DTYPE
|
||||
from fastvideo.models.utils import set_weight_attrs
|
||||
|
||||
|
||||
class MiniMaxH3SerializedFP8Config(QuantizationConfig):
|
||||
"""Serialized 128x128 block-FP8 contract for the H3 text encoder."""
|
||||
|
||||
def __init__(self, weight_block_size: tuple[int, int]) -> None:
|
||||
super().__init__()
|
||||
if weight_block_size != (128, 128):
|
||||
raise ValueError("MiniMax-H3 serialized FP8 requires weight_block_size=[128, 128], "
|
||||
f"got {list(weight_block_size)}")
|
||||
self.weight_block_size = weight_block_size
|
||||
self.is_checkpoint_fp8_serialized = True
|
||||
self.activation_scheme = "dynamic"
|
||||
|
||||
@classmethod
|
||||
def get_name(cls) -> str:
|
||||
return "fp8"
|
||||
|
||||
@classmethod
|
||||
def get_supported_act_dtypes(cls) -> list[torch.dtype]:
|
||||
return [torch.bfloat16]
|
||||
|
||||
@classmethod
|
||||
def get_min_capability(cls) -> int:
|
||||
return 100
|
||||
|
||||
@staticmethod
|
||||
def get_config_filenames() -> list[str]:
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config: dict[str, Any]) -> "MiniMaxH3SerializedFP8Config":
|
||||
quant_method = str(config.get("quant_method", "")).lower()
|
||||
if quant_method != "fp8":
|
||||
raise ValueError(f"MiniMax-H3 only supports serialized FP8 text-encoder checkpoints, got {quant_method!r}")
|
||||
if str(config.get("activation_scheme", "")).lower() != "dynamic":
|
||||
raise ValueError("MiniMax-H3 serialized FP8 requires dynamic activation quantization")
|
||||
if str(config.get("fmt", "e4m3")).lower() not in ("e4m3", "float8_e4m3fn"):
|
||||
raise ValueError(f"MiniMax-H3 serialized FP8 requires E4M3 weights, got {config.get('fmt')!r}")
|
||||
block_size = config.get("weight_block_size")
|
||||
if not isinstance(block_size, list | tuple) or len(block_size) != 2:
|
||||
raise ValueError("MiniMax-H3 serialized FP8 requires a two-dimensional weight_block_size")
|
||||
ignored_layers = config.get("modules_to_not_convert", config.get("ignored_layers", []))
|
||||
if not isinstance(ignored_layers, list | tuple):
|
||||
raise ValueError("MiniMax-H3 serialized FP8 modules_to_not_convert must be a sequence")
|
||||
language_exclusions = [
|
||||
name for name in ignored_layers
|
||||
if isinstance(name, str) and (name.startswith("language_model.") or ".language_model." in name)
|
||||
]
|
||||
if language_exclusions:
|
||||
raise ValueError("MiniMax-H3 does not support partially quantized language stacks; "
|
||||
f"ignored language layers: {language_exclusions[:3]}")
|
||||
if not any(isinstance(name, str) and "visual" in name for name in ignored_layers):
|
||||
raise ValueError("MiniMax-H3 serialized FP8 requires the vision stack to be listed in "
|
||||
"modules_to_not_convert")
|
||||
return cls((int(block_size[0]), int(block_size[1])))
|
||||
|
||||
def validate_runtime(self, device: torch.device) -> None:
|
||||
if device.type != "cuda":
|
||||
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 requires a CUDA device; "
|
||||
f"got {device.type!r}")
|
||||
capability = torch.cuda.get_device_capability(device)
|
||||
capability_number = capability[0] * 10 + capability[1]
|
||||
if capability_number < self.get_min_capability():
|
||||
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 requires GPU capability "
|
||||
f"sm{self.get_min_capability()} or newer, got sm{capability_number}")
|
||||
if capability[0] not in (10, 12):
|
||||
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 currently adapts SGLang's Blackwell "
|
||||
f"FlashInfer path; got unsupported sm{capability_number}")
|
||||
_require_sglang_per_token_group_fp8_quantization()
|
||||
_get_flashinfer_groupwise_fp8_gemm()
|
||||
|
||||
def get_quant_method(self, layer: torch.nn.Module, prefix: str):
|
||||
if isinstance(layer, LinearBase) and ".language_model.layers." in prefix:
|
||||
return MiniMaxH3SerializedFP8LinearMethod(self.weight_block_size)
|
||||
return None
|
||||
|
||||
|
||||
# Copyright 2024 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0.
|
||||
# Adapted from SGLang's per-token-group quantization kernels and Blackwell
|
||||
# FlashInfer dispatch at commit f99c62063c7dcfcd06784b885dc08cb52cf23865:
|
||||
# https://github.com/sgl-project/sglang/blob/f99c62063c7dcfcd06784b885dc08cb52cf23865/python/sglang/kernels/ops/quantization/fp8_kernel.py
|
||||
# https://github.com/sgl-project/sglang/blob/f99c62063c7dcfcd06784b885dc08cb52cf23865/python/sglang/srt/layers/quantization/fp8_utils.py
|
||||
if triton is not None:
|
||||
|
||||
@triton.jit
|
||||
def _h3_per_token_group_quant_fp8_row_major(
|
||||
input_ptr,
|
||||
output_ptr,
|
||||
scale_ptr,
|
||||
group_size,
|
||||
eps,
|
||||
fp8_min,
|
||||
fp8_max,
|
||||
BLOCK: tl.constexpr,
|
||||
):
|
||||
group_id = tl.program_id(0)
|
||||
input_ptr += group_id.to(tl.int64) * group_size
|
||||
output_ptr += group_id.to(tl.int64) * group_size
|
||||
scale_ptr += group_id
|
||||
|
||||
offsets = tl.arange(0, BLOCK)
|
||||
mask = offsets < group_size
|
||||
values = tl.load(input_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
|
||||
absmax = tl.maximum(tl.max(tl.abs(values)), eps)
|
||||
scale = absmax / fp8_max
|
||||
quantized = tl.clamp(values / scale, fp8_min, fp8_max).to(output_ptr.dtype.element_ty)
|
||||
|
||||
tl.store(output_ptr + offsets, quantized, mask=mask)
|
||||
tl.store(scale_ptr, scale)
|
||||
|
||||
@triton.jit
|
||||
def _h3_per_token_group_quant_fp8_column_major(
|
||||
input_ptr,
|
||||
output_ptr,
|
||||
scale_ptr,
|
||||
group_size,
|
||||
input_columns,
|
||||
scale_column_stride,
|
||||
eps,
|
||||
fp8_min,
|
||||
fp8_max,
|
||||
BLOCK: tl.constexpr,
|
||||
):
|
||||
group_id = tl.program_id(0)
|
||||
input_ptr += group_id.to(tl.int64) * group_size
|
||||
output_ptr += group_id.to(tl.int64) * group_size
|
||||
|
||||
groups_per_row = input_columns // group_size
|
||||
scale_column = group_id % groups_per_row
|
||||
scale_row = group_id // groups_per_row
|
||||
scale_ptr += scale_column * scale_column_stride + scale_row
|
||||
|
||||
offsets = tl.arange(0, BLOCK)
|
||||
mask = offsets < group_size
|
||||
values = tl.load(input_ptr + offsets, mask=mask, other=0.0).to(tl.float32)
|
||||
absmax = tl.maximum(tl.max(tl.abs(values)), eps)
|
||||
scale = absmax / fp8_max
|
||||
quantized = tl.clamp(values / scale, fp8_min, fp8_max).to(output_ptr.dtype.element_ty)
|
||||
|
||||
tl.store(output_ptr + offsets, quantized, mask=mask)
|
||||
tl.store(scale_ptr, scale)
|
||||
else:
|
||||
_h3_per_token_group_quant_fp8_row_major = None
|
||||
_h3_per_token_group_quant_fp8_column_major = None
|
||||
|
||||
|
||||
def _require_sglang_per_token_group_fp8_quantization() -> None:
|
||||
if (triton is None or _h3_per_token_group_quant_fp8_row_major is None
|
||||
or _h3_per_token_group_quant_fp8_column_major is None):
|
||||
raise RuntimeError(
|
||||
"MiniMax-H3 serialized blockwise FP8 requires Triton for SGLang-compatible "
|
||||
"per-token-group activation quantization")
|
||||
|
||||
|
||||
def _sglang_per_token_group_quant_fp8(
|
||||
input_tensor: torch.Tensor,
|
||||
group_size: int,
|
||||
*,
|
||||
column_major_scales: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""SGLang-compatible dynamic FP8 quantization for contiguous 2-D activations."""
|
||||
_require_sglang_per_token_group_fp8_quantization()
|
||||
if input_tensor.ndim != 2:
|
||||
raise ValueError(f"per-token-group FP8 quantization expects 2-D input, got {input_tensor.ndim}-D")
|
||||
if not input_tensor.is_contiguous():
|
||||
raise ValueError("per-token-group FP8 quantization requires contiguous input")
|
||||
if input_tensor.shape[-1] % group_size:
|
||||
raise ValueError(f"activation width {input_tensor.shape[-1]} is not divisible by group_size={group_size}")
|
||||
|
||||
quantized = torch.empty_like(input_tensor, dtype=FP8_DTYPE)
|
||||
rows, columns = input_tensor.shape
|
||||
groups_per_row = columns // group_size
|
||||
if column_major_scales:
|
||||
scales = torch.empty(
|
||||
(groups_per_row, rows),
|
||||
device=input_tensor.device,
|
||||
dtype=torch.float32,
|
||||
).permute(1, 0)
|
||||
else:
|
||||
scales = torch.empty(
|
||||
(rows, groups_per_row),
|
||||
device=input_tensor.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
|
||||
if rows:
|
||||
num_groups = input_tensor.numel() // group_size
|
||||
block = triton.next_power_of_2(group_size)
|
||||
num_warps = min(max(block // 256, 1), 8)
|
||||
if column_major_scales:
|
||||
_h3_per_token_group_quant_fp8_column_major[(num_groups,)](
|
||||
input_tensor,
|
||||
quantized,
|
||||
scales,
|
||||
group_size,
|
||||
columns,
|
||||
scales.stride(1),
|
||||
1e-10,
|
||||
-448.0,
|
||||
448.0,
|
||||
BLOCK=block,
|
||||
num_warps=num_warps,
|
||||
num_stages=1,
|
||||
)
|
||||
else:
|
||||
_h3_per_token_group_quant_fp8_row_major[(num_groups,)](
|
||||
input_tensor,
|
||||
quantized,
|
||||
scales,
|
||||
group_size,
|
||||
1e-10,
|
||||
-448.0,
|
||||
448.0,
|
||||
BLOCK=block,
|
||||
num_warps=num_warps,
|
||||
num_stages=1,
|
||||
)
|
||||
return quantized, scales
|
||||
|
||||
|
||||
def _get_flashinfer_groupwise_fp8_gemm():
|
||||
try:
|
||||
from flashinfer.gemm import gemm_fp8_nt_groupwise
|
||||
except (AttributeError, ImportError) as error:
|
||||
raise RuntimeError(
|
||||
"MiniMax-H3 serialized blockwise FP8 requires "
|
||||
"flashinfer.gemm.gemm_fp8_nt_groupwise (validated with flashinfer-python==0.6.8). "
|
||||
"FastVideo will not re-quantize this checkpoint to tensorwise FP8.") from error
|
||||
return gemm_fp8_nt_groupwise
|
||||
|
||||
|
||||
def _get_flashinfer_groupwise_backend(device: torch.device) -> str:
|
||||
capability = torch.cuda.get_device_capability(device)
|
||||
if capability[0] >= 12:
|
||||
return "cutlass"
|
||||
if capability[0] == 10:
|
||||
return "trtllm"
|
||||
capability_number = capability[0] * 10 + capability[1]
|
||||
raise RuntimeError(f"FlashInfer groupwise FP8 requires a Blackwell GPU, got sm{capability_number}")
|
||||
|
||||
|
||||
def _flashinfer_gemm_w8a8_block_fp8_linear_with_fallback(
|
||||
input_tensor: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
block_size: tuple[int, int],
|
||||
weight_scale: torch.Tensor,
|
||||
bias: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
input_2d = input_tensor.view(-1, input_tensor.shape[-1])
|
||||
output_shape = [*input_tensor.shape[:-1], weight.shape[0]]
|
||||
backend = _get_flashinfer_groupwise_backend(input_tensor.device)
|
||||
if input_2d.dtype != torch.bfloat16:
|
||||
raise RuntimeError("MiniMax-H3 FlashInfer groupwise FP8 requires BF16 activations; "
|
||||
f"got {input_2d.dtype}. The SGLang FP16 Triton GEMM fallback is not enabled for H3.")
|
||||
if backend == "trtllm" and input_2d.shape[1] < 256:
|
||||
raise RuntimeError("MiniMax-H3 FlashInfer TRTLLM groupwise FP8 requires K >= 256; "
|
||||
f"got K={input_2d.shape[1]}. The SGLang Triton GEMM fallback is not enabled for H3.")
|
||||
|
||||
gemm_fp8_nt_groupwise = _get_flashinfer_groupwise_fp8_gemm()
|
||||
block_n, block_k = block_size
|
||||
q_input, x_scale = _sglang_per_token_group_quant_fp8(
|
||||
input_2d,
|
||||
block_k,
|
||||
column_major_scales=(backend == "trtllm"),
|
||||
)
|
||||
if backend == "cutlass":
|
||||
m, k = input_2d.shape
|
||||
n = weight.shape[0]
|
||||
expected_x_scale_shape = (k // block_k, m)
|
||||
expected_weight_scale_shape = (k // block_k, n // block_n)
|
||||
if x_scale.shape == (m, k // block_k):
|
||||
x_scale = x_scale.transpose(-1, -2).contiguous()
|
||||
if weight_scale.shape == (n // block_n, k // block_k):
|
||||
weight_scale = weight_scale.transpose(-1, -2).contiguous()
|
||||
if x_scale.shape != expected_x_scale_shape or weight_scale.shape != expected_weight_scale_shape:
|
||||
raise RuntimeError("FlashInfer CUTLASS block-FP8 scale layout mismatch: "
|
||||
f"x_scale={tuple(x_scale.shape)}, weight_scale={tuple(weight_scale.shape)}, "
|
||||
f"expected={expected_x_scale_shape}/{expected_weight_scale_shape}")
|
||||
if x_scale.dtype != torch.float32 or weight_scale.dtype != torch.float32:
|
||||
raise RuntimeError("FlashInfer CUTLASS block-FP8 scales must be float32")
|
||||
output = gemm_fp8_nt_groupwise(
|
||||
q_input,
|
||||
weight,
|
||||
x_scale.contiguous(),
|
||||
weight_scale.contiguous(),
|
||||
out_dtype=input_2d.dtype,
|
||||
backend="cutlass",
|
||||
scale_major_mode="MN",
|
||||
)
|
||||
else:
|
||||
expected_x_scale_shape = (input_2d.shape[0], input_2d.shape[1] // block_k)
|
||||
expected_weight_scale_shape = (weight.shape[0] // block_n, weight.shape[1] // block_k)
|
||||
if x_scale.shape != expected_x_scale_shape or x_scale.stride(0) != 1:
|
||||
raise RuntimeError("FlashInfer TRTLLM block-FP8 activation scale layout mismatch: "
|
||||
f"shape={tuple(x_scale.shape)}, stride={x_scale.stride()}, "
|
||||
f"expected column-major {expected_x_scale_shape}")
|
||||
if weight_scale.shape != expected_weight_scale_shape:
|
||||
raise RuntimeError("FlashInfer TRTLLM block-FP8 weight scale layout mismatch: "
|
||||
f"shape={tuple(weight_scale.shape)}, expected={expected_weight_scale_shape}")
|
||||
output = gemm_fp8_nt_groupwise(
|
||||
q_input,
|
||||
weight,
|
||||
x_scale,
|
||||
weight_scale,
|
||||
out_dtype=input_2d.dtype,
|
||||
backend="trtllm",
|
||||
)
|
||||
if bias is not None:
|
||||
output += bias
|
||||
return output.to(dtype=input_2d.dtype).view(*output_shape)
|
||||
|
||||
|
||||
class MiniMaxH3SerializedFP8LinearMethod(LinearMethodBase):
|
||||
"""Execute serialized 128x128 block-FP8 weights without re-quantizing them."""
|
||||
|
||||
def __init__(self, weight_block_size: tuple[int, int]) -> None:
|
||||
super().__init__()
|
||||
self.weight_block_size = weight_block_size
|
||||
|
||||
def create_weights(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
input_size_per_partition: int,
|
||||
output_partition_sizes: list[int],
|
||||
input_size: int,
|
||||
output_size: int,
|
||||
params_dtype: torch.dtype,
|
||||
**extra_weight_attrs,
|
||||
) -> None:
|
||||
output_size_per_partition = sum(output_partition_sizes)
|
||||
block_n, block_k = self.weight_block_size
|
||||
tp_size = get_tp_world_size()
|
||||
if tp_size > 1 and input_size // input_size_per_partition == tp_size:
|
||||
if input_size_per_partition % block_k:
|
||||
raise ValueError(f"Weight input_size_per_partition={input_size_per_partition} is not divisible "
|
||||
f"by block_k={block_k}")
|
||||
if tp_size > 1 and output_size // output_size_per_partition == tp_size:
|
||||
for output_partition_size in output_partition_sizes:
|
||||
if output_partition_size % block_n:
|
||||
raise ValueError(f"Weight output_partition_size={output_partition_size} is not divisible "
|
||||
f"by block_n={block_n}")
|
||||
|
||||
layer.logical_widths = output_partition_sizes
|
||||
layer.input_size_per_partition = input_size_per_partition
|
||||
layer.output_size_per_partition = output_size_per_partition
|
||||
layer.orig_dtype = params_dtype
|
||||
|
||||
weight_loader = extra_weight_attrs.get("weight_loader")
|
||||
weight = Parameter(
|
||||
torch.empty(output_size_per_partition, input_size_per_partition, dtype=FP8_DTYPE),
|
||||
requires_grad=False,
|
||||
)
|
||||
set_weight_attrs(weight, {
|
||||
"input_dim": 1,
|
||||
"output_dim": 0,
|
||||
"weight_loader": weight_loader,
|
||||
})
|
||||
layer.register_parameter("weight", weight)
|
||||
|
||||
scale = Parameter(
|
||||
torch.empty((output_size_per_partition + block_n - 1) // block_n,
|
||||
(input_size_per_partition + block_k - 1) // block_k,
|
||||
dtype=torch.float32),
|
||||
requires_grad=False,
|
||||
)
|
||||
set_weight_attrs(scale, {
|
||||
"input_dim": 1,
|
||||
"output_dim": 0,
|
||||
"weight_loader": weight_loader,
|
||||
})
|
||||
scale.data.fill_(torch.finfo(torch.float32).min)
|
||||
layer.register_parameter("weight_scale_inv", scale)
|
||||
layer.register_parameter("input_scale", None)
|
||||
|
||||
def process_weights_after_loading(self, layer: nn.Module) -> None:
|
||||
weight = getattr(layer, "weight", None)
|
||||
block_scales = getattr(layer, "weight_scale_inv", None)
|
||||
if weight is None or block_scales is None:
|
||||
raise ValueError("Serialized MiniMax-H3 FP8 linear is missing weight or weight_scale_inv")
|
||||
if weight.dtype != FP8_DTYPE:
|
||||
raise ValueError(f"Serialized MiniMax-H3 FP8 weight must be {FP8_DTYPE}, got {weight.dtype}")
|
||||
if block_scales.dtype != torch.float32:
|
||||
raise ValueError("Serialized MiniMax-H3 FP8 weight_scale_inv must be float32, "
|
||||
f"got {block_scales.dtype}")
|
||||
|
||||
block_n, block_k = self.weight_block_size
|
||||
output_size, input_size = weight.shape
|
||||
if output_size % block_n or input_size % block_k:
|
||||
raise ValueError("Serialized MiniMax-H3 FP8 weight dimensions must be divisible by the 128x128 block size; "
|
||||
f"got {tuple(weight.shape)}")
|
||||
expected_scale_shape = (output_size // block_n, input_size // block_k)
|
||||
if tuple(block_scales.shape) != expected_scale_shape:
|
||||
raise ValueError("Serialized MiniMax-H3 FP8 scale shape mismatch: "
|
||||
f"expected {expected_scale_shape}, got {tuple(block_scales.shape)}")
|
||||
if not bool(torch.isfinite(block_scales).all()) or bool((block_scales <= 0).any()):
|
||||
raise ValueError("Serialized MiniMax-H3 FP8 weight_scale_inv must contain finite positive values")
|
||||
layer.weight.data = weight.data
|
||||
layer.weight_scale_inv.data = block_scales.data
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer: torch.nn.Module,
|
||||
x: torch.Tensor,
|
||||
bias: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
if x.device.type != "cuda":
|
||||
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 execution requires CUDA")
|
||||
|
||||
capability = torch.cuda.get_device_capability(x.device)
|
||||
capability_number = capability[0] * 10 + capability[1]
|
||||
if capability_number < MiniMaxH3SerializedFP8Config.get_min_capability():
|
||||
raise RuntimeError("MiniMax-H3 serialized blockwise FP8 requires GPU capability "
|
||||
f"sm{MiniMaxH3SerializedFP8Config.get_min_capability()} or newer, "
|
||||
f"got sm{capability_number}")
|
||||
|
||||
if not x.is_contiguous():
|
||||
x = x.contiguous()
|
||||
return _flashinfer_gemm_w8a8_block_fp8_linear_with_fallback(
|
||||
x,
|
||||
layer.weight,
|
||||
self.weight_block_size,
|
||||
layer.weight_scale_inv,
|
||||
bias,
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"MiniMaxH3SerializedFP8Config",
|
||||
"MiniMaxH3SerializedFP8LinearMethod",
|
||||
]
|
||||
@@ -8,13 +8,13 @@ import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput
|
||||
from fastvideo.configs.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConfig
|
||||
from fastvideo.distributed import get_tp_world_size
|
||||
from fastvideo.layers.layernorm import RMSNorm
|
||||
from fastvideo.layers.linear import ColumnParallelLinear, RowParallelLinear
|
||||
from fastvideo.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||
from fastvideo.models.encoders.base import TextEncoder
|
||||
from fastvideo.models.encoders.minimax_h3_checkpoint_fp8 import MiniMaxH3SerializedFP8Config
|
||||
from fastvideo.models.loader.weight_utils import default_weight_loader
|
||||
|
||||
|
||||
@@ -227,19 +227,13 @@ class MiniMaxH3Qwen3VLLanguageModel(nn.Module):
|
||||
org_num_embeddings=config.vocab_size,
|
||||
quant_config=quant_config,
|
||||
)
|
||||
# Build only as far as the consumer reads. The hidden-state tuple records
|
||||
# each layer's input, so stopping after N layers still yields entry N,
|
||||
# the output of layer N-1, unchanged. Everything above it exists only to
|
||||
# feed `last_hidden_state`, which nothing consumes.
|
||||
override = config.num_hidden_layers_override
|
||||
self.num_layers = (config.num_hidden_layers
|
||||
if override is None else min(config.num_hidden_layers, override))
|
||||
self.output_hidden_state_index = config.output_hidden_state_index
|
||||
self.layers = nn.ModuleList(
|
||||
MiniMaxH3Qwen3VLTextDecoderLayer(config, prefix=f"{config.prefix}.language_model.layers.{index}")
|
||||
for index in range(self.num_layers))
|
||||
# The final norm sits above the tapped layer, so a truncated stack drops
|
||||
# it. Keeping it would overwrite the tapped entry with a normalised
|
||||
# tensor and change conditioning without raising anything.
|
||||
self.norm = (RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||
if self.num_layers == config.num_hidden_layers else None)
|
||||
self.rotary_emb = MiniMaxH3Qwen3VLTextRotaryEmbedding(config)
|
||||
@@ -249,18 +243,14 @@ class MiniMaxH3Qwen3VLLanguageModel(nn.Module):
|
||||
inputs_embeds: torch.Tensor,
|
||||
position_ids: torch.Tensor,
|
||||
attention_mask: torch.Tensor | None,
|
||||
output_hidden_states: bool,
|
||||
visual_pos_masks: torch.Tensor | None,
|
||||
deepstack_visual_embeds: list[torch.Tensor] | None,
|
||||
) -> BaseEncoderOutput:
|
||||
) -> torch.Tensor:
|
||||
if attention_mask is not None and bool(attention_mask.to(torch.bool).all()):
|
||||
attention_mask = None
|
||||
position_embeddings = self.rotary_emb(inputs_embeds, position_ids)
|
||||
hidden_states = inputs_embeds
|
||||
all_hidden_states: tuple[torch.Tensor, ...] | None = () if output_hidden_states else None
|
||||
for layer_index, layer in enumerate(self.layers):
|
||||
if all_hidden_states is not None:
|
||||
all_hidden_states += (hidden_states, )
|
||||
hidden_states = layer(hidden_states, position_embeddings, attention_mask)
|
||||
if deepstack_visual_embeds is not None and layer_index < len(deepstack_visual_embeds):
|
||||
if visual_pos_masks is None:
|
||||
@@ -269,13 +259,9 @@ class MiniMaxH3Qwen3VLLanguageModel(nn.Module):
|
||||
visual = deepstack_visual_embeds[layer_index].to(hidden_states.device, hidden_states.dtype)
|
||||
updated = hidden_states[mask].clone() + visual
|
||||
hidden_states[mask] = updated
|
||||
if self.norm is not None:
|
||||
hidden_states = self.norm(hidden_states)
|
||||
# Truncated or not, the last entry is appended here, so the tapped index
|
||||
# lands in the same place either way.
|
||||
if all_hidden_states is not None:
|
||||
all_hidden_states += (hidden_states, )
|
||||
return BaseEncoderOutput(last_hidden_state=hidden_states, hidden_states=all_hidden_states)
|
||||
if layer_index + 1 == self.output_hidden_state_index:
|
||||
return hidden_states
|
||||
raise RuntimeError(f"MiniMax-H3 text stack did not reach hidden_states[{self.output_hidden_state_index}]")
|
||||
|
||||
|
||||
class MiniMaxH3Qwen3VLVisionPatchEmbed(nn.Module):
|
||||
@@ -513,10 +499,18 @@ class MiniMaxH3Qwen3VLVisionModel(nn.Module):
|
||||
return self.merger(hidden_states), deepstack_features
|
||||
|
||||
|
||||
class MiniMaxH3Qwen3VLConditioner(TextEncoder):
|
||||
"""FastVideo-native Qwen3-VL body without the unused language-model head."""
|
||||
class MiniMaxH3Qwen3VLConditioner(TextEncoder[torch.Tensor]):
|
||||
"""H3 conditioner returning the unnormalized layer-50 hidden tensor."""
|
||||
|
||||
supports_hf_from_pretrained = False
|
||||
supported_checkpoint_quantization_methods = frozenset({"fp8"})
|
||||
|
||||
@classmethod
|
||||
def checkpoint_quantization_config_from_metadata(
|
||||
cls,
|
||||
metadata: dict[str, Any],
|
||||
) -> MiniMaxH3SerializedFP8Config:
|
||||
return MiniMaxH3SerializedFP8Config.from_config(metadata)
|
||||
|
||||
def __init__(self, config: MiniMaxH3Qwen3VLConfig) -> None:
|
||||
super().__init__(config)
|
||||
@@ -530,15 +524,12 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
|
||||
|
||||
@property
|
||||
def num_hidden_layers(self) -> int:
|
||||
"""The checkpoint architecture's nominal depth, matching its config.json.
|
||||
|
||||
When ``num_hidden_layers_override`` truncates the stack at the
|
||||
conditioning tap, fewer layers exist; the built count is
|
||||
``self.language_model.num_layers``, and the hidden-state tuple has
|
||||
``num_layers + 1`` entries, not ``num_hidden_layers + 1``.
|
||||
"""
|
||||
return self.config.num_hidden_layers
|
||||
|
||||
@property
|
||||
def num_built_hidden_layers(self) -> int:
|
||||
return self.language_model.num_layers
|
||||
|
||||
def _get_rope_index(
|
||||
self,
|
||||
input_ids: torch.Tensor,
|
||||
@@ -631,35 +622,39 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
|
||||
f"tokens={int(mask.sum())}, features={features.shape[0]}")
|
||||
return mask
|
||||
|
||||
def forward(
|
||||
# no_grad, NOT inference_mode: with text_encoder_cpu_offload=True (the
|
||||
# FastVideoArgs default) the loader FSDP2-shards this conditioner, and
|
||||
# FSDP2's wait_for_unshard reads tensor._version via
|
||||
# _unsafe_preserve_version_counter - inference tensors do not track
|
||||
# version counters, so inference_mode crashes the first encode. no_grad
|
||||
# frees the same activation memory and keeps prompt_embeds ordinary
|
||||
# tensors (safe for any future backward through the conditioning).
|
||||
@torch.no_grad()
|
||||
def encode_ids(
|
||||
self,
|
||||
input_ids: torch.Tensor | None,
|
||||
position_ids: torch.Tensor | None = None,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
inputs_embeds: torch.Tensor | None = None,
|
||||
output_hidden_states: bool | None = None,
|
||||
input_ids: torch.Tensor,
|
||||
*,
|
||||
pixel_values: torch.Tensor | None = None,
|
||||
pixel_values_videos: torch.Tensor | None = None,
|
||||
image_grid_thw: torch.Tensor | None = None,
|
||||
pixel_values_videos: torch.Tensor | None = None,
|
||||
video_grid_thw: torch.Tensor | None = None,
|
||||
mm_token_type_ids: torch.Tensor | None = None,
|
||||
**kwargs: Any,
|
||||
) -> BaseEncoderOutput:
|
||||
del mm_token_type_ids, kwargs
|
||||
if (input_ids is None) == (inputs_embeds is None):
|
||||
raise ValueError("Exactly one of input_ids or inputs_embeds is required")
|
||||
if inputs_embeds is None:
|
||||
assert input_ids is not None
|
||||
inputs_embeds = self.language_model.embed_tokens(input_ids)
|
||||
if input_ids is None and (pixel_values is not None or pixel_values_videos is not None):
|
||||
raise ValueError("Multimodal Qwen3-VL inputs require input_ids for placeholder matching")
|
||||
) -> torch.Tensor:
|
||||
if input_ids.ndim != 1:
|
||||
raise ValueError(f"MiniMax-H3 slim forward expects 1-D input_ids, got shape={tuple(input_ids.shape)}")
|
||||
if (pixel_values is None) != (image_grid_thw is None):
|
||||
raise ValueError("pixel_values and image_grid_thw must be provided together")
|
||||
if (pixel_values_videos is None) != (video_grid_thw is None):
|
||||
raise ValueError("pixel_values_videos and video_grid_thw must be provided together")
|
||||
|
||||
input_ids = input_ids.unsqueeze(0)
|
||||
inputs_embeds = self.language_model.embed_tokens(input_ids)
|
||||
|
||||
image_mask = None
|
||||
video_mask = None
|
||||
image_deepstack = None
|
||||
video_deepstack = None
|
||||
if pixel_values is not None:
|
||||
if input_ids is None or image_grid_thw is None:
|
||||
if image_grid_thw is None:
|
||||
raise ValueError("pixel_values require input_ids and image_grid_thw")
|
||||
image_features, image_deepstack = self._visual_features(pixel_values, image_grid_thw)
|
||||
image_features = image_features.to(inputs_embeds.device, inputs_embeds.dtype)
|
||||
@@ -667,7 +662,7 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
|
||||
"image")
|
||||
inputs_embeds = inputs_embeds.masked_scatter(image_mask.unsqueeze(-1), image_features)
|
||||
if pixel_values_videos is not None:
|
||||
if input_ids is None or video_grid_thw is None:
|
||||
if video_grid_thw is None:
|
||||
raise ValueError("pixel_values_videos require input_ids and video_grid_thw")
|
||||
video_features, video_deepstack = self._visual_features(pixel_values_videos, video_grid_thw)
|
||||
video_features = video_features.to(inputs_embeds.device, inputs_embeds.dtype)
|
||||
@@ -695,50 +690,34 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
|
||||
visual_mask = video_mask
|
||||
deepstack_features = video_deepstack
|
||||
|
||||
if position_ids is None:
|
||||
if input_ids is None:
|
||||
sequence_length = inputs_embeds.shape[1]
|
||||
position_ids = torch.arange(sequence_length,
|
||||
device=inputs_embeds.device).view(1, 1,
|
||||
-1).expand(3, inputs_embeds.shape[0], -1)
|
||||
else:
|
||||
position_ids = self._get_rope_index(input_ids, image_grid_thw, video_grid_thw, attention_mask)
|
||||
output_hidden_states = self.config.output_hidden_states if output_hidden_states is None else output_hidden_states
|
||||
outputs = self.language_model(
|
||||
position_ids = self._get_rope_index(input_ids, image_grid_thw, video_grid_thw, None)
|
||||
hidden_states = self.language_model(
|
||||
inputs_embeds,
|
||||
position_ids,
|
||||
attention_mask,
|
||||
output_hidden_states,
|
||||
None,
|
||||
visual_mask,
|
||||
deepstack_features,
|
||||
)
|
||||
outputs.attention_mask = attention_mask
|
||||
return outputs
|
||||
if hidden_states.ndim != 3 or hidden_states.shape[0] != 1:
|
||||
raise RuntimeError(f"MiniMax-H3 language model returned unexpected shape={tuple(hidden_states.shape)}")
|
||||
return hidden_states[0]
|
||||
|
||||
def _is_above_the_tap(self, name: str) -> bool:
|
||||
"""Whether this checkpoint key belongs to a layer we did not build.
|
||||
|
||||
A truncated language stack still ships every layer in the checkpoint, and
|
||||
the unexpected-key check below is strict on purpose, so the surplus keys
|
||||
have to be dropped here rather than by relaxing it.
|
||||
"""
|
||||
language_model = self.language_model
|
||||
# The final norm is dropped exactly when the stack is truncated, so its
|
||||
# absence is the signal.
|
||||
if language_model.norm is not None:
|
||||
return False
|
||||
if name == "language_model.norm.weight":
|
||||
return True
|
||||
prefix = "language_model.layers."
|
||||
if not name.startswith(prefix):
|
||||
return False
|
||||
index = name[len(prefix):].split(".", 1)[0]
|
||||
if not index.isdigit():
|
||||
return False
|
||||
# Only drop indexes the full stack would have built. Anything at or
|
||||
# above the checkpoint's own num_hidden_layers is corrupt and must
|
||||
# still raise below, exactly as it does without truncation.
|
||||
return language_model.num_layers <= int(index) < self.config.num_hidden_layers
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor,
|
||||
*,
|
||||
pixel_values: torch.Tensor | None = None,
|
||||
image_grid_thw: torch.Tensor | None = None,
|
||||
pixel_values_videos: torch.Tensor | None = None,
|
||||
video_grid_thw: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
return self.encode_ids(
|
||||
input_ids,
|
||||
pixel_values=pixel_values,
|
||||
image_grid_thw=image_grid_thw,
|
||||
pixel_values_videos=pixel_values_videos,
|
||||
video_grid_thw=video_grid_thw,
|
||||
)
|
||||
|
||||
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
|
||||
parameters = dict(self.named_parameters())
|
||||
@@ -748,7 +727,7 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
|
||||
if source_name == "lm_head.weight":
|
||||
continue
|
||||
name = source_name[6:] if source_name.startswith("model.") else source_name
|
||||
if self._is_above_the_tap(name):
|
||||
if self._is_omitted_checkpoint_key(name):
|
||||
continue
|
||||
if name not in parameters:
|
||||
raise ValueError(f"Unexpected MiniMax-H3 Qwen3-VL checkpoint key: {source_name}")
|
||||
@@ -758,7 +737,23 @@ class MiniMaxH3Qwen3VLConditioner(TextEncoder):
|
||||
loaded.add(name)
|
||||
return loaded
|
||||
|
||||
def _is_omitted_checkpoint_key(self, name: str) -> bool:
|
||||
"""Return whether a valid checkpoint key belongs to an unbuilt layer."""
|
||||
language_model = self.language_model
|
||||
if language_model.norm is not None:
|
||||
return False
|
||||
if name == "language_model.norm.weight":
|
||||
return True
|
||||
prefix = "language_model.layers."
|
||||
if not name.startswith(prefix):
|
||||
return False
|
||||
index = name[len(prefix):].split(".", 1)[0]
|
||||
return (index.isdigit() and language_model.num_layers <= int(index) < self.config.num_hidden_layers)
|
||||
|
||||
|
||||
EntryClass = MiniMaxH3Qwen3VLConditioner
|
||||
|
||||
__all__ = ["MiniMaxH3Qwen3VLConditioner"]
|
||||
__all__ = [
|
||||
"MiniMaxH3Qwen3VLConditioner",
|
||||
"MiniMaxH3SerializedFP8Config",
|
||||
]
|
||||
|
||||
@@ -9,7 +9,7 @@ from abc import ABC, abstractmethod
|
||||
from collections.abc import Generator, Iterable
|
||||
from contextlib import nullcontext
|
||||
from copy import deepcopy
|
||||
from typing import cast
|
||||
from typing import Any, cast
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
@@ -30,9 +30,13 @@ from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.layers.quantization import get_quantization_config
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.encoders.base import TextEncoder
|
||||
from fastvideo.models.hf_transformer_utils import get_diffusers_config
|
||||
from fastvideo.models.loader.fsdp_load import maybe_load_fsdp_model, shard_model
|
||||
from fastvideo.models.loader.text_encoder_quantization import (
|
||||
_configure_text_encoder_quantization,
|
||||
_process_quantized_text_encoder_weights,
|
||||
_resolve_text_encoder_checkpoint_path,
|
||||
)
|
||||
from fastvideo.models.loader.utils import set_default_torch_dtype
|
||||
from fastvideo.models.loader.weight_utils import (
|
||||
filter_duplicate_safetensors_files,
|
||||
@@ -347,22 +351,46 @@ class TextEncoderLoader(ComponentLoader):
|
||||
if cpu_offload is None:
|
||||
cpu_offload = fastvideo_args.text_encoder_cpu_offload
|
||||
use_cpu_offload = (cpu_offload and len(getattr(model_config, "_fsdp_shard_conditions", [])) > 0)
|
||||
runtime_device = get_local_torch_device()
|
||||
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
if cpu_offload:
|
||||
target_device = (torch.device("mps") if current_platform.is_mps() else torch.device("cpu"))
|
||||
|
||||
# Set quantization config if specified
|
||||
if (use_text_encoder_override and fastvideo_args.override_text_encoder_quant is not None):
|
||||
if fastvideo_args.override_text_encoder_safetensors is None:
|
||||
raise ValueError("override_text_encoder_quant is set but override_text_encoder_safetensors is None")
|
||||
quant_cls = get_quantization_config(fastvideo_args.override_text_encoder_quant)
|
||||
model_config.quant_config = quant_cls()
|
||||
|
||||
with set_default_torch_dtype(PRECISION_TO_TYPE[dtype]):
|
||||
architectures = getattr(model_config, "architectures", [])
|
||||
model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
|
||||
checkpoint_path = _resolve_text_encoder_checkpoint_path(
|
||||
model_path,
|
||||
fastvideo_args,
|
||||
use_text_encoder_override,
|
||||
)
|
||||
checkpoint_quant_config = _configure_text_encoder_quantization(
|
||||
model_config,
|
||||
model_cls,
|
||||
checkpoint_path,
|
||||
)
|
||||
if checkpoint_quant_config is not None:
|
||||
if fastvideo_args.override_text_encoder_quant is not None:
|
||||
raise ValueError("Serialized checkpoint quantization is selected from checkpoint metadata; "
|
||||
"override_text_encoder_quant is an online conversion option and must be unset")
|
||||
requested_dtype = PRECISION_TO_TYPE[dtype]
|
||||
if requested_dtype not in checkpoint_quant_config.get_supported_act_dtypes():
|
||||
raise ValueError(f"Serialized {checkpoint_quant_config.get_name()} text encoder does not support "
|
||||
f"activation dtype {requested_dtype}")
|
||||
checkpoint_quant_config.validate_runtime(runtime_device)
|
||||
logger.info(
|
||||
"Selected serialized %s text-encoder checkpoint execution from %s",
|
||||
checkpoint_quant_config.get_name(),
|
||||
checkpoint_path,
|
||||
)
|
||||
elif use_text_encoder_override and fastvideo_args.override_text_encoder_quant is not None:
|
||||
if fastvideo_args.override_text_encoder_safetensors is None:
|
||||
raise ValueError("override_text_encoder_quant is set but override_text_encoder_safetensors is None")
|
||||
quant_cls = get_quantization_config(fastvideo_args.override_text_encoder_quant)
|
||||
model_config.quant_config = quant_cls()
|
||||
|
||||
if getattr(model_cls, "supports_hf_from_pretrained", False):
|
||||
model = model_cls.from_pretrained_local( # type: ignore[attr-defined]
|
||||
model_path,
|
||||
@@ -381,11 +409,20 @@ class TextEncoderLoader(ComponentLoader):
|
||||
|
||||
weights_to_load = {name for name, _ in model.named_parameters()}
|
||||
if (use_text_encoder_override and fastvideo_args.override_text_encoder_safetensors is not None):
|
||||
loaded_weights: set[str] = model.load_weights(
|
||||
safetensors_weights_iterator(
|
||||
[fastvideo_args.override_text_encoder_safetensors],
|
||||
if os.path.isdir(checkpoint_path):
|
||||
override_weights = self._get_all_weights(
|
||||
model,
|
||||
checkpoint_path,
|
||||
to_cpu=bool(cpu_offload),
|
||||
)
|
||||
else:
|
||||
if self.counter_before_loading_weights == 0.0:
|
||||
self.counter_before_loading_weights = time.perf_counter()
|
||||
override_weights = safetensors_weights_iterator(
|
||||
[checkpoint_path],
|
||||
to_cpu=use_cpu_offload,
|
||||
)) # type: ignore
|
||||
)
|
||||
loaded_weights: set[str] = model.load_weights(override_weights) # type: ignore
|
||||
else:
|
||||
loaded_weights: set[str] = model.load_weights(
|
||||
self._get_all_weights(
|
||||
@@ -400,6 +437,10 @@ class TextEncoderLoader(ComponentLoader):
|
||||
self.counter_after_loading_weights - self.counter_before_loading_weights,
|
||||
)
|
||||
|
||||
if checkpoint_quant_config is not None:
|
||||
processed_linears = _process_quantized_text_encoder_weights(model, runtime_device)
|
||||
logger.info("Validated %d serialized blockwise FP8 text-encoder linears", processed_linears)
|
||||
|
||||
# Explicitly move model to target device after loading weights
|
||||
model = model.to(target_device)
|
||||
|
||||
@@ -442,7 +483,7 @@ class TextEncoderLoader(ComponentLoader):
|
||||
# that have loaded weights tracking currently.
|
||||
# if loaded_weights is not None:
|
||||
weights_not_loaded = weights_to_load - loaded_weights
|
||||
if weights_not_loaded and model_config.quant_config is None:
|
||||
if weights_not_loaded and (model_config.quant_config is None or checkpoint_quant_config is not None):
|
||||
raise ValueError("Following weights were not initialized from "
|
||||
f"checkpoint: {weights_not_loaded}")
|
||||
|
||||
@@ -1057,7 +1098,12 @@ class TransformerLoader(ComponentLoader):
|
||||
# so recording here makes the decision readable from the loaded
|
||||
# transformer — and records the narrowed one for teacher/critic.
|
||||
resolved = record_resolved_attention_backend(dit_config)
|
||||
logger.info("transformer attention backend: %s", resolved.name if resolved else "automatic selection")
|
||||
# Every worker records its resolved backend so distributed profile
|
||||
# snapshots can prove that all ranks use the requested kernels.
|
||||
logger.info("Worker %s transformer attention backend: %s",
|
||||
os.environ.get("RANK", "0"),
|
||||
resolved.name if resolved else "automatic selection",
|
||||
local_main_process_only=False)
|
||||
model = maybe_load_fsdp_model(
|
||||
model_cls=model_cls,
|
||||
init_params={
|
||||
|
||||
@@ -0,0 +1,127 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Checkpoint-serialized quantization lifecycle for native text encoders."""
|
||||
|
||||
import json
|
||||
import os
|
||||
from itertools import chain
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from safetensors.torch import safe_open
|
||||
|
||||
from fastvideo.configs.models import EncoderConfig
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.layers.linear import LinearBase, UnquantizedLinearMethod
|
||||
from fastvideo.layers.quantization.base_config import QuantizationConfig
|
||||
from fastvideo.models.encoders.base import TextEncoder
|
||||
|
||||
|
||||
def _resolve_text_encoder_checkpoint_path(
|
||||
model_path: str,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
use_text_encoder_override: bool,
|
||||
) -> str:
|
||||
override = fastvideo_args.override_text_encoder_safetensors if use_text_encoder_override else None
|
||||
checkpoint_path = override or model_path
|
||||
if not os.path.exists(checkpoint_path):
|
||||
raise FileNotFoundError(f"Text-encoder checkpoint does not exist: {checkpoint_path}")
|
||||
if not os.path.isdir(checkpoint_path) and not os.path.isfile(checkpoint_path):
|
||||
raise ValueError(f"Text-encoder checkpoint must be a file or directory: {checkpoint_path}")
|
||||
return checkpoint_path
|
||||
|
||||
|
||||
def _read_text_encoder_checkpoint_quantization_config(checkpoint_path: str) -> dict[str, Any] | None:
|
||||
checkpoint_dir = checkpoint_path if os.path.isdir(checkpoint_path) else os.path.dirname(checkpoint_path)
|
||||
config_path = os.path.join(checkpoint_dir, "config.json")
|
||||
if os.path.isfile(config_path):
|
||||
try:
|
||||
with open(config_path, encoding="utf-8") as config_file:
|
||||
checkpoint_config = json.load(config_file)
|
||||
except json.JSONDecodeError as error:
|
||||
raise ValueError(f"Invalid text-encoder checkpoint config: {config_path}") from error
|
||||
quantization_config = checkpoint_config.get("quantization_config")
|
||||
if quantization_config is not None:
|
||||
if not isinstance(quantization_config, dict):
|
||||
raise ValueError(f"quantization_config in {config_path} must be an object")
|
||||
return quantization_config
|
||||
|
||||
if not os.path.isfile(checkpoint_path) or not checkpoint_path.endswith(".safetensors"):
|
||||
return None
|
||||
with safe_open(checkpoint_path, framework="pt", device="cpu") as checkpoint_file:
|
||||
metadata = checkpoint_file.metadata() or {}
|
||||
for key in ("quantization_config", "_quantization_metadata"):
|
||||
serialized = metadata.get(key)
|
||||
if serialized is None:
|
||||
continue
|
||||
try:
|
||||
quantization_config = json.loads(serialized)
|
||||
except json.JSONDecodeError as error:
|
||||
raise ValueError(f"Invalid {key} metadata in {checkpoint_path}") from error
|
||||
if not isinstance(quantization_config, dict):
|
||||
raise ValueError(f"{key} metadata in {checkpoint_path} must decode to an object")
|
||||
return quantization_config
|
||||
return None
|
||||
|
||||
|
||||
def _configure_text_encoder_quantization(
|
||||
model_config: EncoderConfig,
|
||||
model_cls: type[nn.Module],
|
||||
checkpoint_path: str,
|
||||
) -> QuantizationConfig | None:
|
||||
if not issubclass(model_cls, TextEncoder):
|
||||
return None
|
||||
checkpoint_quantization = _read_text_encoder_checkpoint_quantization_config(checkpoint_path)
|
||||
if checkpoint_quantization is None:
|
||||
return None
|
||||
|
||||
quant_method = str(checkpoint_quantization.get("quant_method", "")).lower()
|
||||
if not quant_method:
|
||||
raise ValueError(f"Quantized text-encoder checkpoint {checkpoint_path} does not declare quant_method")
|
||||
supported_methods = getattr(model_cls, "supported_checkpoint_quantization_methods", frozenset())
|
||||
if quant_method not in supported_methods:
|
||||
supported = ", ".join(sorted(supported_methods)) or "none"
|
||||
raise ValueError(f"Text encoder {model_cls.__name__} does not support serialized {quant_method!r} "
|
||||
f"checkpoints (supported: {supported})")
|
||||
|
||||
factory = getattr(model_cls, "checkpoint_quantization_config_from_metadata", None)
|
||||
if not callable(factory):
|
||||
raise ValueError(f"Text encoder {model_cls.__name__} advertises serialized {quant_method!r} support "
|
||||
"without a checkpoint quantization factory")
|
||||
quant_config = factory(checkpoint_quantization)
|
||||
model_config.quant_config = quant_config
|
||||
return quant_config
|
||||
|
||||
|
||||
def _module_tensor_device(module: nn.Module) -> torch.device | None:
|
||||
devices = {
|
||||
tensor.device
|
||||
for tensor in chain(
|
||||
module.parameters(recurse=False),
|
||||
module.buffers(recurse=False),
|
||||
)
|
||||
}
|
||||
if len(devices) > 1:
|
||||
raise ValueError(f"Quantized text-encoder module {type(module).__name__} spans multiple devices: {devices}")
|
||||
return next(iter(devices), None)
|
||||
|
||||
|
||||
def _process_quantized_text_encoder_weights(model: nn.Module, process_device: torch.device) -> int:
|
||||
"""Run quantized post-load hooks one linear at a time on ``process_device``."""
|
||||
processed = 0
|
||||
for module in model.modules():
|
||||
if not isinstance(module, LinearBase) or isinstance(module.quant_method, UnquantizedLinearMethod):
|
||||
continue
|
||||
if module.quant_method is None:
|
||||
continue
|
||||
original_device = _module_tensor_device(module)
|
||||
try:
|
||||
module.to(process_device)
|
||||
module.quant_method.process_weights_after_loading(module)
|
||||
finally:
|
||||
if original_device is not None:
|
||||
module.to(original_device)
|
||||
processed += 1
|
||||
if processed == 0:
|
||||
raise ValueError("Serialized quantized text-encoder checkpoint selected, but no quantized linear layers exist")
|
||||
return processed
|
||||
@@ -392,9 +392,16 @@ class MiniMaxH3AudioBigVGANDecoder(nn.Module):
|
||||
return torch.clamp(hidden_states, min=-1.0, max=1.0)
|
||||
|
||||
|
||||
def _is_minimax_h3_audio_vae_decoder(name: str, submodule: nn.Module) -> bool:
|
||||
"""Select the audio decoder that serves the H3 VAE ``decode`` path."""
|
||||
return name == "decoder" and isinstance(submodule, MiniMaxH3AudioBigVGANDecoder)
|
||||
|
||||
|
||||
class MiniMaxH3AudioVAE(nn.Module):
|
||||
"""DAC encoder plus BigVGAN decoder for mono 32 kHz waveforms."""
|
||||
|
||||
_compile_conditions = [_is_minimax_h3_audio_vae_decoder]
|
||||
|
||||
def __init__(self, config: MiniMaxH3AudioVAEConfig):
|
||||
super().__init__()
|
||||
self.config = config
|
||||
|
||||
@@ -0,0 +1,378 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Sequence-parallel chunk scheduling for the MiniMax-H3 video VAE.
|
||||
|
||||
The H3 video VAE decodes a video as a series of temporal-chunk decoder
|
||||
forwards whose outputs are joined by a short deterministic frame blend
|
||||
(``AutoencoderKLMiniMaxH3._decode_chunks``), and encodes videos as fully
|
||||
independent ``clip_length``-frame encoder forwards. Neither the chunk decode
|
||||
nor the clip encode has any cross-chunk data dependency — only the *joining*
|
||||
of decoded chunks (overlap blending, frame trimming) is sequential. This
|
||||
module round-robins the chunk/clip forwards across the ranks of a
|
||||
sequence-parallel group and replays the serial joining logic on the
|
||||
assembling rank, reproducing the serial result bit for bit.
|
||||
|
||||
Bit-exactness contract:
|
||||
- every rank holds an identical copy of the inputs (the H3 DiT all-gathers
|
||||
its outputs, and reference pixels are prepared identically on all ranks);
|
||||
- a chunk decoded on any rank is bitwise the tensor the serial loop would
|
||||
produce (identical weights, inputs, and deterministic kernels on identical
|
||||
GPUs), and NCCL transports it bitwise;
|
||||
- every serialization point of the serial algorithm (overlap blending, frame
|
||||
trimming, pixel denormalization, output-buffer copies, moment
|
||||
concatenation and token-drop trimming) runs on the assembling rank in
|
||||
serial order via the same VAE methods the serial path uses.
|
||||
|
||||
Collective safety: all group ranks must call these functions together with
|
||||
identically shaped inputs. Work proceeds in rounds of one collective each;
|
||||
ranks without a chunk in the final round contribute a placeholder tensor, so
|
||||
participation is uniform by construction and no rank-dependent branch guards
|
||||
a collective.
|
||||
|
||||
Caveat — compiled decoders (``enable_torch_compile_vae``): inductor autotunes
|
||||
kernel configs per process at first call, so a compiled decoder is only
|
||||
deterministic WITHIN a process, not across processes. Chunks decoded on other
|
||||
ranks then differ from the serial rank's decode of the same chunk exactly as
|
||||
two serial runs in different processes would (measured on GB200 at 124f:
|
||||
max 63/255 on <0.5% of pixels, mean ~1e-2/255, first chunk bit-identical).
|
||||
With the eager decoder — the pipeline default — parallel output is bitwise
|
||||
equal to serial ``decode_to_pixels``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.models.vaes.minimax_h3_video import (
|
||||
AutoencoderKLMiniMaxH3,
|
||||
AutoencoderKLOutput,
|
||||
DiagonalGaussianDistribution,
|
||||
)
|
||||
from fastvideo.profiler import nvtx_range
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastvideo.distributed.parallel_state import GroupCoordinator
|
||||
|
||||
# Collective used to move decoded chunk segments to the assembling rank.
|
||||
# "gather" moves each segment once (destination-only); "all_gather" also
|
||||
# leaves every rank with every segment. Both are exact; the default is the
|
||||
# faster one measured on GB200 NVL72 (see the PR notes).
|
||||
DECODE_GATHER_STRATEGIES = ("gather", "all_gather")
|
||||
DEFAULT_DECODE_GATHER_STRATEGY = "gather"
|
||||
|
||||
|
||||
def parallel_chunk_indices(num_chunks: int, world_size: int, rank_in_group: int) -> list[int]:
|
||||
"""Round-robin chunk ownership: chunk ``i`` belongs to rank ``i % world_size``."""
|
||||
if num_chunks < 0:
|
||||
raise ValueError(f"num_chunks must be non-negative, got {num_chunks}.")
|
||||
if world_size < 1:
|
||||
raise ValueError(f"world_size must be positive, got {world_size}.")
|
||||
if not 0 <= rank_in_group < world_size:
|
||||
raise ValueError(f"rank_in_group {rank_in_group} out of range for world_size {world_size}.")
|
||||
return list(range(rank_in_group, num_chunks, world_size))
|
||||
|
||||
|
||||
def _num_rounds(num_chunks: int, world_size: int) -> int:
|
||||
return -(-num_chunks // world_size)
|
||||
|
||||
|
||||
def _decode_segment(vae: AutoencoderKLMiniMaxH3, z_padded: torch.Tensor, chunk_index: int) -> torch.Tensor:
|
||||
"""Decode one temporal chunk's clip and keep the frames the join consumes.
|
||||
|
||||
The serial loop uses two spans of each decoded clip: the chunk body
|
||||
``clip[:, :, frame_pre_padding:chunk_num_frames]`` and (when
|
||||
``token_drop > 0``) the blend tail
|
||||
``clip[:, :, chunk_num_frames + frame_pre_padding:]``. Everything from
|
||||
``frame_pre_padding`` on covers both, so one contiguous slice per chunk
|
||||
travels over the wire. ``.contiguous()`` also detaches the segment from
|
||||
any decoder-owned storage (e.g. a compiled decoder's reuse pools) before
|
||||
the next chunk decode can overwrite it.
|
||||
"""
|
||||
start = chunk_index * vae.tokens_chunk_size
|
||||
with nvtx_range(f"minimax_h3.vae.parallel_chunk.{chunk_index}"):
|
||||
clip = vae._decode_clip(z_padded[:, :, start:start + vae.tokens_chunk_size + vae.token_overlap])
|
||||
return clip[:, :, vae.frame_pre_padding:].contiguous()
|
||||
|
||||
|
||||
class _ChunkAssembler:
|
||||
"""Replay the serial chunk-joining semantics of ``_decode_chunks`` +
|
||||
``_decode_to_pixels`` on gathered chunk segments, in chunk order.
|
||||
|
||||
On CUDA the joining kernels and output copies run on a dedicated side
|
||||
stream: they depend only on already-gathered segments, so running them
|
||||
off the main stream keeps the assembling rank's next chunk decode (and
|
||||
therefore every other rank's next collective) off the assembly's tail.
|
||||
Stream placement cannot change values — the ops and their order are
|
||||
identical — so bit-exactness with the serial path is unaffected.
|
||||
"""
|
||||
|
||||
def __init__(self, vae: AutoencoderKLMiniMaxH3, output: torch.Tensor, output_num_frames: int,
|
||||
non_blocking: bool, device: torch.device) -> None:
|
||||
self._vae = vae
|
||||
self._output = output
|
||||
self._output_num_frames = output_num_frames
|
||||
self._non_blocking = non_blocking
|
||||
self._body_frames = vae.tokens_chunk_size * vae.temporal_compression_ratio - vae.frame_pre_padding
|
||||
self._overlap: torch.Tensor | None = None
|
||||
self._frame_start = 0
|
||||
self._stream = torch.cuda.Stream(device) if device.type == "cuda" else None
|
||||
|
||||
def push(self, segment: torch.Tensor) -> None:
|
||||
"""Consume the next chunk's segment (``clip[:, :, frame_pre_padding:]``)."""
|
||||
if self._stream is None:
|
||||
self._push(segment)
|
||||
return
|
||||
# The segment is produced on the current (collective) stream; hand it
|
||||
# to the assembly stream and pin its storage until assembly reads it.
|
||||
self._stream.wait_stream(torch.cuda.current_stream(segment.device))
|
||||
segment.record_stream(self._stream)
|
||||
with torch.cuda.stream(self._stream):
|
||||
self._push(segment)
|
||||
|
||||
def _push(self, segment: torch.Tensor) -> None:
|
||||
vae = self._vae
|
||||
chunk = segment[:, :, :self._body_frames]
|
||||
if self._overlap is not None:
|
||||
chunk = vae._blend(self._overlap, chunk, vae.frame_overlap, dim=-3)
|
||||
num_frames = min(chunk.shape[2], self._output_num_frames - self._frame_start)
|
||||
chunk = chunk[:, :, :num_frames]
|
||||
# The tail past the body (and its pre-padding gap) is the next
|
||||
# chunk's blend overlap — the serial loop's ``next_overlap``.
|
||||
self._overlap = segment[:, :, self._body_frames + vae.frame_pre_padding:] if vae.config.token_drop > 0 else None
|
||||
if num_frames > 0:
|
||||
self._emit(chunk)
|
||||
|
||||
def finalize(self) -> None:
|
||||
"""Emit the final overlap tail exactly as the serial generator does."""
|
||||
if self._overlap is not None and self._frame_start < self._output_num_frames:
|
||||
tail = self._overlap[:, :, :self._output_num_frames - self._frame_start]
|
||||
if self._stream is None:
|
||||
self._emit(tail)
|
||||
else:
|
||||
with torch.cuda.stream(self._stream):
|
||||
self._emit(tail)
|
||||
if self._frame_start != self._output.shape[2]:
|
||||
raise RuntimeError(
|
||||
f"MiniMax-H3 decode wrote {self._frame_start} frames into an output buffer expecting "
|
||||
f"{self._output.shape[2]}.")
|
||||
|
||||
def synchronize(self) -> None:
|
||||
"""Drain assembly kernels and output copies before the buffer is read."""
|
||||
if self._stream is not None:
|
||||
self._stream.synchronize()
|
||||
|
||||
def _emit(self, chunk: torch.Tensor) -> None:
|
||||
pixels = self._vae.denormalize_pixels(chunk.float()).clamp_(0, 1)
|
||||
self._vae._copy_chunk_pixels(pixels, self._output, self._frame_start, self._non_blocking)
|
||||
self._frame_start += pixels.shape[2]
|
||||
|
||||
|
||||
def _broadcast_segment_meta(group: "GroupCoordinator",
|
||||
segment: torch.Tensor | None) -> tuple[torch.dtype, tuple[int, ...]]:
|
||||
"""Share the leader's real segment dtype/shape so placeholder tensors match.
|
||||
|
||||
The decoder's output dtype depends on the surrounding autocast context;
|
||||
deriving it on the leader from an actually decoded segment (instead of
|
||||
predicting it) keeps collective dtypes correct by construction.
|
||||
"""
|
||||
meta = (segment.dtype, tuple(segment.shape)) if segment is not None else None
|
||||
meta = group.broadcast_object(meta, src=0)
|
||||
if meta is None:
|
||||
raise RuntimeError("MiniMax-H3 parallel VAE meta broadcast returned no leader metadata.")
|
||||
return meta
|
||||
|
||||
|
||||
def decode_to_pixels_parallel(
|
||||
vae: AutoencoderKLMiniMaxH3,
|
||||
z: torch.Tensor,
|
||||
output: torch.Tensor | None,
|
||||
group: "GroupCoordinator",
|
||||
strategy: str = DEFAULT_DECODE_GATHER_STRATEGY,
|
||||
) -> torch.Tensor | None:
|
||||
"""Chunk-parallel ``decode_to_pixels`` across a sequence-parallel group.
|
||||
|
||||
All group ranks call this together with identical ``z``. Temporal chunks
|
||||
are decoded round-robin across the group and their segments move to the
|
||||
group's first rank, which assembles bitwise the serial
|
||||
``decode_to_pixels`` result into ``output``. Only the first rank passes
|
||||
``output`` (validated exactly like the serial API); other ranks pass
|
||||
``None`` and receive ``None``.
|
||||
"""
|
||||
if strategy not in DECODE_GATHER_STRATEGIES:
|
||||
raise ValueError(f"Unknown parallel-decode strategy {strategy!r}; expected one of {DECODE_GATHER_STRATEGIES}.")
|
||||
is_leader = group.rank_in_group == 0
|
||||
if is_leader:
|
||||
if output is None:
|
||||
raise ValueError("The first sequence-parallel rank must provide the CPU output buffer.")
|
||||
expected_shape = vae.decoded_pixel_shape(z.shape)
|
||||
if output.device.type != "cpu" or output.dtype != torch.float32 or tuple(output.shape) != expected_shape:
|
||||
raise ValueError(
|
||||
"`output` must be a CPU float32 tensor with shape "
|
||||
f"{expected_shape}, got device={output.device}, dtype={output.dtype}, shape={tuple(output.shape)}.")
|
||||
elif output is not None:
|
||||
raise ValueError("Only the first sequence-parallel rank may provide an output buffer.")
|
||||
if group.world_size == 1:
|
||||
return vae.decode_to_pixels(z, output)
|
||||
|
||||
try:
|
||||
if vae.use_slicing and z.shape[0] > 1:
|
||||
for batch_index, z_slice in enumerate(z.split(1)):
|
||||
slice_output = output[batch_index:batch_index + 1] if output is not None else None
|
||||
_decode_single_parallel(vae, z_slice, slice_output, group, strategy)
|
||||
else:
|
||||
_decode_single_parallel(vae, z, output, group, strategy)
|
||||
finally:
|
||||
# Drain the leader's async chunk copies before the caller (or an
|
||||
# exception handler) can read or release the pinned buffer.
|
||||
if output is not None and vae._streams_chunk_copies(z, output):
|
||||
torch.cuda.current_stream(z.device).synchronize()
|
||||
return output
|
||||
|
||||
|
||||
def _decode_single_parallel(
|
||||
vae: AutoencoderKLMiniMaxH3,
|
||||
z: torch.Tensor,
|
||||
output: torch.Tensor | None,
|
||||
group: "GroupCoordinator",
|
||||
strategy: str,
|
||||
) -> None:
|
||||
pad_tokens, num_chunks, output_num_frames = vae._temporal_decode_plan(z.shape[2])
|
||||
if pad_tokens > 0:
|
||||
z = torch.cat([z, z[:, :, -1:].repeat(1, 1, pad_tokens, 1, 1)], dim=2)
|
||||
world_size = group.world_size
|
||||
rank = group.rank_in_group
|
||||
|
||||
# Every rank decodes its round-0 chunk BEFORE the metadata rendezvous so
|
||||
# the first decodes run concurrently (a rank that waited on the broadcast
|
||||
# first would idle a full chunk-decode behind the leader). The leader
|
||||
# owns chunk 0 under round-robin assignment, so its segment supplies real
|
||||
# dtype/shape for placeholder rounds instead of guessing autocast state.
|
||||
first_segment = _decode_segment(vae, z, rank) if rank < num_chunks else None
|
||||
segment_dtype, segment_shape = _broadcast_segment_meta(group, first_segment if rank == 0 else None)
|
||||
|
||||
assembler = None
|
||||
if output is not None:
|
||||
non_blocking = vae._streams_chunk_copies(z, output)
|
||||
assembler = _ChunkAssembler(vae, output, output_num_frames, non_blocking, z.device)
|
||||
|
||||
try:
|
||||
segment_frames = segment_shape[2]
|
||||
for round_index in range(_num_rounds(num_chunks, world_size)):
|
||||
chunk_index = round_index * world_size + rank
|
||||
if chunk_index >= num_chunks:
|
||||
segment = torch.zeros(segment_shape, dtype=segment_dtype, device=z.device)
|
||||
elif round_index == 0 and first_segment is not None:
|
||||
segment = first_segment
|
||||
else:
|
||||
segment = _decode_segment(vae, z, chunk_index)
|
||||
with nvtx_range(f"minimax_h3.vae.parallel_{strategy}.{round_index}"):
|
||||
if strategy == "gather":
|
||||
gathered = group.gather(segment, dst=0, dim=2)
|
||||
else:
|
||||
gathered = group.all_gather(segment, dim=2)
|
||||
if assembler is None or gathered is None:
|
||||
continue
|
||||
for slot in range(world_size):
|
||||
if round_index * world_size + slot >= num_chunks:
|
||||
break
|
||||
assembler.push(gathered.narrow(2, slot * segment_frames, segment_frames))
|
||||
if assembler is not None:
|
||||
assembler.finalize()
|
||||
finally:
|
||||
# Drain assembly-stream copies into ``output`` even on the error path
|
||||
# so an exception cannot leave an in-flight DMA into a buffer the
|
||||
# caller may release.
|
||||
if assembler is not None:
|
||||
assembler.synchronize()
|
||||
|
||||
|
||||
def _encode_clip_moments(vae: AutoencoderKLMiniMaxH3, pixels: torch.Tensor, clip_index: int) -> torch.Tensor:
|
||||
"""Encode one ``clip_length``-frame clip exactly as ``_encode_pixels`` does."""
|
||||
clip_length = vae.config.clip_length
|
||||
frame_start = clip_index * clip_length
|
||||
with nvtx_range(f"minimax_h3.vae.parallel_encode_clip.{clip_index}"):
|
||||
clip = pixels[:, :, frame_start:frame_start + clip_length].to(
|
||||
device=vae.pixel_mean.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
if pixels.dtype == torch.uint8:
|
||||
clip = clip / 255.0
|
||||
if clip.shape[2] < clip_length:
|
||||
pad_frames = clip[:, :, -1:].repeat(1, 1, clip_length - clip.shape[2], 1, 1)
|
||||
clip = torch.cat([clip, pad_frames], dim=2)
|
||||
clip = vae.normalize_pixels(clip)
|
||||
return vae._encode_clip(clip).contiguous()
|
||||
|
||||
|
||||
def encode_pixels_parallel(
|
||||
vae: AutoencoderKLMiniMaxH3,
|
||||
pixels: torch.Tensor,
|
||||
group: "GroupCoordinator",
|
||||
) -> AutoencoderKLOutput:
|
||||
"""Clip-parallel ``encode_pixels`` across a sequence-parallel group.
|
||||
|
||||
Encoder clips have no cross-clip dependency (no overlap, no blending), so
|
||||
ranks encode disjoint clips and all-gather the per-clip moment tensors.
|
||||
Every rank returns the identical full posterior — preserving the serial
|
||||
contract that all ranks hold the same encoded latents — bitwise equal to
|
||||
``vae.encode_pixels(pixels)``. Moments are latent-sized (a few MB per
|
||||
clip), so the all-gather is negligible next to the clip forwards.
|
||||
"""
|
||||
if pixels.ndim != 5 or pixels.shape[1] != vae.config.in_channels or pixels.shape[2] <= 0:
|
||||
raise ValueError(
|
||||
f"`pixels` must have shape [B, {vae.config.in_channels}, T, H, W] with T > 0, "
|
||||
f"got {tuple(pixels.shape)}.")
|
||||
if pixels.device.type != "cpu":
|
||||
raise ValueError(f"`pixels` must remain on CPU, got device={pixels.device}.")
|
||||
if pixels.dtype != torch.uint8 and not pixels.is_floating_point():
|
||||
raise TypeError(f"`pixels` must use uint8 or a floating-point dtype, got {pixels.dtype}.")
|
||||
if group.world_size == 1:
|
||||
return vae.encode_pixels(pixels)
|
||||
if vae.use_slicing and pixels.shape[0] > 1:
|
||||
moments = torch.cat([_encode_single_parallel(vae, pixel_slice, group) for pixel_slice in pixels.split(1)])
|
||||
else:
|
||||
moments = _encode_single_parallel(vae, pixels, group)
|
||||
return AutoencoderKLOutput(latent_dist=DiagonalGaussianDistribution(moments))
|
||||
|
||||
|
||||
def _encode_single_parallel(vae: AutoencoderKLMiniMaxH3, pixels: torch.Tensor,
|
||||
group: "GroupCoordinator") -> torch.Tensor:
|
||||
clip_length = vae.config.clip_length
|
||||
num_clips = -(-pixels.shape[2] // clip_length)
|
||||
world_size = group.world_size
|
||||
rank = group.rank_in_group
|
||||
|
||||
# Same first-work-then-rendezvous ordering as the decode path: encode the
|
||||
# round-0 clip before the metadata broadcast so first encodes overlap.
|
||||
first_moments = _encode_clip_moments(vae, pixels, rank) if rank < num_clips else None
|
||||
moment_dtype, moment_shape = _broadcast_segment_meta(group, first_moments if rank == 0 else None)
|
||||
|
||||
moment_tokens = moment_shape[2]
|
||||
parts: list[torch.Tensor] = []
|
||||
for round_index in range(_num_rounds(num_clips, world_size)):
|
||||
clip_index = round_index * world_size + rank
|
||||
if clip_index >= num_clips:
|
||||
moments = torch.zeros(moment_shape, dtype=moment_dtype, device=vae.pixel_mean.device)
|
||||
elif round_index == 0 and first_moments is not None:
|
||||
moments = first_moments
|
||||
else:
|
||||
moments = _encode_clip_moments(vae, pixels, clip_index)
|
||||
gathered = group.all_gather(moments, dim=2)
|
||||
for slot in range(world_size):
|
||||
if round_index * world_size + slot >= num_clips:
|
||||
break
|
||||
parts.append(gathered.narrow(2, slot * moment_tokens, moment_tokens))
|
||||
encoded = torch.cat(parts, dim=2)
|
||||
if vae.config.token_drop > 0:
|
||||
encoded = encoded[:, :, :-vae.config.token_drop]
|
||||
return encoded
|
||||
|
||||
|
||||
__all__ = [
|
||||
"DECODE_GATHER_STRATEGIES",
|
||||
"DEFAULT_DECODE_GATHER_STRATEGY",
|
||||
"decode_to_pixels_parallel",
|
||||
"encode_pixels_parallel",
|
||||
"parallel_chunk_indices",
|
||||
]
|
||||
@@ -15,7 +15,10 @@ import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
|
||||
from fastvideo.attention import get_attn_backend
|
||||
from fastvideo.configs.models.vaes.minimax_h3_video import MiniMaxH3VideoVAEConfig
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.profiler import nvtx_range
|
||||
|
||||
|
||||
class DiagonalGaussianDistribution:
|
||||
@@ -292,6 +295,7 @@ class MiniMaxH3VideoRotaryPosEmbed(nn.Module):
|
||||
class MiniMaxH3VideoAttention(nn.Module):
|
||||
|
||||
def __init__(self, dim: int, heads: int, dim_head: int, eps: float = 1e-5, bias: bool = True) -> None:
|
||||
"""Build projections and the selected dense FastVideo attention implementation."""
|
||||
super().__init__()
|
||||
self.heads = heads
|
||||
self.dim_head = dim_head
|
||||
@@ -303,12 +307,34 @@ class MiniMaxH3VideoAttention(nn.Module):
|
||||
self.to_k = nn.Linear(dim, inner_dim, bias=bias)
|
||||
self.to_v = nn.Linear(dim, inner_dim, bias=bias)
|
||||
self.to_out = nn.ModuleList([nn.Linear(inner_dim, dim, bias=bias), nn.Dropout(0.0)])
|
||||
self.attn_impl = None
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
if current_platform.is_cuda_alike():
|
||||
attention_backend = get_attn_backend(
|
||||
dim_head,
|
||||
# FlashAttention executes the FP32 VAE activations in BF16 and
|
||||
# restores FP32 output, so resolve against the kernel dtype.
|
||||
torch.bfloat16,
|
||||
supported_attention_backends=(
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
),
|
||||
)
|
||||
self.attn_impl = attention_backend.get_impl_cls()(
|
||||
num_heads=heads,
|
||||
head_size=dim_head,
|
||||
softmax_scale=dim_head**-0.5,
|
||||
num_kv_heads=heads,
|
||||
causal=False,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
rotary_emb: tuple[torch.Tensor, torch.Tensor] | None = None,
|
||||
) -> torch.Tensor:
|
||||
"""Apply dense self-attention to one spatial VAE token sequence."""
|
||||
query = self.to_q(hidden_states).unflatten(2, (self.heads, -1))
|
||||
key = self.to_k(hidden_states).unflatten(2, (self.heads, -1))
|
||||
value = self.to_v(hidden_states).unflatten(2, (self.heads, -1))
|
||||
@@ -329,9 +355,17 @@ class MiniMaxH3VideoAttention(nn.Module):
|
||||
query = torch.cat([query_rotary * cos + query_rotated * sin, query_pass], dim=-1)
|
||||
key = torch.cat([key_rotary * cos + key_rotated * sin, key_pass], dim=-1)
|
||||
|
||||
query, key, value = (tensor.permute(0, 2, 1, 3) for tensor in (query, key, value))
|
||||
hidden_states = F.scaled_dot_product_attention(query, key, value)
|
||||
hidden_states = hidden_states.permute(0, 2, 1, 3).flatten(2, 3)
|
||||
if self.attn_impl is not None and query.device.type != "cpu":
|
||||
# VAE decoding has no diffusion-step metadata, so call the selected
|
||||
# backend implementation directly with dense BSHD tensors.
|
||||
hidden_states = self.attn_impl.forward(query, key, value, None)
|
||||
hidden_states = hidden_states.flatten(2, 3)
|
||||
else:
|
||||
# Keep CPU construction and execution available without requiring
|
||||
# an accelerator attention backend.
|
||||
query, key, value = (tensor.permute(0, 2, 1, 3) for tensor in (query, key, value))
|
||||
hidden_states = F.scaled_dot_product_attention(query, key, value)
|
||||
hidden_states = hidden_states.permute(0, 2, 1, 3).flatten(2, 3)
|
||||
return self.to_out[0](hidden_states)
|
||||
|
||||
|
||||
@@ -434,6 +468,7 @@ class MiniMaxH3VideoViTDecoder3d(nn.Module):
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
"""Decode one latent spatial input through the H3 video transformer."""
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
hidden_states = hidden_states.permute(0, 2, 3, 4, 1).reshape(
|
||||
batch_size,
|
||||
@@ -483,6 +518,11 @@ class MiniMaxH3VideoViTDecoder3d(nn.Module):
|
||||
)
|
||||
|
||||
|
||||
def _is_minimax_h3_video_vae_decoder(name: str, submodule: nn.Module) -> bool:
|
||||
"""Select the video decoder that serves the H3 VAE ``decode`` path."""
|
||||
return name == "decoder" and isinstance(submodule, MiniMaxH3VideoViTDecoder3d)
|
||||
|
||||
|
||||
class AutoencoderKLMiniMaxH3(nn.Module):
|
||||
"""MiniMax-H3 causal encoder and ViT decoder with exact release geometry."""
|
||||
|
||||
@@ -490,6 +530,7 @@ class AutoencoderKLMiniMaxH3(nn.Module):
|
||||
_no_split_modules = ["MiniMaxH3VideoResnetBlock3d", "MiniMaxH3VideoTransformerBlock"]
|
||||
_repeated_blocks = ["MiniMaxH3VideoTransformerBlock"]
|
||||
_keep_in_fp32_modules = ["encoder", "decoder", "quant_conv", "post_quant_conv"]
|
||||
_compile_conditions = [_is_minimax_h3_video_vae_decoder]
|
||||
|
||||
def __init__(self, config: MiniMaxH3VideoVAEConfig) -> None:
|
||||
super().__init__()
|
||||
@@ -655,12 +696,15 @@ class AutoencoderKLMiniMaxH3(nn.Module):
|
||||
slice_rest[dim] = slice(blend_extent, None)
|
||||
return torch.cat([blended, b[tuple(slice_rest)]], dim=dim)
|
||||
|
||||
# The fixed spatial tile grid reuses one compiled blend-and-concatenate graph.
|
||||
@torch.compile(backend="inductor", mode="reduce-overhead", dynamic=False)
|
||||
def _stitch_tiles(
|
||||
self,
|
||||
tiles: list[list[torch.Tensor]],
|
||||
height_overlaps: list[int],
|
||||
width_overlaps: list[int],
|
||||
) -> torch.Tensor:
|
||||
"""Blend decoded tile overlaps and concatenate the spatial canvas."""
|
||||
result_rows = []
|
||||
for row_index, row in enumerate(tiles):
|
||||
result_row = []
|
||||
@@ -677,6 +721,12 @@ class AutoencoderKLMiniMaxH3(nn.Module):
|
||||
result_rows.append(torch.cat(result_row, dim=-1))
|
||||
return torch.cat(result_rows, dim=-2)
|
||||
|
||||
# Each fixed-shape latent tile reuses one compiled decoder-input projection.
|
||||
@torch.compile(backend="inductor", mode="reduce-overhead", dynamic=False)
|
||||
def _project_decoder_tile(self, tile: torch.Tensor) -> torch.Tensor:
|
||||
"""Project one spatial latent tile into the decoder input channels."""
|
||||
return self.post_quant_conv(tile)
|
||||
|
||||
def _encode_clip(self, x: torch.Tensor) -> torch.Tensor:
|
||||
if not self.use_tiling:
|
||||
return self.quant_conv(self.encoder(x))
|
||||
@@ -700,36 +750,64 @@ class AutoencoderKLMiniMaxH3(nn.Module):
|
||||
rows.append(row)
|
||||
latent_y_overlaps = [overlap // self.spatial_compression_ratio for overlap in y_overlaps]
|
||||
latent_x_overlaps = [overlap // self.spatial_compression_ratio for overlap in x_overlaps]
|
||||
return self._stitch_tiles(rows, latent_y_overlaps, latent_x_overlaps)
|
||||
# Under mode="reduce-overhead" the stitched canvas is a CUDA-graph
|
||||
# static buffer that the next _stitch_tiles replay overwrites. Callers
|
||||
# (_encode/_encode_pixels/encode_keyframe) collect per-clip results
|
||||
# across replays before concatenating, so hand them a caller-owned
|
||||
# tensor instead of cudagraph-pooled storage.
|
||||
return self._stitch_tiles(rows, latent_y_overlaps, latent_x_overlaps).clone()
|
||||
|
||||
def _decode_clip(self, z: torch.Tensor) -> torch.Tensor:
|
||||
if not self.use_tiling:
|
||||
return self.decoder(self.post_quant_conv(z))
|
||||
height = z.shape[-2] * self.spatial_compression_ratio
|
||||
width = z.shape[-1] * self.spatial_compression_ratio
|
||||
y_indices, y_lengths, y_overlaps = self._split_tiles(
|
||||
height,
|
||||
self.tile_sample_min_height,
|
||||
self.tile_sample_min_overlap_height,
|
||||
)
|
||||
x_indices, x_lengths, x_overlaps = self._split_tiles(
|
||||
width,
|
||||
self.tile_sample_min_width,
|
||||
self.tile_sample_min_overlap_width,
|
||||
)
|
||||
ratio = self.spatial_compression_ratio
|
||||
rows = []
|
||||
for y_position, y_length in zip(y_indices, y_lengths):
|
||||
row = []
|
||||
for x_position, x_length in zip(x_indices, x_lengths):
|
||||
tile = z[
|
||||
...,
|
||||
y_position // ratio:y_position // ratio + y_length // ratio,
|
||||
x_position // ratio:x_position // ratio + x_length // ratio,
|
||||
]
|
||||
row.append(self.decoder(self.post_quant_conv(tile)))
|
||||
rows.append(row)
|
||||
return self._stitch_tiles(rows, y_overlaps, x_overlaps)
|
||||
"""Decode one temporal clip, with optional overlapping spatial tiles."""
|
||||
with nvtx_range("minimax_h3.vae.decode_clip"):
|
||||
if not self.use_tiling:
|
||||
with nvtx_range("minimax_h3.vae.decode_clip.no_s_tile.post_quant_conv"):
|
||||
projected_clip = self.post_quant_conv(z)
|
||||
with nvtx_range("minimax_h3.vae.decode_clip.no_s_tile.decoder_forward"):
|
||||
return self.decoder(projected_clip)
|
||||
|
||||
height = z.shape[-2] * self.spatial_compression_ratio
|
||||
width = z.shape[-1] * self.spatial_compression_ratio
|
||||
with nvtx_range("minimax_h3.vae.decode_clip.split_tiles"):
|
||||
y_indices, y_lengths, y_overlaps = self._split_tiles(
|
||||
height,
|
||||
self.tile_sample_min_height,
|
||||
self.tile_sample_min_overlap_height,
|
||||
)
|
||||
x_indices, x_lengths, x_overlaps = self._split_tiles(
|
||||
width,
|
||||
self.tile_sample_min_width,
|
||||
self.tile_sample_min_overlap_width,
|
||||
)
|
||||
|
||||
ratio = self.spatial_compression_ratio
|
||||
rows = []
|
||||
# The eager tile driver owns NVTX so each marker remains outside
|
||||
# the compiled decoder graph.
|
||||
with nvtx_range("minimax_h3.vae.decode_clip.decode_tiles"):
|
||||
for row_index, (y_position, y_length) in enumerate(zip(y_indices, y_lengths)):
|
||||
row = []
|
||||
for column_index, (x_position, x_length) in enumerate(zip(x_indices, x_lengths)):
|
||||
with nvtx_range(f"minimax_h3.vae.decode_clip.tile.{row_index}.{column_index}"):
|
||||
tile = z[
|
||||
...,
|
||||
y_position // ratio:y_position // ratio + y_length // ratio,
|
||||
x_position // ratio:x_position // ratio + x_length // ratio,
|
||||
]
|
||||
projected_tile = self._project_decoder_tile(tile)
|
||||
with nvtx_range("minimax_h3.vae.decode_clip.tile.decoder_forward"):
|
||||
decoded_tile = self.decoder(projected_tile)
|
||||
row.append(decoded_tile)
|
||||
rows.append(row)
|
||||
|
||||
with nvtx_range("minimax_h3.vae.decode_clip.stitch_tiles"):
|
||||
# Same CUDA-graph output-ownership contract as _encode_clip:
|
||||
# _decode collects chunks across _stitch_tiles replays before
|
||||
# torch.cat, so the pooled canvas must not escape this driver.
|
||||
# (The streaming _decode_to_pixels path copies each chunk out
|
||||
# before the next decode and never held stale storage; the
|
||||
# clone keeps that path correct too at one D2D copy per chunk.)
|
||||
return self._stitch_tiles(rows, y_overlaps, x_overlaps).clone()
|
||||
|
||||
def _encode(self, x: torch.Tensor) -> torch.Tensor:
|
||||
clip_length = self.config.clip_length
|
||||
@@ -809,21 +887,27 @@ class AutoencoderKLMiniMaxH3(nn.Module):
|
||||
|
||||
output_frame_start = 0
|
||||
overlap = None
|
||||
for index in range(num_chunks):
|
||||
start = index * tokens_chunk_size
|
||||
clip = self._decode_clip(z[:, :, start:start + tokens_chunk_size + self.token_overlap])
|
||||
chunk = clip[:, :, self.frame_pre_padding:chunk_num_frames]
|
||||
next_overlap = None
|
||||
if self.config.token_drop > 0:
|
||||
next_overlap = clip[:, :, chunk_num_frames + self.frame_pre_padding:].clone()
|
||||
if overlap is not None:
|
||||
chunk = self._blend(overlap, chunk, self.frame_overlap, dim=-3)
|
||||
for chunk_index in range(num_chunks):
|
||||
with nvtx_range(f"minimax_h3.vae.temporal_chunk.{chunk_index}"):
|
||||
start = chunk_index * tokens_chunk_size
|
||||
clip = self._decode_clip(z[:, :, start:start + tokens_chunk_size + self.token_overlap])
|
||||
with nvtx_range(f"minimax_h3.vae.temporal_chunk.{chunk_index}.frame_segment.0"):
|
||||
chunk = clip[:, :, self.frame_pre_padding:chunk_num_frames]
|
||||
if overlap is not None:
|
||||
chunk = self._blend(overlap, chunk, self.frame_overlap, dim=-3)
|
||||
num_frames = min(chunk.shape[2], output_num_frames - output_frame_start)
|
||||
chunk = chunk[:, :, :num_frames]
|
||||
|
||||
num_frames = min(chunk.shape[2], output_num_frames - output_frame_start)
|
||||
if num_frames > 0:
|
||||
yield chunk[:, :, :num_frames]
|
||||
output_frame_start += num_frames
|
||||
next_overlap = None
|
||||
if self.config.token_drop > 0:
|
||||
with nvtx_range(f"minimax_h3.vae.temporal_chunk.{chunk_index}.frame_segment.1"):
|
||||
next_overlap = clip[:, :, chunk_num_frames + self.frame_pre_padding:].clone()
|
||||
|
||||
# Yield after the ranges close so consumer-side CPU copies do not inflate decoder timing.
|
||||
overlap = next_overlap
|
||||
if num_frames > 0:
|
||||
output_frame_start += num_frames
|
||||
yield chunk
|
||||
|
||||
if overlap is not None and output_frame_start < output_num_frames:
|
||||
yield overlap[:, :, :output_num_frames - output_frame_start]
|
||||
@@ -849,31 +933,42 @@ class AutoencoderKLMiniMaxH3(nn.Module):
|
||||
"""Whether finalized chunks copy to ``output`` asynchronously on the current CUDA stream."""
|
||||
return z.device.type == "cuda" and output.is_pinned()
|
||||
|
||||
def _decode_to_pixels(self, z: torch.Tensor, output: torch.Tensor) -> None:
|
||||
"""Decode temporal chunks and immediately copy finalized pixels to CPU.
|
||||
@staticmethod
|
||||
def _copy_chunk_pixels(pixels: torch.Tensor, output: torch.Tensor, frame_start: int, non_blocking: bool) -> None:
|
||||
"""Copy one finalized fp32 pixel chunk into the CPU ``output`` buffer.
|
||||
|
||||
Device-to-host copies run per (batch, channel) plane: the temporal
|
||||
slice of ``output`` is strided across channels, but each plane is
|
||||
contiguous on both sides, so every transfer stays a direct memcpy
|
||||
instead of staging through a pageable CPU temporary. With a pinned
|
||||
``output`` the copies are additionally asynchronous and overlap the
|
||||
next chunk's decode; ``decode_to_pixels`` synchronizes once before
|
||||
returning.
|
||||
``output`` and ``non_blocking=True`` the copies are additionally
|
||||
asynchronous on the current CUDA stream; callers synchronize once
|
||||
before releasing the buffer.
|
||||
"""
|
||||
target = output[:, :, frame_start:frame_start + pixels.shape[2]]
|
||||
if pixels.device.type == "cuda":
|
||||
pixels = pixels.contiguous()
|
||||
for batch_index in range(pixels.shape[0]):
|
||||
for channel_index in range(pixels.shape[1]):
|
||||
target[batch_index, channel_index].copy_(pixels[batch_index, channel_index],
|
||||
non_blocking=non_blocking)
|
||||
else:
|
||||
target.copy_(pixels)
|
||||
|
||||
def _decode_to_pixels(self, z: torch.Tensor, output: torch.Tensor) -> None:
|
||||
"""Decode temporal chunks and immediately copy finalized pixels to CPU.
|
||||
|
||||
Each finalized chunk streams through ``_copy_chunk_pixels`` (direct
|
||||
per-plane memcpys; asynchronous with a pinned ``output``) so the
|
||||
copies overlap the next chunk's decode; ``decode_to_pixels``
|
||||
synchronizes once before returning.
|
||||
"""
|
||||
non_blocking = self._streams_chunk_copies(z, output)
|
||||
output_frame_start = 0
|
||||
for chunk in self._decode_chunks(z):
|
||||
num_frames = chunk.shape[2]
|
||||
pixels = self.denormalize_pixels(chunk.float()).clamp_(0, 1)
|
||||
target = output[:, :, output_frame_start:output_frame_start + num_frames]
|
||||
if z.device.type == "cuda":
|
||||
pixels = pixels.contiguous()
|
||||
for batch_index in range(pixels.shape[0]):
|
||||
for channel_index in range(pixels.shape[1]):
|
||||
target[batch_index, channel_index].copy_(pixels[batch_index, channel_index],
|
||||
non_blocking=non_blocking)
|
||||
else:
|
||||
target.copy_(pixels)
|
||||
self._copy_chunk_pixels(pixels, output, output_frame_start, non_blocking)
|
||||
output_frame_start += num_frames
|
||||
if output_frame_start != output.shape[2]:
|
||||
raise RuntimeError(
|
||||
|
||||
@@ -12,9 +12,9 @@ from torch.distributed.tensor import DTensor
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
|
||||
from fastvideo.profiler import nvtx_range
|
||||
from fastvideo.pipelines.basic.minimax_h3.packing import (
|
||||
MINIMAX_H3_IMAGE_PAD_TOKEN,
|
||||
MINIMAX_H3_TEXT_ENCODER_LAYER,
|
||||
MINIMAX_H3_TEXT_TAG,
|
||||
MINIMAX_H3_VIDEO_PAD_TOKEN,
|
||||
MINIMAX_H3_VIDEO_TAG,
|
||||
@@ -42,25 +42,6 @@ def _token_ids(tokenized: Any) -> list[int]:
|
||||
return [int(token_id) for token_id in input_ids]
|
||||
|
||||
|
||||
def _create_mm_token_type_ids(processor: Any, token_ids: list[int]) -> list[list[int]]:
|
||||
"""Build Qwen3-VL modality IDs across old and new Transformers releases."""
|
||||
create_ids = getattr(processor, "create_mm_token_type_ids", None)
|
||||
if callable(create_ids):
|
||||
return create_ids([token_ids])
|
||||
|
||||
modality_ids = [0] * len(token_ids)
|
||||
for modality, modality_type in (("image", 1), ("video", 2), ("audio", 3)):
|
||||
special_ids = getattr(processor, f"{modality}_token_ids", None)
|
||||
if special_ids is None:
|
||||
special_id = getattr(processor, f"{modality}_token_id", None)
|
||||
special_ids = [] if special_id is None else [special_id]
|
||||
resolved_ids = {int(special_id) for special_id in special_ids if special_id is not None}
|
||||
for index, token_id in enumerate(token_ids):
|
||||
if token_id in resolved_ids:
|
||||
modality_ids[index] = modality_type
|
||||
return [modality_ids]
|
||||
|
||||
|
||||
def build_ref2va_presentation(
|
||||
tokenizer: Any,
|
||||
prompt: str,
|
||||
@@ -155,20 +136,10 @@ class MiniMaxH3ConditioningStage(PipelineStage):
|
||||
device: torch.device,
|
||||
**vision_inputs: torch.Tensor | None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
hidden_state_index = MINIMAX_H3_TEXT_ENCODER_LAYER
|
||||
input_ids = torch.tensor([token_ids], dtype=torch.long, device=device)
|
||||
mm_token_type_ids = torch.as_tensor(
|
||||
_create_mm_token_type_ids(self.processor, token_ids),
|
||||
dtype=torch.long,
|
||||
device=device,
|
||||
)
|
||||
input_ids = torch.tensor(token_ids, dtype=torch.long, device=device)
|
||||
dtype = self.conditioner.dtype
|
||||
outputs = self.conditioner(
|
||||
input_ids=input_ids,
|
||||
attention_mask=torch.ones_like(input_ids),
|
||||
mm_token_type_ids=mm_token_type_ids,
|
||||
use_cache=False,
|
||||
output_hidden_states=True,
|
||||
prompt_embeds = self.conditioner(
|
||||
input_ids,
|
||||
**{
|
||||
name:
|
||||
None if value is None else value.to(
|
||||
@@ -178,10 +149,10 @@ class MiniMaxH3ConditioningStage(PipelineStage):
|
||||
for name, value in vision_inputs.items()
|
||||
},
|
||||
)
|
||||
if outputs.hidden_states is None or len(outputs.hidden_states) <= hidden_state_index:
|
||||
raise ValueError(f"Qwen3-VL did not return `hidden_states[{hidden_state_index}]`.")
|
||||
if prompt_embeds.ndim != 2 or prompt_embeds.shape[0] != len(token_ids):
|
||||
raise ValueError(f"MiniMax-H3 slim text encoder returned unexpected shape={tuple(prompt_embeds.shape)}")
|
||||
return (
|
||||
outputs.hidden_states[hidden_state_index].to(device=device, dtype=dtype),
|
||||
prompt_embeds.unsqueeze(0).to(device=device, dtype=dtype),
|
||||
torch.tensor(token_tags, dtype=torch.long),
|
||||
)
|
||||
|
||||
@@ -286,6 +257,7 @@ class MiniMaxH3ConditioningStage(PipelineStage):
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
"""Encode one H3 prompt presentation and attach its packed text features."""
|
||||
device = get_local_torch_device()
|
||||
first_param = next(self.conditioner.parameters(), None)
|
||||
moved_for_forward = (fastvideo_args.text_encoder_cpu_offload and first_param is not None
|
||||
@@ -293,10 +265,13 @@ class MiniMaxH3ConditioningStage(PipelineStage):
|
||||
if moved_for_forward:
|
||||
self.conditioner.to(device)
|
||||
try:
|
||||
if self.ref2va:
|
||||
prompt_embeds, text_token_tags = self._encode_ref2va(batch, device)
|
||||
else:
|
||||
prompt_embeds, text_token_tags = self._encode_fl2va(batch, device)
|
||||
# Keep both H3 prompt-presentation modes under one text-encoding
|
||||
# range so Nsight Systems exposes their complete conditioning cost.
|
||||
with nvtx_range("minimax_h3.text_encoding"):
|
||||
if self.ref2va:
|
||||
prompt_embeds, text_token_tags = self._encode_ref2va(batch, device)
|
||||
else:
|
||||
prompt_embeds, text_token_tags = self._encode_fl2va(batch, device)
|
||||
finally:
|
||||
if moved_for_forward:
|
||||
self.conditioner.to("cpu")
|
||||
|
||||
@@ -7,10 +7,13 @@ from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device, get_world_group, model_parallel_is_initialized
|
||||
from fastvideo.distributed import get_local_torch_device, get_sp_group, model_parallel_is_initialized
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.vaes.minimax_h3_audio import MiniMaxH3AudioVAE
|
||||
from fastvideo.models.vaes.minimax_h3_parallel import DEFAULT_DECODE_GATHER_STRATEGY, decode_to_pixels_parallel
|
||||
from fastvideo.models.vaes.minimax_h3_video import AutoencoderKLMiniMaxH3
|
||||
from fastvideo.profiler import nvtx_range
|
||||
from fastvideo.pipelines.basic.minimax_h3.packing import (
|
||||
MiniMaxH3PackedLayout,
|
||||
unpack_audio_tokens,
|
||||
@@ -23,6 +26,8 @@ from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.utils import is_pin_memory_available
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def _layout(batch: ForwardBatch) -> MiniMaxH3PackedLayout:
|
||||
layout = batch.extra.get(MINIMAX_H3_LAYOUT_KEY)
|
||||
@@ -31,6 +36,23 @@ def _layout(batch: ForwardBatch) -> MiniMaxH3PackedLayout:
|
||||
return layout
|
||||
|
||||
|
||||
def _decode_participation(fastvideo_args: FastVideoArgs, want_parallel: bool) -> tuple[Any, bool, bool]:
|
||||
"""Resolve (sp_group, is_output_rank, parallel) for the VAE decode stages.
|
||||
|
||||
The executors consume rank 0's ForwardBatch and the training validation
|
||||
callback consumes each sequence-parallel group leader's, so the output
|
||||
rank is the SP group's first rank (identical to world rank 0 in the
|
||||
single-group e2e case). ``parallel`` is only true when every group rank
|
||||
will run the decode body — the collectives inside require uniform
|
||||
participation, so no rank-dependent branch may guard them.
|
||||
"""
|
||||
if not model_parallel_is_initialized():
|
||||
return None, True, False
|
||||
sp_group = get_sp_group()
|
||||
parallel = bool(want_parallel) and sp_group.world_size > 1
|
||||
return sp_group, sp_group.is_first_rank, parallel
|
||||
|
||||
|
||||
class MiniMaxH3VideoDecodingStage(PipelineStage):
|
||||
"""Drop visual condition rows, unpatchify, and decode the target video."""
|
||||
|
||||
@@ -55,11 +77,14 @@ class MiniMaxH3VideoDecodingStage(PipelineStage):
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if model_parallel_is_initialized() and not get_world_group().is_first_rank:
|
||||
# Distributed executors consume rank 0's ForwardBatch. Keep a
|
||||
"""Decode H3 video latents into normalized CPU pixels."""
|
||||
placeholder = torch.empty((0, 3, 0, 0, 0), device="cpu", dtype=torch.float32)
|
||||
sp_group, is_output_rank, parallel = _decode_participation(fastvideo_args, fastvideo_args.vae_parallel_decode)
|
||||
if not is_output_rank and not parallel:
|
||||
# Consumers read the output rank's ForwardBatch. Keep a
|
||||
# verifier-compatible placeholder on other ranks and avoid
|
||||
# duplicating the full VAE decode and CPU output buffer.
|
||||
batch.output = torch.empty((0, 3, 0, 0, 0), device="cpu", dtype=torch.float32)
|
||||
batch.output = placeholder
|
||||
return batch
|
||||
|
||||
layout = _layout(batch)
|
||||
@@ -79,19 +104,33 @@ class MiniMaxH3VideoDecodingStage(PipelineStage):
|
||||
try:
|
||||
latents = self.vae.denormalize_latents(latents.to(device=device, dtype=torch.float32))
|
||||
if fastvideo_args.output_type == "latent":
|
||||
batch.output = latents.detach().float().cpu()
|
||||
# No collectives on this path, so uniform participation is
|
||||
# trivial: every rank returns here.
|
||||
batch.output = latents.detach().float().cpu() if is_output_rank else placeholder
|
||||
return batch
|
||||
|
||||
output = torch.empty(
|
||||
self.vae.decoded_pixel_shape(latents.shape),
|
||||
device="cpu",
|
||||
dtype=torch.float32,
|
||||
pin_memory=fastvideo_args.pin_cpu_memory and is_pin_memory_available(),
|
||||
)
|
||||
# The published decode recipe uses FP16 autocast over FP32 weights.
|
||||
with torch.autocast(device_type=device.type, dtype=torch.float16, enabled=device.type == "cuda"):
|
||||
self.vae.decode_to_pixels(latents, output)
|
||||
batch.output = output
|
||||
output = None
|
||||
if is_output_rank:
|
||||
output = torch.empty(
|
||||
self.vae.decoded_pixel_shape(latents.shape),
|
||||
device="cpu",
|
||||
dtype=torch.float32,
|
||||
pin_memory=fastvideo_args.pin_cpu_memory and is_pin_memory_available(),
|
||||
)
|
||||
# Attribute the streamed decoder computation while retaining
|
||||
# per-chunk device-to-host transfer and pinned-buffer reuse.
|
||||
with (
|
||||
nvtx_range("minimax_h3.vae"),
|
||||
torch.autocast(device_type=device.type, dtype=torch.float16, enabled=device.type == "cuda"),
|
||||
):
|
||||
if parallel:
|
||||
strategy = fastvideo_args.vae_parallel_decode_strategy or DEFAULT_DECODE_GATHER_STRATEGY
|
||||
logger.info_once(f"MiniMax-H3 VAE decode: sequence-parallel chunks across "
|
||||
f"{sp_group.world_size} ranks ({strategy})")
|
||||
decode_to_pixels_parallel(self.vae, latents, output, sp_group, strategy=strategy)
|
||||
else:
|
||||
self.vae.decode_to_pixels(latents, output)
|
||||
batch.output = output if is_output_rank else placeholder
|
||||
return batch
|
||||
finally:
|
||||
if fastvideo_args.vae_cpu_offload:
|
||||
@@ -121,7 +160,10 @@ class MiniMaxH3AudioDecodingStage(PipelineStage):
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if model_parallel_is_initialized() and not get_world_group().is_first_rank:
|
||||
"""Decode H3 audio latents into a stereo CPU waveform."""
|
||||
# Audio decode is sub-second, so it always runs serially on the SP
|
||||
# group's first rank (the rank whose ForwardBatch consumers read).
|
||||
if model_parallel_is_initialized() and not get_sp_group().is_first_rank:
|
||||
batch.extra["audio"] = torch.empty((0, 2), device="cpu", dtype=torch.float32)
|
||||
batch.extra["audio_sample_rate"] = self.audio_vae.sampling_rate
|
||||
self._clear_runtime(batch)
|
||||
@@ -144,7 +186,10 @@ class MiniMaxH3AudioDecodingStage(PipelineStage):
|
||||
self._clear_runtime(batch)
|
||||
return batch
|
||||
|
||||
decoded = self.audio_vae.decode(latents).sample.float()
|
||||
# The range isolates waveform synthesis from packing and runtime
|
||||
# cleanup so the audio decoder has one stable timeline boundary.
|
||||
with nvtx_range("minimax_h3.audio_vae"):
|
||||
decoded = self.audio_vae.decode(latents).sample.float()
|
||||
if decoded.ndim != 3 or decoded.shape[0] != 2 or decoded.shape[1] != 1:
|
||||
raise ValueError("MiniMax-H3 audio VAE must decode stereo channels as two mono batch items; "
|
||||
f"got {tuple(decoded.shape)}.")
|
||||
|
||||
@@ -11,8 +11,8 @@ from fastvideo.attention.selector import component_attention_backend, get_attn_b
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.profiler import profiler_region
|
||||
from fastvideo.hooks.activation_trace import trace_step
|
||||
from fastvideo.profiler import nvtx_range, profiler_region
|
||||
from fastvideo.pipelines.basic.minimax_h3.packing import (
|
||||
MINIMAX_H3_KEYFRAME_NOISE_AUG,
|
||||
MiniMaxH3PackedLayout,
|
||||
@@ -89,6 +89,7 @@ class MiniMaxH3DenoisingStage(PipelineStage):
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
"""Denoise the packed H3 video and audio streams over one shared schedule."""
|
||||
layout = batch.extra.get(MINIMAX_H3_LAYOUT_KEY)
|
||||
if not isinstance(layout, MiniMaxH3PackedLayout):
|
||||
raise ValueError("MiniMax-H3 packed layout is missing before denoising.")
|
||||
@@ -145,9 +146,15 @@ class MiniMaxH3DenoisingStage(PipelineStage):
|
||||
vsa_exempt = vsa_mode == "exempt"
|
||||
vsa_dense_layers = tuple(batch.extra.get("vsa_dense_layers", ()))
|
||||
vsa_dense_first_n = int(batch.extra.get("vsa_dense_first_n_steps", 0))
|
||||
# Run-level tile geometry (256 default, 64 = native Triton path),
|
||||
# plumbed like the run-level sparsity; the builder validates the
|
||||
# value against VSA_H3_TILE_SHAPES.
|
||||
vsa_tile_size = int(fastvideo_args.VSA_tile_size)
|
||||
|
||||
try:
|
||||
with profiler_region("inference_denoising"):
|
||||
# The stage range groups the complete denoising loop while the
|
||||
# indexed model ranges retain timing detail for every H3 block.
|
||||
with profiler_region("inference_denoising"), nvtx_range("minimax_h3.dit"):
|
||||
for index, (video_timestep,
|
||||
audio_timestep) in enumerate(zip(video_timesteps, audio_timesteps, strict=True)):
|
||||
unique_timesteps, timestep_indices = row_timestep_plan[index]
|
||||
@@ -167,6 +174,7 @@ class MiniMaxH3DenoisingStage(PipelineStage):
|
||||
device=device,
|
||||
exempt=vsa_exempt,
|
||||
dense_layers=vsa_dense_layers,
|
||||
tile_size=vsa_tile_size,
|
||||
)
|
||||
# Under torch.compile(mode="reduce-overhead") each denoising
|
||||
# step must be marked, or cudagraph trees flag cross-step
|
||||
|
||||
@@ -9,8 +9,10 @@ import numpy as np
|
||||
import torch
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.distributed import get_local_torch_device, get_sp_group, model_parallel_is_initialized
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.vaes.minimax_h3_parallel import encode_pixels_parallel
|
||||
from fastvideo.pipelines.basic.minimax_h3.packing import (
|
||||
MINIMAX_H3_AUDIO_CHANNELS,
|
||||
MINIMAX_H3_KEYFRAME_ENCODE_SEED,
|
||||
@@ -36,6 +38,8 @@ from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.pipelines.stages.validators import VerificationResult
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
MINIMAX_H3_LAYOUT_KEY = "minimax_h3_layout"
|
||||
|
||||
|
||||
@@ -105,8 +109,20 @@ class MiniMaxH3LatentPreparationStage(PipelineStage):
|
||||
self,
|
||||
references: list[MiniMaxH3PreparedReference],
|
||||
device: torch.device,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> list[torch.Tensor]:
|
||||
patch_size = self.transformer.patch_size
|
||||
# Reference encode runs on every rank (all ranks hold identical
|
||||
# prepared references), so clip-parallel encode keeps participation
|
||||
# uniform by construction: each rank encodes a clip subset and the
|
||||
# all-gather leaves the identical full posterior everywhere.
|
||||
parallel_group = None
|
||||
if fastvideo_args.vae_parallel_encode and model_parallel_is_initialized():
|
||||
sp_group = get_sp_group()
|
||||
if sp_group.world_size > 1:
|
||||
parallel_group = sp_group
|
||||
logger.info_once(f"MiniMax-H3 reference VAE encode: sequence-parallel clips across "
|
||||
f"{sp_group.world_size} ranks")
|
||||
rows: list[torch.Tensor] = []
|
||||
for reference in references:
|
||||
if reference.media_type == "audio":
|
||||
@@ -120,7 +136,10 @@ class MiniMaxH3LatentPreparationStage(PipelineStage):
|
||||
raise ValueError("MiniMax-H3 reference video frames are missing.")
|
||||
frames = reference.frames[:trim_reference_num_frames(reference.frames.shape[0])]
|
||||
pixels = torch.from_numpy(np.ascontiguousarray(frames)).permute(3, 0, 1, 2)[None]
|
||||
posterior = self.vae.encode_pixels(pixels).latent_dist
|
||||
if parallel_group is not None:
|
||||
posterior = encode_pixels_parallel(self.vae, pixels, parallel_group).latent_dist
|
||||
else:
|
||||
posterior = self.vae.encode_pixels(pixels).latent_dist
|
||||
latents = self.vae.normalize_latents(_sample_visual_posterior(posterior).to(
|
||||
torch.float16).float()).cpu()
|
||||
reference.num_latent_frames = int(latents.shape[2])
|
||||
@@ -201,7 +220,7 @@ class MiniMaxH3LatentPreparationStage(PipelineStage):
|
||||
vae_device = get_local_torch_device()
|
||||
self.vae.to(vae_device)
|
||||
try:
|
||||
video_rows = self._encode_visual_rows(references, vae_device)
|
||||
video_rows = self._encode_visual_rows(references, vae_device, fastvideo_args)
|
||||
finally:
|
||||
if fastvideo_args.vae_cpu_offload:
|
||||
self.vae.to("cpu")
|
||||
|
||||
@@ -38,6 +38,25 @@ logger = init_logger(__name__)
|
||||
_GLOBAL_CONTROLLER: TorchProfilerController | None = None
|
||||
|
||||
|
||||
@contextlib.contextmanager
|
||||
def nvtx_range(name: str):
|
||||
"""Emit one optional NVTX range for an external CUDA profiler.
|
||||
|
||||
``FASTVIDEO_NVTX_PROFILE=1`` enables the marker. The context manager stays
|
||||
a no-op without CUDA so call sites can remain shared with CPU tests.
|
||||
"""
|
||||
enabled = envs.FASTVIDEO_NVTX_PROFILE and torch.cuda.is_available()
|
||||
if not enabled:
|
||||
yield
|
||||
return
|
||||
|
||||
torch.cuda.nvtx.range_push(name)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
torch.cuda.nvtx.range_pop()
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ProfilerRegion:
|
||||
"""Metadata describing a profiler region."""
|
||||
|
||||
@@ -5,20 +5,29 @@ reference. The same reference doubles as the GPU kernel parity oracle."""
|
||||
|
||||
import math
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.attention.backends.video_sparse_attn_h3 import (_TILE_ELEMS, MiniMaxH3VSAImpl,
|
||||
MiniMaxH3VSAMetadataBuilder, _build_block_mask,
|
||||
_pool_tiles, token_tile_and_valid)
|
||||
_pool_tiles, _validate_h3_tile_geometry,
|
||||
token_tile_and_valid)
|
||||
|
||||
_720P = dict(raw_latent_shape=(30, 44, 80), patch_size=(1, 2, 2), prefix_segments=(512, 1760, 400))
|
||||
_TINY = dict(raw_latent_shape=(8, 8, 12), patch_size=(1, 2, 2), prefix_segments=(7, 5, 3))
|
||||
# (4,4,4) coverage: dit grid (9, 10, 13) is ragged in all three dims
|
||||
# (t: 4+4+1, h: 4+4+2, w: 4+4+4+1) and every prefix segment leaves a
|
||||
# partial tail tile at 64 (70 -> 64+6, 5 -> 5, 130 -> 64+64+2).
|
||||
_TINY64 = dict(raw_latent_shape=(9, 20, 26), patch_size=(1, 2, 2), prefix_segments=(70, 5, 130))
|
||||
# production-shape request: 768x1344, 124 frames -> latents (37, 48, 84),
|
||||
# patch (1,2,2) -> token grid (37, 24, 42); text 300 + audio 414 rows.
|
||||
_PROD = dict(raw_latent_shape=(37, 48, 84), patch_size=(1, 2, 2), prefix_segments=(300, 0, 414))
|
||||
|
||||
_CPU = torch.device("cpu")
|
||||
|
||||
|
||||
def _build(spec, sparsity=0.0, device=_CPU):
|
||||
def _build(spec, sparsity=0.0, device=_CPU, tile_size=_TILE_ELEMS):
|
||||
return MiniMaxH3VSAMetadataBuilder().build(
|
||||
current_timestep=0,
|
||||
raw_latent_shape=spec["raw_latent_shape"],
|
||||
@@ -26,6 +35,7 @@ def _build(spec, sparsity=0.0, device=_CPU):
|
||||
VSA_sparsity=sparsity,
|
||||
prefix_segments=spec["prefix_segments"],
|
||||
device=device,
|
||||
tile_size=tile_size,
|
||||
)
|
||||
|
||||
|
||||
@@ -36,7 +46,7 @@ def _impl():
|
||||
def reference_sparse_attention(query, key, value, mask, meta):
|
||||
"""Token-level oracle: SDPA over the padded tile buffer with the block
|
||||
mask expanded to tokens. query/key/value: tiled [B, S_pad, H, D]."""
|
||||
token_tile, token_valid = token_tile_and_valid(meta.variable_block_sizes)
|
||||
token_tile, token_valid = token_tile_and_valid(meta.variable_block_sizes, meta.tile_elems)
|
||||
out = torch.empty_like(query)
|
||||
for b in range(query.shape[0]):
|
||||
for h in range(query.shape[2]):
|
||||
@@ -133,9 +143,117 @@ def test_prefix_queries_stay_dense_at_high_sparsity():
|
||||
"video rows should actually be sparse at 75%"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 64-token (4,4,4) tile geometry
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_geometry_tile64_ragged_tails():
|
||||
"""Hand-computed (4,4,4) oracle on a grid ragged in all three dims."""
|
||||
meta = _build(_TINY64, tile_size=64)
|
||||
assert meta.tile_elems == 64
|
||||
t, h, w = 9, 10, 13 # raw latents (9, 20, 26) under patch (1, 2, 2)
|
||||
n_t, n_h, n_w = 3, 3, 4
|
||||
prefix_len = sum(_TINY64["prefix_segments"])
|
||||
seq = prefix_len + t * h * w
|
||||
assert meta.total_seq_length == seq
|
||||
assert meta.num_prefix_tiles == 2 + 1 + 3
|
||||
assert meta.num_video_tiles == n_t * n_h * n_w
|
||||
assert int(meta.variable_block_sizes.sum()) == seq
|
||||
assert int(meta.variable_block_sizes.max()) <= 64
|
||||
assert meta.variable_block_sizes[:meta.num_prefix_tiles].tolist() == [64, 6, 5, 64, 64, 2]
|
||||
|
||||
# per-tile valid sizes: product of the per-dim clamped tails
|
||||
expected = torch.tensor([
|
||||
min(4, t - 4 * tt) * min(4, h - 4 * hh) * min(4, w - 4 * ww) for tt in range(n_t) for hh in range(n_h)
|
||||
for ww in range(n_w)
|
||||
],
|
||||
dtype=torch.long)
|
||||
assert torch.equal(meta.variable_block_sizes[meta.num_prefix_tiles:], expected)
|
||||
assert int(expected.min()) == 1 * 2 * 1 # the (t,h,w) ragged corner
|
||||
|
||||
# every packed video row lands in the 3D tile its (t,h,w) coordinate says
|
||||
idx = meta.untile_combined_index
|
||||
row = torch.arange(t * h * w)
|
||||
row_t, row_h, row_w = row // (h * w), (row // w) % h, row % w
|
||||
expected_tile = meta.num_prefix_tiles + ((row_t // 4) * n_h + row_h // 4) * n_w + row_w // 4
|
||||
assert torch.equal(idx[prefix_len:] // 64, expected_tile)
|
||||
# and in a non-pad slot of that tile
|
||||
assert bool((idx % 64 < meta.variable_block_sizes[idx // 64]).all())
|
||||
|
||||
# untile(tile(x)) == x on the 64-wide padded buffer
|
||||
x = torch.randn(1, seq, 2, 4)
|
||||
buf = _impl().tile(x, meta)
|
||||
assert buf.shape[1] == meta.variable_block_sizes.numel() * 64
|
||||
assert torch.equal(buf[:, idx], x)
|
||||
|
||||
|
||||
def test_geometry_tile64_production_shape():
|
||||
"""Production latents (37, 48, 84): ragged t and w tails at (4,4,4)."""
|
||||
meta64 = _build(_PROD, tile_size=64)
|
||||
assert meta64.num_prefix_tiles == 5 + 7 # 300 -> 4x64+44, 414 -> 6x64+30
|
||||
assert meta64.num_video_tiles == 10 * 6 * 11 # (37, 24, 42) / (4, 4, 4)
|
||||
assert meta64.total_seq_length == 300 + 414 + 37 * 24 * 42
|
||||
assert int(meta64.variable_block_sizes.sum()) == meta64.total_seq_length
|
||||
sizes_vid = meta64.variable_block_sizes[meta64.num_prefix_tiles:]
|
||||
assert int(sizes_vid.max()) == 64 and int(sizes_vid.min()) == 1 * 4 * 2 # (t, w) ragged corner
|
||||
|
||||
# same packed sequence under the default 256 geometry, fewer tiles
|
||||
meta256 = _build(_PROD)
|
||||
assert meta256.tile_elems == _TILE_ELEMS
|
||||
assert meta256.num_prefix_tiles == 2 + 2
|
||||
assert meta256.num_video_tiles == 10 * 3 * 6
|
||||
assert meta256.total_seq_length == meta64.total_seq_length
|
||||
|
||||
x = torch.randn(1, meta64.total_seq_length, 2, 4)
|
||||
buf = _impl().tile(x, meta64)
|
||||
assert torch.equal(buf[:, meta64.untile_combined_index], x)
|
||||
|
||||
|
||||
def test_sparsity_zero_matches_dense_sdpa_tile64():
|
||||
torch.manual_seed(2)
|
||||
meta = _build(_TINY64, tile_size=64)
|
||||
seq = meta.total_seq_length
|
||||
q, k, v = (torch.randn(1, seq, 2, 8) for _ in range(3))
|
||||
impl = _impl()
|
||||
tq, tk, tv = (impl.tile(t, meta).clone() for t in (q, k, v))
|
||||
|
||||
scores = torch.matmul(_pool_tiles(tq, meta.variable_block_sizes, meta.tile_elems),
|
||||
_pool_tiles(tk, meta.variable_block_sizes, meta.tile_elems).transpose(-2, -1))
|
||||
mask = _build_block_mask(scores, meta.num_prefix_tiles, meta.num_video_tiles, 0.0, exempt=True)
|
||||
sparse_out = impl.postprocess_output(reference_sparse_attention(tq, tk, tv, mask, meta), meta)
|
||||
|
||||
dense_out = F.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)).transpose(1, 2)
|
||||
assert torch.allclose(sparse_out, dense_out, atol=1e-5), (sparse_out - dense_out).abs().max()
|
||||
|
||||
|
||||
def test_geometry_guard_enforces_tile64_bound():
|
||||
"""A 65-token tile passes the 256 bound but must fail the 64 one."""
|
||||
meta = _build(_TINY64, tile_size=64)
|
||||
prefix = tuple(s for s in _TINY64["prefix_segments"] if s > 0)
|
||||
dit_shape = (9, 10, 13)
|
||||
sizes = meta.variable_block_sizes.clone()
|
||||
sizes[0] = 65
|
||||
with pytest.raises(ValueError, match="tile sizes out of bounds"):
|
||||
_validate_h3_tile_geometry(prefix, dit_shape, sizes, meta.untile_combined_index, 64)
|
||||
# the untampered tile-64 geometry passes its own bound
|
||||
_validate_h3_tile_geometry(prefix, dit_shape, meta.variable_block_sizes, meta.untile_combined_index, 64)
|
||||
|
||||
|
||||
def test_builder_rejects_unknown_tile_size():
|
||||
for bad in (0, 128, 512):
|
||||
with pytest.raises(ValueError, match="tile_size"):
|
||||
_build(_TINY, tile_size=bad)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_geometry_720p()
|
||||
test_mask_policy()
|
||||
test_sparsity_zero_matches_dense_sdpa()
|
||||
test_prefix_queries_stay_dense_at_high_sparsity()
|
||||
test_geometry_tile64_ragged_tails()
|
||||
test_geometry_tile64_production_shape()
|
||||
test_sparsity_zero_matches_dense_sdpa_tile64()
|
||||
test_geometry_guard_enforces_tile64_bound()
|
||||
test_builder_rejects_unknown_tile_size()
|
||||
print("all VSA-H3 CPU checks passed")
|
||||
|
||||
@@ -0,0 +1,181 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""CPU checks for the VSA-H3 tile-64 sm_100a route selection.
|
||||
|
||||
The opt-in third kernel route (``FASTVIDEO_VSA_SM100A=1``) must (a) stay off by
|
||||
default, (b) engage only when the extension is present, the device qualifies,
|
||||
and the forward carries no grad, and (c) fall back to the Triton-64 entry with
|
||||
one warning when the env is set but a precondition fails. All device/extension
|
||||
probes are monkeypatched; no GPU or kernel install needed.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
import fastvideo.attention.backends.video_sparse_attn_h3 as vsa_h3
|
||||
from fastvideo.attention.backends.video_sparse_attn_h3 import (VSA_SM100A_ENV, MiniMaxH3VSAImpl,
|
||||
MiniMaxH3VSAMetadataBuilder, _sm100a_unavailable_reason)
|
||||
|
||||
# Small tile-64 geometry: 2 prefix segments + a (4,4,8)-token video grid.
|
||||
_SPEC = dict(raw_latent_shape=(4, 8, 16), patch_size=(1, 2, 2), prefix_segments=(70, 30))
|
||||
_HEADS, _DIM = 2, 128
|
||||
|
||||
|
||||
def _build_meta():
|
||||
return MiniMaxH3VSAMetadataBuilder().build(
|
||||
current_timestep=0,
|
||||
raw_latent_shape=_SPEC["raw_latent_shape"],
|
||||
patch_size=_SPEC["patch_size"],
|
||||
VSA_sparsity=0.0,
|
||||
prefix_segments=_SPEC["prefix_segments"],
|
||||
device=torch.device("cpu"),
|
||||
tile_size=64,
|
||||
)
|
||||
|
||||
|
||||
def _tiled_qkv(meta, requires_grad=False):
|
||||
# bf16 like the real tiled buffers, so forward()'s dtype-cast warning
|
||||
# stays out of the warning assertions below.
|
||||
s_pad = meta.variable_block_sizes.numel() * 64
|
||||
return tuple(
|
||||
torch.randn(1, s_pad, _HEADS, _DIM, dtype=torch.bfloat16, requires_grad=requires_grad) for _ in range(3))
|
||||
|
||||
|
||||
class _FakeSm100a:
|
||||
"""Stands in for fastvideo_kernel.block_sparse_attn_sm100a."""
|
||||
|
||||
def __init__(self, supported=True):
|
||||
self.supported = supported
|
||||
self.calls = []
|
||||
|
||||
def is_supported(self, q, variable_block_sizes):
|
||||
return self.supported
|
||||
|
||||
def block_sparse_attn_sm100a(self, q, k, v, q2k_idx, q2k_num, variable_block_sizes, need_lse=True):
|
||||
self.calls.append(dict(q=q, q2k_idx=q2k_idx, q2k_num=q2k_num, vbs=variable_block_sizes,
|
||||
need_lse=need_lse))
|
||||
return q.clone(), None
|
||||
|
||||
|
||||
def _fake_map_to_index(block_map):
|
||||
"""Pure-torch stand-in for the Triton map_to_index (same contract)."""
|
||||
b, h, t, n = block_map.shape
|
||||
idx = torch.full((b, h, t, n), -1, dtype=torch.int32)
|
||||
num = block_map.sum(dim=-1, dtype=torch.int32)
|
||||
for bi in range(b):
|
||||
for hi in range(h):
|
||||
for ti in range(t):
|
||||
cols = torch.nonzero(block_map[bi, hi, ti], as_tuple=False).flatten()
|
||||
idx[bi, hi, ti, :cols.numel()] = cols.to(torch.int32)
|
||||
return idx, num
|
||||
|
||||
|
||||
class _FakeTriton:
|
||||
def __init__(self):
|
||||
self.calls = 0
|
||||
|
||||
def __call__(self, q, k, v, mask, variable_block_sizes):
|
||||
self.calls += 1
|
||||
return q.clone(), None
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def routed(monkeypatch):
|
||||
"""Backend with both kernel entries faked; returns (fakes, run)."""
|
||||
fake_sm = _FakeSm100a()
|
||||
fake_triton = _FakeTriton()
|
||||
monkeypatch.setattr(vsa_h3, "_sm100a", fake_sm)
|
||||
monkeypatch.setattr(vsa_h3, "block_sparse_attn_64_bhsd", fake_triton)
|
||||
monkeypatch.setattr(vsa_h3, "map_to_index", _fake_map_to_index)
|
||||
meta = _build_meta()
|
||||
impl = MiniMaxH3VSAImpl(num_heads=_HEADS, head_size=_DIM, causal=False, softmax_scale=_DIM**-0.5)
|
||||
|
||||
def run(requires_grad=False):
|
||||
q, k, v = _tiled_qkv(meta, requires_grad=requires_grad)
|
||||
return impl.forward(q, k, v, None, meta)
|
||||
|
||||
return fake_sm, fake_triton, run, meta
|
||||
|
||||
|
||||
def test_reason_covers_every_precondition():
|
||||
q = torch.randn(1, _HEADS, 128, _DIM)
|
||||
vbs = torch.full((2, ), 64, dtype=torch.long)
|
||||
assert "not installed" in _sm100a_unavailable_reason(None, q, vbs, grad_mode=False)
|
||||
ok = _FakeSm100a(supported=True)
|
||||
assert "forward-only" in _sm100a_unavailable_reason(ok, q, vbs, grad_mode=True)
|
||||
bad = _FakeSm100a(supported=False)
|
||||
assert "is_supported" in _sm100a_unavailable_reason(bad, q, vbs, grad_mode=False)
|
||||
assert _sm100a_unavailable_reason(ok, q, vbs, grad_mode=False) is None
|
||||
|
||||
|
||||
def test_default_off_routes_triton(routed, monkeypatch):
|
||||
fake_sm, fake_triton, run, _ = routed
|
||||
monkeypatch.delenv(VSA_SM100A_ENV, raising=False)
|
||||
run()
|
||||
assert fake_triton.calls == 1
|
||||
assert fake_sm.calls == []
|
||||
|
||||
|
||||
def test_env_on_routes_sm100a_with_index_metadata(routed, monkeypatch):
|
||||
fake_sm, fake_triton, run, meta = routed
|
||||
monkeypatch.setenv(VSA_SM100A_ENV, "1")
|
||||
out = run()
|
||||
assert fake_triton.calls == 0
|
||||
assert len(fake_sm.calls) == 1
|
||||
call = fake_sm.calls[0]
|
||||
n_tiles = meta.variable_block_sizes.numel()
|
||||
# sparsity 0 -> all-True mask -> every row's count is n_tiles
|
||||
assert call["q2k_num"].dtype == torch.int32 and (call["q2k_num"] == n_tiles).all()
|
||||
assert call["q2k_idx"].shape[-1] == n_tiles and call["q2k_idx"].dtype == torch.int32
|
||||
assert call["vbs"].dtype == torch.int32
|
||||
assert call["need_lse"] is False
|
||||
# BHSD kernel result comes back in the backend's BSHD layout
|
||||
assert out.shape == (1, n_tiles * 64, _HEADS, _DIM)
|
||||
|
||||
|
||||
def test_env_on_grad_inputs_fall_back_to_triton(routed, monkeypatch):
|
||||
fake_sm, fake_triton, run, _ = routed
|
||||
monkeypatch.setenv(VSA_SM100A_ENV, "1")
|
||||
run(requires_grad=True)
|
||||
assert fake_triton.calls == 1
|
||||
assert fake_sm.calls == []
|
||||
# ...but the same process still routes no-grad forwards to sm_100a
|
||||
run(requires_grad=False)
|
||||
assert len(fake_sm.calls) == 1
|
||||
|
||||
|
||||
def test_env_on_unsupported_warns_once_and_falls_back(routed, monkeypatch):
|
||||
fake_sm, fake_triton, run, _ = routed
|
||||
fake_sm.supported = False
|
||||
monkeypatch.setenv(VSA_SM100A_ENV, "1")
|
||||
warnings = []
|
||||
monkeypatch.setattr(vsa_h3.logger, "warning_once", warnings.append)
|
||||
run()
|
||||
run()
|
||||
assert fake_triton.calls == 2
|
||||
assert fake_sm.calls == []
|
||||
assert len(warnings) == 2 # warning_once dedups by message; both carry the same one line
|
||||
assert warnings[0] == warnings[1]
|
||||
assert VSA_SM100A_ENV in warnings[0] and "is_supported" in warnings[0]
|
||||
|
||||
|
||||
def test_env_on_missing_module_warns_and_falls_back(routed, monkeypatch):
|
||||
fake_sm, fake_triton, run, _ = routed
|
||||
monkeypatch.setattr(vsa_h3, "_sm100a", None)
|
||||
monkeypatch.setenv(VSA_SM100A_ENV, "1")
|
||||
warnings = []
|
||||
monkeypatch.setattr(vsa_h3.logger, "warning_once", warnings.append)
|
||||
run()
|
||||
assert fake_triton.calls == 1
|
||||
assert warnings and "not installed" in warnings[0]
|
||||
|
||||
|
||||
def test_env_on_no_grad_context_detaches_route_from_leaf_flags(routed, monkeypatch):
|
||||
"""A requires_grad leaf under torch.no_grad() is still a no-grad forward."""
|
||||
fake_sm, fake_triton, run, meta = routed
|
||||
monkeypatch.setenv(VSA_SM100A_ENV, "1")
|
||||
impl = MiniMaxH3VSAImpl(num_heads=_HEADS, head_size=_DIM, causal=False, softmax_scale=_DIM**-0.5)
|
||||
q, k, v = _tiled_qkv(meta, requires_grad=True)
|
||||
with torch.no_grad():
|
||||
impl.forward(q, k, v, None, meta)
|
||||
assert len(fake_sm.calls) == 1
|
||||
assert fake_triton.calls == 0
|
||||
@@ -15,6 +15,12 @@ import json
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.profiler import nvtx_range
|
||||
|
||||
# Five-window child: ops before any region, inside a region, between regions,
|
||||
# inside a second (short-named) region, after the last region. Exits without
|
||||
@@ -105,3 +111,73 @@ def test_noop_without_profiler_dir(tmp_path):
|
||||
proc = subprocess.run([sys.executable, "-c", child], env=env,
|
||||
capture_output=True, text=True, timeout=300)
|
||||
assert proc.returncode == 0, proc.stderr
|
||||
|
||||
|
||||
def test_nvtx_range_disabled_is_noop(monkeypatch):
|
||||
"""Keep CUDA NVTX untouched when external profiling is disabled."""
|
||||
range_push = Mock()
|
||||
range_pop = Mock()
|
||||
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "0")
|
||||
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
|
||||
monkeypatch.setattr(torch.cuda.nvtx, "range_push", range_push)
|
||||
monkeypatch.setattr(torch.cuda.nvtx, "range_pop", range_pop)
|
||||
|
||||
with nvtx_range("disabled"):
|
||||
body_executed = True
|
||||
|
||||
assert body_executed is True
|
||||
range_push.assert_not_called()
|
||||
range_pop.assert_not_called()
|
||||
|
||||
|
||||
def test_nvtx_range_without_cuda_is_noop(monkeypatch):
|
||||
"""Keep NVTX untouched when profiling is enabled on a CPU-only process."""
|
||||
range_push = Mock()
|
||||
range_pop = Mock()
|
||||
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "1")
|
||||
monkeypatch.setattr(torch.cuda, "is_available", lambda: False)
|
||||
monkeypatch.setattr(torch.cuda.nvtx, "range_push", range_push)
|
||||
monkeypatch.setattr(torch.cuda.nvtx, "range_pop", range_pop)
|
||||
|
||||
with nvtx_range("cpu-only"):
|
||||
body_executed = True
|
||||
|
||||
assert body_executed is True
|
||||
range_push.assert_not_called()
|
||||
range_pop.assert_not_called()
|
||||
|
||||
|
||||
def test_nvtx_range_enabled_orders_push_body_pop(monkeypatch):
|
||||
"""Place the profiled body between one matching NVTX push and pop."""
|
||||
events = []
|
||||
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "1")
|
||||
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
|
||||
monkeypatch.setattr(torch.cuda.nvtx, "range_push", lambda name: events.append(("push", name)))
|
||||
monkeypatch.setattr(torch.cuda.nvtx, "range_pop", lambda: events.append(("pop", None)))
|
||||
|
||||
with nvtx_range("minimax_h3.test"):
|
||||
events.append(("body", None))
|
||||
|
||||
assert events == [
|
||||
("push", "minimax_h3.test"),
|
||||
("body", None),
|
||||
("pop", None),
|
||||
]
|
||||
|
||||
|
||||
def test_nvtx_range_body_exception_pops_and_propagates(monkeypatch):
|
||||
"""Balance the NVTX stack while preserving a body exception."""
|
||||
events = []
|
||||
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "1")
|
||||
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
|
||||
monkeypatch.setattr(torch.cuda.nvtx, "range_push", lambda name: events.append(("push", name)))
|
||||
monkeypatch.setattr(torch.cuda.nvtx, "range_pop", lambda: events.append(("pop", None)))
|
||||
|
||||
with pytest.raises(RuntimeError, match="profile body failed"):
|
||||
with nvtx_range("minimax_h3.failure"):
|
||||
raise RuntimeError("profile body failed")
|
||||
|
||||
assert events == [
|
||||
("push", "minimax_h3.failure"),
|
||||
("pop", None),
|
||||
]
|
||||
|
||||
@@ -0,0 +1,281 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29514")
|
||||
|
||||
import fastvideo.models.encoders.minimax_h3_checkpoint_fp8 as h3_fp8
|
||||
from fastvideo.configs.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConfig
|
||||
from fastvideo.layers.linear import ColumnParallelLinear, UnquantizedLinearMethod
|
||||
from fastvideo.layers.vocab_parallel_embedding import UnquantizedEmbeddingMethod, VocabParallelEmbedding
|
||||
from fastvideo.models.encoders.base import TextEncoder
|
||||
from fastvideo.models.encoders.minimax_h3_checkpoint_fp8 import (
|
||||
MiniMaxH3SerializedFP8Config,
|
||||
MiniMaxH3SerializedFP8LinearMethod,
|
||||
)
|
||||
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
|
||||
from fastvideo.models.loader.text_encoder_quantization import (
|
||||
_configure_text_encoder_quantization,
|
||||
_process_quantized_text_encoder_weights,
|
||||
_read_text_encoder_checkpoint_quantization_config,
|
||||
)
|
||||
|
||||
|
||||
def _checkpoint_quantization_config(**overrides) -> dict:
|
||||
config = {
|
||||
"quant_method": "fp8",
|
||||
"activation_scheme": "dynamic",
|
||||
"fmt": "e4m3",
|
||||
"weight_block_size": [128, 128],
|
||||
"modules_to_not_convert": ["model.visual", "lm_head"],
|
||||
}
|
||||
config.update(overrides)
|
||||
return config
|
||||
|
||||
|
||||
def test_h3_accepts_only_the_serialized_blockwise_checkpoint_contract() -> None:
|
||||
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
|
||||
assert config.weight_block_size == (128, 128)
|
||||
assert config.get_supported_act_dtypes() == [torch.bfloat16]
|
||||
|
||||
with pytest.raises(ValueError, match=r"weight_block_size=\[128, 128\]"):
|
||||
MiniMaxH3SerializedFP8Config.from_config(
|
||||
_checkpoint_quantization_config(weight_block_size=[1, 128]))
|
||||
with pytest.raises(ValueError, match="dynamic activation"):
|
||||
MiniMaxH3SerializedFP8Config.from_config(
|
||||
_checkpoint_quantization_config(activation_scheme="static"))
|
||||
with pytest.raises(ValueError, match="vision stack"):
|
||||
MiniMaxH3SerializedFP8Config.from_config(
|
||||
_checkpoint_quantization_config(modules_to_not_convert=["lm_head"]))
|
||||
with pytest.raises(ValueError, match="partially quantized language"):
|
||||
MiniMaxH3SerializedFP8Config.from_config(
|
||||
_checkpoint_quantization_config(modules_to_not_convert=["model.visual", "language_model.layers.3"]))
|
||||
|
||||
|
||||
def test_serialized_fp8_allocates_checkpoint_weight_and_scale_without_requantization(distributed_setup) -> None:
|
||||
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
|
||||
layer = ColumnParallelLinear(
|
||||
input_size=128,
|
||||
output_size=256,
|
||||
bias=False,
|
||||
quant_config=config,
|
||||
prefix="minimax_h3_qwen3_vl.language_model.layers.0.self_attn.q_proj",
|
||||
)
|
||||
|
||||
assert isinstance(layer.quant_method, MiniMaxH3SerializedFP8LinearMethod)
|
||||
assert layer.weight.dtype == torch.float8_e4m3fn
|
||||
assert layer.weight.shape == (256, 128)
|
||||
assert layer.weight_scale_inv.dtype == torch.float32
|
||||
assert layer.weight_scale_inv.shape == (2, 1)
|
||||
|
||||
layer.weight.data.zero_()
|
||||
layer.weight_scale_inv.data.fill_(0.25)
|
||||
weight_pointer = layer.weight.data_ptr()
|
||||
scale_pointer = layer.weight_scale_inv.data_ptr()
|
||||
layer.quant_method.process_weights_after_loading(layer)
|
||||
|
||||
assert layer.weight.data_ptr() == weight_pointer
|
||||
assert layer.weight_scale_inv.data_ptr() == scale_pointer
|
||||
assert not hasattr(layer, "_fp8_weight")
|
||||
|
||||
|
||||
def test_serialized_fp8_quantizes_only_language_linears(distributed_setup) -> None:
|
||||
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
|
||||
visual_linear = ColumnParallelLinear(
|
||||
input_size=128,
|
||||
output_size=128,
|
||||
bias=False,
|
||||
quant_config=config,
|
||||
prefix="minimax_h3_qwen3_vl.visual.blocks.0.attn.proj",
|
||||
)
|
||||
embedding = VocabParallelEmbedding(
|
||||
num_embeddings=128,
|
||||
embedding_dim=128,
|
||||
org_num_embeddings=128,
|
||||
quant_config=config,
|
||||
prefix="minimax_h3_qwen3_vl.language_model.embed_tokens",
|
||||
)
|
||||
|
||||
assert isinstance(visual_linear.quant_method, UnquantizedLinearMethod)
|
||||
assert visual_linear.weight.dtype == torch.get_default_dtype()
|
||||
assert isinstance(embedding.quant_method, UnquantizedEmbeddingMethod)
|
||||
assert embedding.weight.dtype == torch.get_default_dtype()
|
||||
|
||||
|
||||
def test_serialized_fp8_cpu_execution_fails_closed(distributed_setup) -> None:
|
||||
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
|
||||
layer = ColumnParallelLinear(
|
||||
input_size=128,
|
||||
output_size=128,
|
||||
bias=False,
|
||||
quant_config=config,
|
||||
prefix="minimax_h3_qwen3_vl.language_model.layers.0.mlp.up_proj",
|
||||
)
|
||||
layer.weight.data.zero_()
|
||||
layer.weight_scale_inv.data.fill_(1.0)
|
||||
assert isinstance(layer.quant_method, MiniMaxH3SerializedFP8LinearMethod)
|
||||
layer.quant_method.process_weights_after_loading(layer)
|
||||
|
||||
with pytest.raises(RuntimeError, match="requires CUDA"):
|
||||
layer(torch.zeros(2, 128, dtype=torch.bfloat16))
|
||||
|
||||
|
||||
def test_runtime_preflight_reports_capability_and_missing_dependencies(monkeypatch) -> None:
|
||||
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
|
||||
monkeypatch.setattr(torch.cuda, "get_device_capability", lambda device: (8, 0))
|
||||
with pytest.raises(RuntimeError, match="sm100 or newer"):
|
||||
config.validate_runtime(torch.device("cuda"))
|
||||
|
||||
monkeypatch.setattr(torch.cuda, "get_device_capability", lambda device: (10, 0))
|
||||
|
||||
def missing_quantizer() -> None:
|
||||
raise RuntimeError("SGLang-compatible Triton quantizer is missing")
|
||||
|
||||
monkeypatch.setattr(h3_fp8, "_require_sglang_per_token_group_fp8_quantization", missing_quantizer)
|
||||
with pytest.raises(RuntimeError, match="Triton quantizer is missing"):
|
||||
config.validate_runtime(torch.device("cuda"))
|
||||
|
||||
monkeypatch.setattr(h3_fp8, "_require_sglang_per_token_group_fp8_quantization", lambda: None)
|
||||
|
||||
def missing_flashinfer():
|
||||
raise RuntimeError("FlashInfer groupwise GEMM is missing")
|
||||
|
||||
monkeypatch.setattr(h3_fp8, "_get_flashinfer_groupwise_fp8_gemm", missing_flashinfer)
|
||||
with pytest.raises(RuntimeError, match="FlashInfer groupwise GEMM is missing"):
|
||||
config.validate_runtime(torch.device("cuda"))
|
||||
|
||||
|
||||
def test_loader_detects_and_capability_gates_checkpoint_metadata(tmp_path) -> None:
|
||||
checkpoint_config = _checkpoint_quantization_config()
|
||||
(tmp_path / "config.json").write_text(
|
||||
json.dumps({"quantization_config": checkpoint_config}),
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
assert _read_text_encoder_checkpoint_quantization_config(str(tmp_path)) == checkpoint_config
|
||||
model_config = MiniMaxH3Qwen3VLConfig()
|
||||
quant_config = _configure_text_encoder_quantization(
|
||||
model_config,
|
||||
MiniMaxH3Qwen3VLConditioner,
|
||||
str(tmp_path),
|
||||
)
|
||||
assert isinstance(quant_config, MiniMaxH3SerializedFP8Config)
|
||||
assert model_config.quant_config is quant_config
|
||||
|
||||
unsupported_config = MiniMaxH3Qwen3VLConfig()
|
||||
with pytest.raises(ValueError, match="does not support serialized 'fp8'"):
|
||||
_configure_text_encoder_quantization(
|
||||
unsupported_config,
|
||||
TextEncoder,
|
||||
str(tmp_path),
|
||||
)
|
||||
|
||||
|
||||
def test_loader_leaves_bf16_checkpoint_path_unchanged(tmp_path) -> None:
|
||||
(tmp_path / "config.json").write_text(json.dumps({"architectures": ["Qwen3VLModel"]}), encoding="utf-8")
|
||||
model_config = MiniMaxH3Qwen3VLConfig()
|
||||
|
||||
quant_config = _configure_text_encoder_quantization(
|
||||
model_config,
|
||||
MiniMaxH3Qwen3VLConditioner,
|
||||
str(tmp_path),
|
||||
)
|
||||
|
||||
assert quant_config is None
|
||||
assert model_config.quant_config is None
|
||||
|
||||
|
||||
def test_post_load_processing_visits_only_serialized_fp8_linears(distributed_setup) -> None:
|
||||
config = MiniMaxH3SerializedFP8Config.from_config(_checkpoint_quantization_config())
|
||||
quantized = ColumnParallelLinear(
|
||||
input_size=128,
|
||||
output_size=128,
|
||||
bias=False,
|
||||
quant_config=config,
|
||||
prefix="minimax_h3_qwen3_vl.language_model.layers.0.self_attn.q_proj",
|
||||
)
|
||||
plain = ColumnParallelLinear(
|
||||
input_size=128,
|
||||
output_size=128,
|
||||
bias=False,
|
||||
prefix="plain",
|
||||
)
|
||||
quantized.weight.data.zero_()
|
||||
quantized.weight_scale_inv.data.fill_(1.0)
|
||||
model = torch.nn.ModuleList([quantized, plain])
|
||||
|
||||
assert _process_quantized_text_encoder_weights(model, torch.device("cpu")) == 1
|
||||
assert quantized.weight.device.type == "cpu"
|
||||
assert plain.weight.device.type == "cpu"
|
||||
|
||||
|
||||
def test_flashinfer_groupwise_path_pins_output_dtype_and_trtllm_scale_layout(monkeypatch) -> None:
|
||||
input_tensor = torch.zeros(2, 256, dtype=torch.bfloat16)
|
||||
weight = torch.zeros(128, 256, dtype=torch.float8_e4m3fn)
|
||||
weight_scale = torch.ones(1, 2, dtype=torch.float32)
|
||||
quantized_input = torch.zeros_like(input_tensor, dtype=torch.float8_e4m3fn)
|
||||
input_scale = torch.empty(2, 2, dtype=torch.float32).t()
|
||||
input_scale.fill_(1.0)
|
||||
receipt: dict[str, object] = {}
|
||||
|
||||
def fake_quantize(
|
||||
value: torch.Tensor,
|
||||
group_size: int,
|
||||
*,
|
||||
column_major_scales: bool,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
assert value.data_ptr() == input_tensor.data_ptr()
|
||||
assert value.shape == input_tensor.shape
|
||||
assert group_size == 128
|
||||
assert column_major_scales is True
|
||||
return quantized_input, input_scale
|
||||
|
||||
def fake_gemm(
|
||||
activation: torch.Tensor,
|
||||
checkpoint_weight: torch.Tensor,
|
||||
activation_scale: torch.Tensor,
|
||||
checkpoint_scale: torch.Tensor,
|
||||
*,
|
||||
out_dtype: torch.dtype,
|
||||
backend: str,
|
||||
) -> torch.Tensor:
|
||||
receipt.update(
|
||||
activation=activation,
|
||||
checkpoint_weight=checkpoint_weight,
|
||||
activation_scale=activation_scale,
|
||||
checkpoint_scale=checkpoint_scale,
|
||||
out_dtype=out_dtype,
|
||||
backend=backend,
|
||||
)
|
||||
return torch.zeros(activation.shape[0], checkpoint_weight.shape[0], dtype=out_dtype)
|
||||
|
||||
monkeypatch.setattr(h3_fp8, "_get_flashinfer_groupwise_backend", lambda device: "trtllm")
|
||||
monkeypatch.setattr(h3_fp8, "_sglang_per_token_group_quant_fp8", fake_quantize)
|
||||
monkeypatch.setattr(h3_fp8, "_get_flashinfer_groupwise_fp8_gemm", lambda: fake_gemm)
|
||||
|
||||
previous_default_dtype = torch.get_default_dtype()
|
||||
torch.set_default_dtype(torch.float32)
|
||||
try:
|
||||
output = h3_fp8._flashinfer_gemm_w8a8_block_fp8_linear_with_fallback(
|
||||
input_tensor,
|
||||
weight,
|
||||
(128, 128),
|
||||
weight_scale,
|
||||
)
|
||||
assert torch.get_default_dtype() == torch.float32
|
||||
finally:
|
||||
torch.set_default_dtype(previous_default_dtype)
|
||||
|
||||
assert output.dtype == torch.bfloat16
|
||||
assert receipt["out_dtype"] == torch.bfloat16
|
||||
assert receipt["backend"] == "trtllm"
|
||||
assert receipt["activation"] is quantized_input
|
||||
assert receipt["checkpoint_weight"] is weight
|
||||
assert receipt["checkpoint_scale"] is weight_scale
|
||||
assert receipt["activation_scale"] is input_scale
|
||||
@@ -1,27 +1,13 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""The Qwen3-VL stack is built only as far as MiniMax H3 reads.
|
||||
|
||||
H3 conditions on one intermediate hidden state. The layers above it were built,
|
||||
weight-loaded and then discarded, which is 13.7 GB in bf16 and the difference
|
||||
between fitting and not fitting on a 121 GB unified-memory device.
|
||||
|
||||
The dangerous part is not the truncation, it is getting the tuple index wrong.
|
||||
`hidden_states` records each layer's *input*, so entry N is the output of layer
|
||||
N-1, and the final entry comes from the norm that sits above the whole stack. A
|
||||
truncated stack that still applies that norm puts a normalised tensor where the
|
||||
raw one belongs: the length check in the conditioning stage still passes, and
|
||||
conditioning silently changes. These tests pin the index, the content, and the
|
||||
constant the two sides agree on.
|
||||
"""
|
||||
"""MiniMax-H3 Qwen3-VL layer truncation and slim-forward tests."""
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
# Matches the other encoder tests: the module registry these build against wants
|
||||
# a process group, and a single-rank one needs a rendezvous address.
|
||||
os.environ.setdefault("MASTER_ADDR", "localhost")
|
||||
os.environ.setdefault("MASTER_PORT", "29513")
|
||||
|
||||
@@ -34,21 +20,15 @@ from fastvideo.pipelines.basic.minimax_h3.packing import MINIMAX_H3_TEXT_ENCODER
|
||||
|
||||
|
||||
def _small_arch(**overrides) -> MiniMaxH3Qwen3VLArchConfig:
|
||||
"""A stack small enough to run on CPU but shaped like the real one.
|
||||
|
||||
Everything goes through the constructor so ``__post_init__`` validates the
|
||||
small shape the same way it validates the real one.
|
||||
"""
|
||||
kwargs: dict = dict(
|
||||
vocab_size=64,
|
||||
hidden_size=16,
|
||||
intermediate_size=32,
|
||||
num_hidden_layers=8,
|
||||
output_hidden_state_index=5,
|
||||
num_attention_heads=2,
|
||||
num_key_value_heads=1,
|
||||
head_dim=8,
|
||||
# __post_init__ reads the sections out of rope_scaling, and they must
|
||||
# cover exactly half of each head.
|
||||
rope_scaling={
|
||||
"mrope_interleaved": True,
|
||||
"mrope_section": [2, 1, 1],
|
||||
@@ -61,25 +41,27 @@ def _small_arch(**overrides) -> MiniMaxH3Qwen3VLArchConfig:
|
||||
|
||||
|
||||
def _small_config(**overrides) -> MiniMaxH3Qwen3VLConfig:
|
||||
"""The outer config, which is what the modules take.
|
||||
|
||||
``ModelConfig.__getattr__`` forwards the architecture fields, so the modules
|
||||
read ``prefix`` off this object and everything else off ``arch_config``.
|
||||
"""
|
||||
config = MiniMaxH3Qwen3VLConfig()
|
||||
config.arch_config = _small_arch(**overrides)
|
||||
return config
|
||||
|
||||
|
||||
def test_default_matches_the_index_the_pipeline_reads() -> None:
|
||||
"""The two sides cannot import each other, so pin them here instead.
|
||||
config = MiniMaxH3Qwen3VLArchConfig()
|
||||
|
||||
`fastvideo/models/` must not import from `fastvideo/pipelines/`, so the tap
|
||||
is written down twice. If they drift, conditioning reads a hidden state that
|
||||
was never built and the run dies with an index error at generation time,
|
||||
after a full model load.
|
||||
"""
|
||||
assert MiniMaxH3Qwen3VLArchConfig().num_hidden_layers_override == MINIMAX_H3_TEXT_ENCODER_LAYER
|
||||
assert config.output_hidden_state_index == MINIMAX_H3_TEXT_ENCODER_LAYER
|
||||
assert config.num_hidden_layers_override == MINIMAX_H3_TEXT_ENCODER_LAYER
|
||||
|
||||
|
||||
def test_rejects_build_depth_that_cannot_reach_the_output() -> None:
|
||||
for override in (0, 4):
|
||||
with pytest.raises(ValueError, match="num_hidden_layers_override"):
|
||||
_small_arch(num_hidden_layers_override=override)
|
||||
|
||||
|
||||
def test_rejects_output_index_above_the_checkpoint_depth() -> None:
|
||||
with pytest.raises(ValueError, match="output_hidden_state_index"):
|
||||
_small_arch(output_hidden_state_index=9, num_hidden_layers_override=None)
|
||||
|
||||
|
||||
def test_builds_only_up_to_the_override(distributed_setup) -> None:
|
||||
@@ -87,7 +69,6 @@ def test_builds_only_up_to_the_override(distributed_setup) -> None:
|
||||
|
||||
assert model.num_layers == 5
|
||||
assert len(model.layers) == 5
|
||||
# The norm sits above the tap, so a truncated stack must not keep it.
|
||||
assert model.norm is None
|
||||
|
||||
|
||||
@@ -98,104 +79,94 @@ def test_override_none_keeps_the_full_stack(distributed_setup) -> None:
|
||||
assert model.norm is not None
|
||||
|
||||
|
||||
def test_nominal_and_built_depths_remain_distinct(distributed_setup) -> None:
|
||||
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
|
||||
|
||||
conditioner = MiniMaxH3Qwen3VLConditioner(_small_config(num_hidden_layers_override=5))
|
||||
|
||||
assert conditioner.num_hidden_layers == 8
|
||||
assert conditioner.num_built_hidden_layers == 5
|
||||
|
||||
|
||||
def test_override_above_the_stack_does_not_over_build(distributed_setup) -> None:
|
||||
# num_hidden_layers comes from the checkpoint's config.json via
|
||||
# update_model_arch, so a smaller variant must clamp rather than ask for
|
||||
# layers that do not exist.
|
||||
model = MiniMaxH3Qwen3VLLanguageModel(_small_config(num_hidden_layers_override=99))
|
||||
|
||||
assert model.num_layers == 8
|
||||
assert model.norm is not None
|
||||
|
||||
|
||||
def test_override_equal_to_the_stack_keeps_the_norm(distributed_setup) -> None:
|
||||
"""The exact boundary of the clamp: a stack cut at its own depth is full.
|
||||
|
||||
A checkpoint with exactly ``override`` layers taps its final layer, whose
|
||||
tuple entry sits after the norm in the full model, so the norm must stay
|
||||
and nothing may be filtered from the checkpoint.
|
||||
"""
|
||||
model = MiniMaxH3Qwen3VLLanguageModel(_small_config(num_hidden_layers_override=8))
|
||||
|
||||
assert model.num_layers == 8
|
||||
assert model.norm is not None
|
||||
|
||||
|
||||
def test_non_positive_override_is_rejected() -> None:
|
||||
"""A non-positive override would build no decoder layers at all.
|
||||
|
||||
Worse, a negative one makes ``num_layers`` disagree with the built stack
|
||||
and the surplus-key filter would then drop every layer key, so the
|
||||
conditioner would load "successfully" with no transformer. Reject it at
|
||||
config construction, and again when update_model_arch re-validates.
|
||||
"""
|
||||
for override in (0, -1):
|
||||
with pytest.raises(ValueError, match="num_hidden_layers_override"):
|
||||
_small_arch(num_hidden_layers_override=override)
|
||||
|
||||
config = _small_config()
|
||||
with pytest.raises(ValueError, match="num_hidden_layers_override"):
|
||||
config.update_model_arch({"num_hidden_layers_override": 0})
|
||||
|
||||
|
||||
def test_tapped_hidden_state_is_unchanged_by_truncation(distributed_setup) -> None:
|
||||
"""The whole point: entry `tap` must be bit-identical either way."""
|
||||
"""The slim model returns the raw output at the selected layer."""
|
||||
tap = 5
|
||||
full = MiniMaxH3Qwen3VLLanguageModel(_small_config(num_hidden_layers_override=None))
|
||||
cut = MiniMaxH3Qwen3VLLanguageModel(_small_config(num_hidden_layers_override=tap))
|
||||
|
||||
# These modules allocate uninitialised storage and expect a checkpoint, so
|
||||
# give them finite weights before running anything through them.
|
||||
torch.manual_seed(0)
|
||||
for parameter in full.parameters():
|
||||
parameter.data.normal_(std=0.02)
|
||||
# Then make the shared prefix identical, which is the only part the tapped
|
||||
# hidden state depends on.
|
||||
for (_, a), (_, b) in zip(full.layers[:tap].named_parameters(),
|
||||
cut.layers[:tap].named_parameters(),
|
||||
strict=True):
|
||||
b.data.copy_(a.data)
|
||||
torch.manual_seed(1)
|
||||
inputs_embeds = torch.randn(1, 6, 16)
|
||||
# mRoPE indexes three axes (t, h, w); text tokens share the same position on
|
||||
# all three.
|
||||
position_ids = torch.arange(6).view(1, 1, 6).expand(3, 1, 6)
|
||||
with torch.no_grad():
|
||||
full_out = full(inputs_embeds, position_ids, None, True, None, None)
|
||||
cut_out = cut(inputs_embeds, position_ids, None, True, None, None)
|
||||
expected = inputs_embeds
|
||||
position_embeddings = full.rotary_emb(inputs_embeds, position_ids)
|
||||
for layer in full.layers[:tap]:
|
||||
expected = layer(expected, position_embeddings, None)
|
||||
full_out = full(inputs_embeds, position_ids, None, None, None)
|
||||
cut_out = cut(inputs_embeds, position_ids, None, None, None)
|
||||
|
||||
assert torch.equal(full_out.hidden_states[tap], cut_out.hidden_states[tap])
|
||||
# And the truncated model must not offer states it never computed.
|
||||
assert len(cut_out.hidden_states) == tap + 1
|
||||
# The whole shared prefix must match, not just the tap: this is the same
|
||||
# comparison the production-loader parity gate runs against the official
|
||||
# model, and it is what catches a truncated stack that still applied the
|
||||
# final norm to its last entry.
|
||||
for index, (cut_state, full_state) in enumerate(zip(cut_out.hidden_states, full_out.hidden_states,
|
||||
strict=False)):
|
||||
assert torch.equal(cut_state, full_state), f"hidden state {index} changed under truncation"
|
||||
assert torch.equal(expected, full_out)
|
||||
assert torch.equal(expected, cut_out)
|
||||
|
||||
|
||||
def test_conditioning_stage_adapts_slim_sequence_output() -> None:
|
||||
from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_conditioning import MiniMaxH3ConditioningStage
|
||||
|
||||
class FakeConditioner:
|
||||
|
||||
dtype = torch.float32
|
||||
|
||||
def __call__(self, input_ids: torch.Tensor, **kwargs) -> torch.Tensor:
|
||||
assert input_ids.ndim == 1
|
||||
assert not kwargs
|
||||
return torch.ones(input_ids.shape[0], 4)
|
||||
|
||||
stage = MiniMaxH3ConditioningStage(conditioner=FakeConditioner(), tokenizer=None, processor=None, ref2va=False)
|
||||
embeddings, tags = stage._encode_tokens([1, 2, 3], [0, 0, 0], torch.device("cpu"))
|
||||
|
||||
assert embeddings.shape == (1, 3, 4)
|
||||
assert tags.shape == (3, )
|
||||
|
||||
|
||||
def test_conditioner_exposes_only_the_slim_forward_contract() -> None:
|
||||
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
|
||||
|
||||
assert tuple(inspect.signature(MiniMaxH3Qwen3VLConditioner.forward).parameters) == (
|
||||
"self",
|
||||
"input_ids",
|
||||
"pixel_values",
|
||||
"image_grid_thw",
|
||||
"pixel_values_videos",
|
||||
"video_grid_thw",
|
||||
)
|
||||
|
||||
|
||||
def test_truncated_model_drops_the_surplus_checkpoint_keys(distributed_setup) -> None:
|
||||
"""The unexpected-key check is strict on purpose, so the surplus keys have
|
||||
to be filtered rather than the check relaxed."""
|
||||
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
|
||||
|
||||
conditioner = MiniMaxH3Qwen3VLConditioner(_small_config(num_hidden_layers_override=5))
|
||||
|
||||
assert conditioner._is_above_the_tap("language_model.layers.5.mlp.gate_proj.weight")
|
||||
assert conditioner._is_above_the_tap("language_model.layers.7.self_attn.q_proj.weight")
|
||||
assert conditioner._is_above_the_tap("language_model.norm.weight")
|
||||
# Kept: layers we built, the embeddings, and the vision tower.
|
||||
assert not conditioner._is_above_the_tap("language_model.layers.4.mlp.gate_proj.weight")
|
||||
assert not conditioner._is_above_the_tap("language_model.embed_tokens.weight")
|
||||
assert not conditioner._is_above_the_tap("visual.blocks.0.attn.qkv.weight")
|
||||
# The filter only drops indexes the full stack would have built. A key at
|
||||
# or above the checkpoint's own num_hidden_layers is corrupt, and it must
|
||||
# keep raising as unexpected exactly as it does without truncation.
|
||||
assert not conditioner._is_above_the_tap("language_model.layers.8.mlp.gate_proj.weight")
|
||||
with pytest.raises(ValueError, match="Unexpected"):
|
||||
conditioner.load_weights([("model.language_model.layers.8.mlp.gate_proj.weight", torch.zeros(1))])
|
||||
assert conditioner._is_omitted_checkpoint_key("language_model.layers.5.mlp.gate_proj.weight")
|
||||
assert conditioner._is_omitted_checkpoint_key("language_model.layers.7.self_attn.q_proj.weight")
|
||||
assert conditioner._is_omitted_checkpoint_key("language_model.norm.weight")
|
||||
assert not conditioner._is_omitted_checkpoint_key("language_model.layers.4.mlp.gate_proj.weight")
|
||||
assert not conditioner._is_omitted_checkpoint_key("language_model.embed_tokens.weight")
|
||||
assert not conditioner._is_omitted_checkpoint_key("visual.blocks.0.attn.qkv.weight")
|
||||
assert not conditioner._is_omitted_checkpoint_key("language_model.layers.8.mlp.gate_proj.weight")
|
||||
|
||||
|
||||
def test_full_stack_filters_nothing(distributed_setup) -> None:
|
||||
@@ -203,5 +174,14 @@ def test_full_stack_filters_nothing(distributed_setup) -> None:
|
||||
|
||||
conditioner = MiniMaxH3Qwen3VLConditioner(_small_config(num_hidden_layers_override=None))
|
||||
|
||||
assert not conditioner._is_above_the_tap("language_model.layers.7.mlp.gate_proj.weight")
|
||||
assert not conditioner._is_above_the_tap("language_model.norm.weight")
|
||||
assert not conditioner._is_omitted_checkpoint_key("language_model.layers.7.mlp.gate_proj.weight")
|
||||
assert not conditioner._is_omitted_checkpoint_key("language_model.norm.weight")
|
||||
|
||||
|
||||
def test_corrupt_layer_above_checkpoint_depth_remains_unexpected(distributed_setup) -> None:
|
||||
from fastvideo.models.encoders.minimax_h3_qwen3_vl import MiniMaxH3Qwen3VLConditioner
|
||||
|
||||
conditioner = MiniMaxH3Qwen3VLConditioner(_small_config(num_hidden_layers_override=5))
|
||||
|
||||
with pytest.raises(ValueError, match="Unexpected"):
|
||||
conditioner.load_weights([("language_model.layers.8.mlp.gate_proj.weight", torch.empty(1))])
|
||||
|
||||
@@ -61,7 +61,8 @@ def test_reference_video_encode_keeps_pixels_on_cpu() -> None:
|
||||
media_type="video",
|
||||
frames=np.zeros((22, 16, 16, 3), dtype=np.uint8),
|
||||
)
|
||||
rows = stage._encode_visual_rows([reference], torch.device("cpu"))
|
||||
args = SimpleNamespace(vae_parallel_encode=False)
|
||||
rows = stage._encode_visual_rows([reference], torch.device("cpu"), args)
|
||||
|
||||
assert observed["pixels"].dtype == torch.uint8
|
||||
assert observed["pixels"].device.type == "cpu"
|
||||
@@ -96,7 +97,7 @@ def test_decode_stage_uses_cpu_output_buffer(monkeypatch) -> None:
|
||||
monkeypatch.setattr(minimax_h3_decoding, "get_local_torch_device", lambda: torch.device("cpu"))
|
||||
result = MiniMaxH3VideoDecodingStage(VAE(), SimpleNamespace(patch_size=(1, 1, 1))).forward(
|
||||
batch,
|
||||
SimpleNamespace(output_type="pil", pin_cpu_memory=False, vae_cpu_offload=False),
|
||||
SimpleNamespace(output_type="pil", pin_cpu_memory=False, vae_cpu_offload=False, vae_parallel_decode=False),
|
||||
)
|
||||
|
||||
torch.testing.assert_close(observed["latents"], latents)
|
||||
@@ -114,8 +115,9 @@ def test_decode_stages_skip_vae_on_non_output_rank(monkeypatch) -> None:
|
||||
raise AssertionError("non-output ranks must not execute a VAE")
|
||||
|
||||
monkeypatch.setattr(minimax_h3_decoding, "model_parallel_is_initialized", lambda: True)
|
||||
monkeypatch.setattr(minimax_h3_decoding, "get_world_group", lambda: SimpleNamespace(is_first_rank=False))
|
||||
args = SimpleNamespace(output_type="pil", pin_cpu_memory=False, vae_cpu_offload=True)
|
||||
monkeypatch.setattr(minimax_h3_decoding, "get_sp_group",
|
||||
lambda: SimpleNamespace(is_first_rank=False, world_size=4, rank_in_group=1))
|
||||
args = SimpleNamespace(output_type="pil", pin_cpu_memory=False, vae_cpu_offload=True, vae_parallel_decode=False)
|
||||
|
||||
video = MiniMaxH3VideoDecodingStage(VAE(), SimpleNamespace()).forward(ForwardBatch(data_type="video"), args)
|
||||
assert video.output.shape == (0, 3, 0, 0, 0)
|
||||
@@ -128,3 +130,56 @@ def test_decode_stages_skip_vae_on_non_output_rank(monkeypatch) -> None:
|
||||
assert audio.latents is None
|
||||
assert audio.audio_latents is None
|
||||
assert MINIMAX_H3_LAYOUT_KEY not in audio.extra
|
||||
|
||||
|
||||
def test_parallel_decode_runs_on_every_rank(monkeypatch) -> None:
|
||||
"""With vae_parallel_decode, non-leader ranks must enter the decode body
|
||||
(the collectives inside require uniform participation) and only the
|
||||
leader owns the CPU output buffer."""
|
||||
latent_shape = (1, 4, 2, 4, 4)
|
||||
rows = patchify_video_latents(torch.randn(latent_shape), (1, 1, 1))
|
||||
calls = []
|
||||
|
||||
class VAE:
|
||||
|
||||
def to(self, device):
|
||||
return self
|
||||
|
||||
def denormalize_latents(self, decoded_latents):
|
||||
return decoded_latents
|
||||
|
||||
def decoded_pixel_shape(self, shape):
|
||||
return (1, 3, 5, 16, 16)
|
||||
|
||||
def fake_parallel(vae, latents, output, group, strategy):
|
||||
calls.append((group.rank_in_group, output, strategy))
|
||||
if output is not None:
|
||||
output.fill_(0.5)
|
||||
return output
|
||||
|
||||
monkeypatch.setattr(minimax_h3_decoding, "get_local_torch_device", lambda: torch.device("cpu"))
|
||||
monkeypatch.setattr(minimax_h3_decoding, "model_parallel_is_initialized", lambda: True)
|
||||
monkeypatch.setattr(minimax_h3_decoding, "decode_to_pixels_parallel", fake_parallel)
|
||||
args = SimpleNamespace(output_type="pil",
|
||||
pin_cpu_memory=False,
|
||||
vae_cpu_offload=False,
|
||||
vae_parallel_decode=True,
|
||||
vae_parallel_decode_strategy="gather")
|
||||
|
||||
for rank, is_first in ((0, True), (2, False)):
|
||||
monkeypatch.setattr(
|
||||
minimax_h3_decoding, "get_sp_group",
|
||||
lambda rank=rank, is_first=is_first: SimpleNamespace(is_first_rank=is_first,
|
||||
world_size=4,
|
||||
rank_in_group=rank))
|
||||
batch = ForwardBatch(data_type="video", latents=rows.clone(), raw_latent_shape=latent_shape)
|
||||
batch.extra[MINIMAX_H3_LAYOUT_KEY] = _layout(rows.shape[0], latent_shape)
|
||||
result = MiniMaxH3VideoDecodingStage(VAE(), SimpleNamespace(patch_size=(1, 1, 1))).forward(batch, args)
|
||||
if is_first:
|
||||
assert result.output.shape == (1, 3, 5, 16, 16)
|
||||
assert torch.all(result.output == 0.5)
|
||||
else:
|
||||
assert result.output.shape == (0, 3, 0, 0, 0)
|
||||
|
||||
assert [(rank, output is not None) for rank, output, _ in calls] == [(0, True), (2, False)]
|
||||
assert all(strategy == "gather" for _, _, strategy in calls)
|
||||
|
||||
@@ -0,0 +1,231 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Focused routing and FA4 integration checks for MiniMax-H3 fusions."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("raw", "expected"),
|
||||
[
|
||||
("", frozenset()),
|
||||
("0", frozenset()),
|
||||
("none", frozenset()),
|
||||
("all", frozenset({"modulate", "qknorm_rope", "swiglu"})),
|
||||
("1", frozenset({"modulate", "qknorm_rope", "swiglu"})),
|
||||
("swiglu, modulate", frozenset({"swiglu", "modulate"})),
|
||||
],
|
||||
)
|
||||
def test_minimax_h3_fusion_selector(raw: str, expected: frozenset[str]) -> None:
|
||||
from fastvideo.models.dits.minimax_h3 import _enabled_minimax_h3_fusions
|
||||
|
||||
assert _enabled_minimax_h3_fusions(raw) == expected
|
||||
|
||||
|
||||
def test_minimax_h3_fusion_selector_rejects_unknown_name() -> None:
|
||||
from fastvideo.models.dits.minimax_h3 import _enabled_minimax_h3_fusions
|
||||
|
||||
with pytest.raises(ValueError, match="Unknown MiniMax H3 fusion"):
|
||||
_enabled_minimax_h3_fusions("swiglu,unknown")
|
||||
|
||||
|
||||
def test_swiglu_fusion_stays_on_eager_path_with_grad(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
import fastvideo.models.dits.minimax_h3 as h3
|
||||
|
||||
def unexpected_kernel(_: torch.Tensor) -> torch.Tensor:
|
||||
raise AssertionError("inference-only fusion ran with grad enabled")
|
||||
|
||||
monkeypatch.setattr(h3, "minimax_h3_swiglu", unexpected_kernel)
|
||||
layer = h3.MiniMaxH3FeedForward(8, 16, fuse_swiglu=True)
|
||||
inputs = torch.randn(2, 3, 8, requires_grad=True)
|
||||
layer(inputs).sum().backward()
|
||||
|
||||
assert inputs.grad is not None
|
||||
|
||||
|
||||
def test_all_minimax_h3_fusions_match_one_eager_block_under_fa4(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
"""Exercise the real block wiring without loading any H3 checkpoint."""
|
||||
if not torch.cuda.is_available() or not torch.cuda.is_bf16_supported():
|
||||
pytest.skip("BF16 CUDA is required")
|
||||
pytest.importorskip("triton")
|
||||
flash_attn = pytest.importorskip("flash_attn")
|
||||
if "fa4" not in getattr(flash_attn, "__version__", "").lower():
|
||||
pytest.skip("the focused integration test requires the FA4 environment")
|
||||
|
||||
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", "FLASH_ATTN")
|
||||
monkeypatch.setenv("FASTVIDEO_FA4", "1")
|
||||
monkeypatch.setenv("MASTER_ADDR", "127.0.0.1")
|
||||
monkeypatch.setenv("MASTER_PORT", "29573")
|
||||
monkeypatch.setenv("RANK", "0")
|
||||
monkeypatch.setenv("WORLD_SIZE", "1")
|
||||
monkeypatch.setenv("LOCAL_RANK", "0")
|
||||
|
||||
from fastvideo.distributed import cleanup_dist_env_and_memory, maybe_init_distributed_environment_and_model_parallel
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.models.dits.minimax_h3 import MiniMaxH3RotaryPosEmbed, MiniMaxH3TransformerBlock
|
||||
|
||||
maybe_init_distributed_environment_and_model_parallel(1, 1)
|
||||
try:
|
||||
kwargs = dict(
|
||||
hidden_size=128,
|
||||
num_attention_heads=1,
|
||||
attention_head_dim=128,
|
||||
ffn_dim=256,
|
||||
time_embed_dim=64,
|
||||
norm_eps=1e-5,
|
||||
qk_norm_eps=1e-5,
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN, ),
|
||||
quant_config=None,
|
||||
prefix="minimax_h3.test_block",
|
||||
)
|
||||
previous_default_dtype = torch.get_default_dtype()
|
||||
torch.set_default_dtype(torch.bfloat16)
|
||||
try:
|
||||
eager = MiniMaxH3TransformerBlock(**kwargs)
|
||||
fused = MiniMaxH3TransformerBlock(
|
||||
**kwargs,
|
||||
fuse_modulate=True,
|
||||
fuse_qknorm_rope=True,
|
||||
fuse_swiglu=True,
|
||||
)
|
||||
finally:
|
||||
torch.set_default_dtype(previous_default_dtype)
|
||||
|
||||
with torch.no_grad():
|
||||
for name, parameter in eager.named_parameters():
|
||||
if "norm" in name and name.endswith("weight"):
|
||||
parameter.fill_(1.0)
|
||||
elif parameter.ndim > 1:
|
||||
torch.nn.init.normal_(parameter, mean=0.0, std=0.02)
|
||||
else:
|
||||
parameter.zero_()
|
||||
fused.load_state_dict(eager.state_dict(), strict=True)
|
||||
|
||||
device = torch.device("cuda")
|
||||
eager = eager.to(device=device, dtype=torch.bfloat16).eval()
|
||||
fused = fused.to(device=device, dtype=torch.bfloat16).eval()
|
||||
generator = torch.Generator(device=device).manual_seed(2026)
|
||||
hidden_states = torch.randn(2, 12, 128, generator=generator, device=device, dtype=torch.bfloat16)
|
||||
temb = torch.randn(2, 64, generator=generator, device=device, dtype=torch.bfloat16)
|
||||
adaln_indices = torch.arange(12, device=device, dtype=torch.long).remainder(6)
|
||||
position_ids = torch.zeros(12, 3, device=device, dtype=torch.float32)
|
||||
position_ids[:, 0] = torch.arange(12, device=device)
|
||||
rotary_emb = MiniMaxH3RotaryPosEmbed(rope_freq_dim=16, rope_theta=10000.0).to(device)(position_ids)
|
||||
inputs = dict(
|
||||
hidden_states=hidden_states,
|
||||
temb=temb,
|
||||
adaln_indices=adaln_indices,
|
||||
rotary_emb=tuple(value.to(torch.bfloat16) for value in rotary_emb),
|
||||
original_seq_len=12,
|
||||
)
|
||||
|
||||
with torch.inference_mode(), set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
eager_output = eager(**inputs)
|
||||
fused_output = fused(**inputs)
|
||||
|
||||
# Sol-Engine keeps fused intermediates in FP32 registers until their
|
||||
# final BF16 stores, so the opt-in path is close but not bit-identical.
|
||||
torch.testing.assert_close(fused_output, eager_output, atol=3e-2, rtol=3e-2)
|
||||
finally:
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
|
||||
def test_minimax_h3_fusions_engage_on_cuda_inference(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""Pin the positive side of the routing guard.
|
||||
|
||||
The parity test above still passes if ``_can_run_minimax_h3_fusion``
|
||||
silently degrades to always-False (both blocks then run the identical
|
||||
eager path), so count the fused-kernel calls: one CUDA inference forward
|
||||
through a fully fused block must hit ``fused_rmsnorm_modulate`` once,
|
||||
``fused_residual_gate_rmsnorm_modulate`` once, ``fused_qknorm_rope``
|
||||
twice (q and k), and ``minimax_h3_swiglu`` once -- and a grad-enabled
|
||||
forward must leave every counter unchanged.
|
||||
"""
|
||||
if not torch.cuda.is_available() or not torch.cuda.is_bf16_supported():
|
||||
pytest.skip("BF16 CUDA is required")
|
||||
pytest.importorskip("triton")
|
||||
|
||||
monkeypatch.setenv("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
|
||||
monkeypatch.setenv("MASTER_ADDR", "127.0.0.1")
|
||||
monkeypatch.setenv("MASTER_PORT", "29574")
|
||||
monkeypatch.setenv("RANK", "0")
|
||||
monkeypatch.setenv("WORLD_SIZE", "1")
|
||||
monkeypatch.setenv("LOCAL_RANK", "0")
|
||||
|
||||
import fastvideo.models.dits.minimax_h3 as h3
|
||||
from fastvideo.distributed import cleanup_dist_env_and_memory, maybe_init_distributed_environment_and_model_parallel
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
|
||||
calls = dict.fromkeys(("rmsnorm_modulate", "residual_gate_rmsnorm_modulate", "qknorm_rope", "swiglu"), 0)
|
||||
|
||||
def _counting(name: str, real):
|
||||
|
||||
def wrapper(*args, **kwargs):
|
||||
calls[name] += 1
|
||||
return real(*args, **kwargs)
|
||||
|
||||
return wrapper
|
||||
|
||||
monkeypatch.setattr(h3, "fused_rmsnorm_modulate", _counting("rmsnorm_modulate", h3.fused_rmsnorm_modulate))
|
||||
monkeypatch.setattr(h3, "fused_residual_gate_rmsnorm_modulate",
|
||||
_counting("residual_gate_rmsnorm_modulate", h3.fused_residual_gate_rmsnorm_modulate))
|
||||
monkeypatch.setattr(h3, "fused_qknorm_rope", _counting("qknorm_rope", h3.fused_qknorm_rope))
|
||||
monkeypatch.setattr(h3, "minimax_h3_swiglu", _counting("swiglu", h3.minimax_h3_swiglu))
|
||||
|
||||
maybe_init_distributed_environment_and_model_parallel(1, 1)
|
||||
try:
|
||||
block = h3.MiniMaxH3TransformerBlock(
|
||||
hidden_size=128,
|
||||
num_attention_heads=1,
|
||||
attention_head_dim=128,
|
||||
ffn_dim=256,
|
||||
time_embed_dim=64,
|
||||
norm_eps=1e-5,
|
||||
qk_norm_eps=1e-5,
|
||||
supported_attention_backends=(AttentionBackendEnum.TORCH_SDPA, ),
|
||||
quant_config=None,
|
||||
prefix="minimax_h3.engagement_block",
|
||||
fuse_modulate=True,
|
||||
fuse_qknorm_rope=True,
|
||||
fuse_swiglu=True,
|
||||
)
|
||||
device = torch.device("cuda")
|
||||
block = block.to(device=device, dtype=torch.bfloat16).eval()
|
||||
|
||||
generator = torch.Generator(device=device).manual_seed(2026)
|
||||
hidden_states = torch.randn(2, 12, 128, generator=generator, device=device, dtype=torch.bfloat16)
|
||||
temb = torch.randn(2, 64, generator=generator, device=device, dtype=torch.bfloat16)
|
||||
adaln_indices = torch.arange(12, device=device, dtype=torch.long).remainder(6)
|
||||
position_ids = torch.zeros(12, 3, device=device, dtype=torch.float32)
|
||||
position_ids[:, 0] = torch.arange(12, device=device)
|
||||
rotary_emb = h3.MiniMaxH3RotaryPosEmbed(rope_freq_dim=16, rope_theta=10000.0).to(device)(position_ids)
|
||||
inputs = dict(
|
||||
hidden_states=hidden_states,
|
||||
temb=temb,
|
||||
adaln_indices=adaln_indices,
|
||||
rotary_emb=tuple(value.to(torch.bfloat16) for value in rotary_emb),
|
||||
original_seq_len=12,
|
||||
)
|
||||
|
||||
with torch.inference_mode(), set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
block(**inputs)
|
||||
engaged = dict(calls)
|
||||
assert engaged == {
|
||||
"rmsnorm_modulate": 1,
|
||||
"residual_gate_rmsnorm_modulate": 1,
|
||||
"qknorm_rope": 2,
|
||||
"swiglu": 1,
|
||||
}, engaged
|
||||
|
||||
grad_inputs = {**inputs, "hidden_states": hidden_states.clone().requires_grad_(True)}
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
block(**grad_inputs)
|
||||
assert dict(calls) == engaged, f"a fusion ran under grad: {calls} vs {engaged}"
|
||||
finally:
|
||||
cleanup_dist_env_and_memory()
|
||||
@@ -0,0 +1,173 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.models.dits.minimax_h3_fusions.modulation import (
|
||||
fused_residual_gate_rmsnorm_modulate,
|
||||
fused_rmsnorm_modulate,
|
||||
)
|
||||
|
||||
|
||||
EPS = 1e-6
|
||||
SOL_ENGINE_BF16_TOLERANCE = 3e-2
|
||||
|
||||
|
||||
def _chunk_tables(rows: int, hidden_size: int, *, device: torch.device | str = "cpu", dtype=torch.float32):
|
||||
wide = torch.randn(rows, 6 * hidden_size, device=device, dtype=dtype)
|
||||
tables = wide.chunk(6, dim=-1)
|
||||
assert all(table.stride() == (6 * hidden_size, 1) for table in tables)
|
||||
assert all(not table.is_contiguous() for table in tables)
|
||||
return tables
|
||||
|
||||
|
||||
def _eager_rmsnorm_modulate(x, weight, scale, shift, index):
|
||||
normed = F.rms_norm(x, (x.shape[-1], ), weight, EPS)
|
||||
return normed * (1.0 + scale.index_select(0, index)) + shift.index_select(0, index)
|
||||
|
||||
|
||||
def _eager_residual_gate_rmsnorm_modulate(residual, branch, gate, weight, scale, shift, index):
|
||||
hidden = residual + gate.index_select(0, index) * branch
|
||||
normed = F.rms_norm(hidden, (hidden.shape[-1], ), weight, EPS)
|
||||
modulated = normed * (1.0 + scale.index_select(0, index)) + shift.index_select(0, index)
|
||||
return hidden, modulated
|
||||
|
||||
|
||||
@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires a CUDA GPU")
|
||||
def test_bf16_sol_engine_fusions_match_production_eager_within_tolerance():
|
||||
"""Sol-Engine keeps fused intermediates in FP32 until its output stores."""
|
||||
pytest.importorskip("triton")
|
||||
torch.manual_seed(2)
|
||||
device = torch.device("cuda")
|
||||
batch, sequence_length, hidden_size, table_rows = 2, 9, 5376, 6
|
||||
residual = torch.randn(batch, sequence_length, hidden_size, device=device, dtype=torch.bfloat16)
|
||||
branch = torch.randn_like(residual)
|
||||
weight = torch.randn(hidden_size, device=device, dtype=torch.bfloat16)
|
||||
shift, scale, gate, shift_mlp, scale_mlp, _ = _chunk_tables(
|
||||
table_rows,
|
||||
hidden_size,
|
||||
device=device,
|
||||
dtype=torch.bfloat16,
|
||||
)
|
||||
index = torch.tensor([5, 0, 4, 1, 3, 2, 5, 1, 0], device=device, dtype=torch.int64)
|
||||
|
||||
expected_norm1 = _eager_rmsnorm_modulate(residual, weight, scale, shift, index)
|
||||
actual_norm1 = fused_rmsnorm_modulate(residual, weight, scale, shift, index, EPS)
|
||||
expected_hidden, expected_norm2 = _eager_residual_gate_rmsnorm_modulate(
|
||||
residual,
|
||||
branch,
|
||||
gate,
|
||||
weight,
|
||||
scale_mlp,
|
||||
shift_mlp,
|
||||
index,
|
||||
)
|
||||
actual_hidden, actual_norm2 = fused_residual_gate_rmsnorm_modulate(
|
||||
residual,
|
||||
branch,
|
||||
gate,
|
||||
weight,
|
||||
scale_mlp,
|
||||
shift_mlp,
|
||||
index,
|
||||
EPS,
|
||||
)
|
||||
|
||||
# Production eager materializes BF16 after each PyTorch operator. The
|
||||
# single-kernel Sol-Engine path deliberately removes those round points.
|
||||
torch.testing.assert_close(
|
||||
actual_norm1,
|
||||
expected_norm1,
|
||||
rtol=SOL_ENGINE_BF16_TOLERANCE,
|
||||
atol=SOL_ENGINE_BF16_TOLERANCE,
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
actual_hidden,
|
||||
expected_hidden,
|
||||
rtol=SOL_ENGINE_BF16_TOLERANCE,
|
||||
atol=SOL_ENGINE_BF16_TOLERANCE,
|
||||
)
|
||||
torch.testing.assert_close(
|
||||
actual_norm2,
|
||||
expected_norm2,
|
||||
rtol=SOL_ENGINE_BF16_TOLERANCE,
|
||||
atol=SOL_ENGINE_BF16_TOLERANCE,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("mutate", "error", "match"),
|
||||
[
|
||||
(lambda args: args | {"x": args["x"][0, 0]}, ValueError, "shape"),
|
||||
(lambda args: args | {"weight": args["weight"][:-1]}, ValueError, "weight"),
|
||||
(lambda args: args | {"index": args["index"][:-1]}, ValueError, "index"),
|
||||
(lambda args: args | {"index": args["index"].float()}, TypeError, "index"),
|
||||
(lambda args: args | {"eps": 0.0}, ValueError, "eps"),
|
||||
],
|
||||
)
|
||||
def test_fusion_rejects_invalid_contracts(mutate, error, match):
|
||||
args = {
|
||||
"x": torch.randn(2, 3, 8),
|
||||
"weight": torch.randn(8),
|
||||
"scale": torch.randn(4, 8),
|
||||
"shift": torch.randn(4, 8),
|
||||
"index": torch.tensor([0, 3, 1]),
|
||||
"eps": EPS,
|
||||
}
|
||||
with pytest.raises(error, match=match):
|
||||
fused_rmsnorm_modulate(**mutate(args))
|
||||
|
||||
|
||||
def test_residual_fusion_rejects_mismatched_branch():
|
||||
residual = torch.randn(2, 3, 8)
|
||||
with pytest.raises(ValueError, match="branch"):
|
||||
fused_residual_gate_rmsnorm_modulate(
|
||||
residual,
|
||||
torch.randn(2, 2, 8),
|
||||
torch.randn(4, 8),
|
||||
torch.randn(8),
|
||||
torch.randn(4, 8),
|
||||
torch.randn(4, 8),
|
||||
torch.tensor([0, 1, 2]),
|
||||
EPS,
|
||||
)
|
||||
|
||||
|
||||
def test_triton_wrappers_fail_explicitly_on_cpu():
|
||||
x = torch.randn(1, 2, 8)
|
||||
branch = torch.randn_like(x)
|
||||
weight = torch.randn(8)
|
||||
gate = torch.randn(3, 8)
|
||||
scale = torch.randn(3, 8)
|
||||
shift = torch.randn(3, 8)
|
||||
index = torch.tensor([0, 2])
|
||||
|
||||
with pytest.raises(RuntimeError, match="Triton|CUDA"):
|
||||
fused_rmsnorm_modulate(x, weight, scale, shift, index, EPS)
|
||||
with pytest.raises(RuntimeError, match="Triton|CUDA"):
|
||||
fused_residual_gate_rmsnorm_modulate(x, branch, gate, weight, scale, shift, index, EPS)
|
||||
|
||||
|
||||
def test_triton_wrappers_reject_autograd_before_backend_check():
|
||||
x = torch.randn(1, 2, 8)
|
||||
branch = torch.randn_like(x, requires_grad=True)
|
||||
weight = torch.randn(8, requires_grad=True)
|
||||
gate = torch.randn(3, 8)
|
||||
scale = torch.randn(3, 8)
|
||||
shift = torch.randn(3, 8)
|
||||
index = torch.tensor([0, 2])
|
||||
|
||||
with pytest.raises(RuntimeError, match="forward-only"):
|
||||
fused_rmsnorm_modulate(x, weight, scale, shift, index, EPS)
|
||||
with pytest.raises(RuntimeError, match="forward-only"):
|
||||
fused_residual_gate_rmsnorm_modulate(
|
||||
x,
|
||||
branch,
|
||||
gate,
|
||||
weight.detach(),
|
||||
scale,
|
||||
shift,
|
||||
index,
|
||||
EPS,
|
||||
)
|
||||
@@ -0,0 +1,192 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.models.dits.minimax_h3 import MiniMaxH3Attention
|
||||
from fastvideo.models.dits.minimax_h3_fusions.qknorm_rope import (
|
||||
HAVE_TRITON,
|
||||
fused_qknorm_rope,
|
||||
)
|
||||
|
||||
|
||||
def _rotary_tables(
|
||||
seq_len: int,
|
||||
rotary_dim: int,
|
||||
*,
|
||||
dtype: torch.dtype,
|
||||
device: torch.device | str,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
angles = torch.randn(seq_len, rotary_dim // 2, dtype=torch.float32, device=device)
|
||||
angles = torch.cat((angles, angles), dim=-1)
|
||||
return angles.cos().to(dtype), angles.sin().to(dtype)
|
||||
|
||||
|
||||
def _eager_qknorm_rope(
|
||||
x: torch.Tensor,
|
||||
weight: torch.Tensor,
|
||||
cos: torch.Tensor,
|
||||
sin: torch.Tensor,
|
||||
eps: float,
|
||||
) -> torch.Tensor:
|
||||
normalized = F.rms_norm(x, (x.shape[-1], ), weight, eps)
|
||||
return MiniMaxH3Attention._apply_rotary_emb(normalized, (cos, sin))
|
||||
|
||||
|
||||
def test_fused_qknorm_rope_rejects_invalid_rotary_dim() -> None:
|
||||
x = torch.randn(2, 3, 4, 128)
|
||||
weight = torch.ones(128)
|
||||
|
||||
cos, sin = _rotary_tables(3, 130, dtype=x.dtype, device=x.device)
|
||||
with pytest.raises(ValueError, match="rotary_dim must not exceed head_dim"):
|
||||
fused_qknorm_rope(x, weight, cos, sin, 1e-6)
|
||||
|
||||
cos = torch.randn(3, 95)
|
||||
sin = torch.randn_like(cos)
|
||||
with pytest.raises(ValueError, match="rotary_dim must be even"):
|
||||
fused_qknorm_rope(x, weight, cos, sin, 1e-6)
|
||||
|
||||
|
||||
def test_fused_qknorm_rope_rejects_shape_mismatches() -> None:
|
||||
x = torch.randn(2, 3, 4, 128)
|
||||
weight = torch.ones(128)
|
||||
cos, sin = _rotary_tables(3, 96, dtype=x.dtype, device=x.device)
|
||||
|
||||
with pytest.raises(ValueError, match="x must have shape"):
|
||||
fused_qknorm_rope(x[0], weight, cos, sin, 1e-6)
|
||||
with pytest.raises(ValueError, match="weight must have shape"):
|
||||
fused_qknorm_rope(x, weight[:-1], cos, sin, 1e-6)
|
||||
with pytest.raises(ValueError, match="sequence length"):
|
||||
fused_qknorm_rope(x, weight, cos[:-1], sin[:-1], 1e-6)
|
||||
with pytest.raises(ValueError, match="sin must match cos shape"):
|
||||
fused_qknorm_rope(x, weight, cos, sin[:, :-2], 1e-6)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("noncontiguous_input", ["weight", "cos", "sin"])
|
||||
def test_fused_qknorm_rope_rejects_noncontiguous_linear_inputs(noncontiguous_input: str) -> None:
|
||||
x = torch.randn(2, 3, 4, 128)
|
||||
weight = torch.ones(256)[::2]
|
||||
cos = torch.randn(3, 192)[:, ::2]
|
||||
sin = torch.randn(3, 192)[:, ::2]
|
||||
assert not weight.is_contiguous()
|
||||
assert not cos.is_contiguous()
|
||||
assert not sin.is_contiguous()
|
||||
|
||||
if noncontiguous_input != "weight":
|
||||
weight = weight.contiguous()
|
||||
if noncontiguous_input != "cos":
|
||||
cos = cos.contiguous()
|
||||
if noncontiguous_input != "sin":
|
||||
sin = sin.contiguous()
|
||||
|
||||
with pytest.raises(ValueError, match="must be contiguous"):
|
||||
fused_qknorm_rope(x, weight, cos, sin, 1e-6)
|
||||
|
||||
|
||||
def test_fused_qknorm_rope_requires_matching_precast_dtype() -> None:
|
||||
x = torch.randn(2, 3, 4, 128, dtype=torch.float32)
|
||||
weight = torch.ones(128, dtype=torch.float32)
|
||||
cos, sin = _rotary_tables(3, 96, dtype=torch.bfloat16, device=x.device)
|
||||
|
||||
with pytest.raises(TypeError, match="cos dtype must match x dtype"):
|
||||
fused_qknorm_rope(x, weight, cos, sin, 1e-6)
|
||||
|
||||
|
||||
def test_fused_qknorm_rope_does_not_accept_missing_rotary_tables() -> None:
|
||||
x = torch.randn(2, 3, 4, 128)
|
||||
weight = torch.ones(128)
|
||||
|
||||
with pytest.raises(TypeError, match="cos must be a torch.Tensor"):
|
||||
fused_qknorm_rope(x, weight, None, None, 1e-6) # type: ignore[arg-type]
|
||||
|
||||
|
||||
def test_fused_qknorm_rope_requires_cuda() -> None:
|
||||
x = torch.randn(2, 3, 4, 128)
|
||||
weight = torch.ones(128)
|
||||
cos, sin = _rotary_tables(3, 96, dtype=x.dtype, device=x.device)
|
||||
|
||||
with pytest.raises(RuntimeError, match="requires CUDA"):
|
||||
fused_qknorm_rope(x, weight, cos, sin, 1e-6)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"rotary_dim,shape,use_input_view",
|
||||
[
|
||||
pytest.param(96, (2, 11, 5, 128), False, id="partial-96-batch2-seq11-heads5"),
|
||||
pytest.param(128, (2, 7, 3, 128), True, id="full-128-noncontiguous-input-view"),
|
||||
],
|
||||
)
|
||||
def test_fused_qknorm_rope_matches_eager_bf16_cuda(
|
||||
rotary_dim: int,
|
||||
shape: tuple[int, ...],
|
||||
use_input_view: bool,
|
||||
) -> None:
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("CUDA is required for the Triton fusion")
|
||||
if not HAVE_TRITON:
|
||||
pytest.skip("Triton is required for the fusion")
|
||||
|
||||
torch.manual_seed(1)
|
||||
device = torch.device("cuda")
|
||||
if use_input_view:
|
||||
x = torch.randn(*shape[:-1], shape[-1] * 2, dtype=torch.bfloat16, device=device)[..., ::2]
|
||||
assert not x.is_contiguous()
|
||||
else:
|
||||
x = torch.randn(shape, dtype=torch.bfloat16, device=device)
|
||||
weight = (1.0 + 0.05 * torch.randn(shape[-1], dtype=torch.bfloat16, device=device)).contiguous()
|
||||
cos, sin = _rotary_tables(shape[1], rotary_dim, dtype=x.dtype, device=device)
|
||||
|
||||
with torch.inference_mode():
|
||||
actual = fused_qknorm_rope(x, weight, cos, sin, 1e-6)
|
||||
expected = _eager_qknorm_rope(x, weight, cos, sin, 1e-6)
|
||||
|
||||
# The fused kernel keeps RMSNorm and both RoPE products in FP32 registers
|
||||
# until its final BF16 store. Eager materializes BF16 intermediates, and
|
||||
# PyTorch/Triton reductions need not use the same summation order.
|
||||
torch.testing.assert_close(actual, expected, atol=2e-2, rtol=2e-2)
|
||||
assert actual.shape == x.shape
|
||||
assert actual.dtype == x.dtype
|
||||
assert actual.is_contiguous()
|
||||
|
||||
|
||||
def test_fused_qknorm_rope_matches_eager_beyond_int32_element_count() -> None:
|
||||
"""Regression: kernel row offsets must be int64.
|
||||
|
||||
With int32 offsets, ``row * head_dim`` wraps once the flattened input
|
||||
crosses 2**31 elements and the kernel reads/writes out of bounds (CUDA
|
||||
illegal memory access). ``(1, 8_500_000, 2, 128)`` is 2.176e9 elements,
|
||||
just past the boundary; for H3's 56 heads x 128 head_dim the equivalent
|
||||
is ``batch*seq >= 299_593`` tokens per rank, reachable at SP=1.
|
||||
|
||||
GPU assumption: needs ~16 GiB free CUDA memory (input + output at
|
||||
bf16 plus the fp32 rotary-table construction); skips below 20 GiB.
|
||||
"""
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("CUDA is required for the Triton fusion")
|
||||
if not HAVE_TRITON:
|
||||
pytest.skip("Triton is required for the fusion")
|
||||
free_bytes, _ = torch.cuda.mem_get_info()
|
||||
if free_bytes < 20 * 1024**3:
|
||||
pytest.skip("needs ~20 GiB free GPU memory for a >2**31-element input")
|
||||
|
||||
heads, head_dim, rotary_dim = 2, 128, 96
|
||||
seq_len = 8_500_000
|
||||
assert seq_len * heads * head_dim > 2**31
|
||||
|
||||
torch.manual_seed(9)
|
||||
device = torch.device("cuda")
|
||||
x = torch.randn(1, seq_len, heads, head_dim, dtype=torch.bfloat16, device=device)
|
||||
weight = (1.0 + 0.05 * torch.randn(head_dim, dtype=torch.bfloat16, device=device)).contiguous()
|
||||
cos, sin = _rotary_tables(seq_len, rotary_dim, dtype=x.dtype, device=device)
|
||||
|
||||
with torch.inference_mode():
|
||||
fused = fused_qknorm_rope(x, weight, cos, sin, 1e-6)
|
||||
torch.cuda.synchronize()
|
||||
# Compare only head/tail slices against eager: a full-tensor eager
|
||||
# reference would double peak memory for no extra coverage, and the
|
||||
# tail rows are exactly the ones an int32 wrap corrupts first.
|
||||
expected_head = _eager_qknorm_rope(x[:, :8], weight, cos[:8], sin[:8], 1e-6)
|
||||
expected_tail = _eager_qknorm_rope(x[:, -8:], weight, cos[-8:], sin[-8:], 1e-6)
|
||||
torch.testing.assert_close(fused[:, :8], expected_head, atol=2e-2, rtol=2e-2)
|
||||
torch.testing.assert_close(fused[:, -8:], expected_tail, atol=2e-2, rtol=2e-2)
|
||||
@@ -0,0 +1,73 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Focused tests for MiniMax H3's value-first packed SwiGLU fusion."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.models.dits.minimax_h3_fusions.swiglu import minimax_h3_swiglu
|
||||
|
||||
|
||||
def _require_bf16_triton_cuda() -> None:
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("MiniMax H3 fused SwiGLU requires CUDA")
|
||||
if not torch.cuda.is_bf16_supported():
|
||||
pytest.skip("MiniMax H3 fused SwiGLU parity requires BF16 support")
|
||||
pytest.importorskip("triton", reason="MiniMax H3 fused SwiGLU requires Triton")
|
||||
|
||||
|
||||
def _assert_bf16_parity(x: torch.Tensor) -> None:
|
||||
value, gate = x.chunk(2, dim=-1)
|
||||
expected = value * F.silu(gate)
|
||||
actual = minimax_h3_swiglu(x)
|
||||
assert actual.shape == (*x.shape[:-1], x.shape[-1] // 2)
|
||||
assert actual.dtype == x.dtype
|
||||
assert actual.device == x.device
|
||||
# Sol-Engine keeps the full SwiGLU expression in FP32 until the output
|
||||
# store, while eager F.silu materializes a BF16 intermediate.
|
||||
torch.testing.assert_close(actual, expected, rtol=2e-2, atol=2e-2)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("last_dim", [0, 7])
|
||||
def test_minimax_h3_swiglu_rejects_nonpositive_or_odd_last_dimension(last_dim: int) -> None:
|
||||
x = torch.empty((2, last_dim), dtype=torch.float32)
|
||||
|
||||
with pytest.raises(ValueError, match="positive even last dimension"):
|
||||
minimax_h3_swiglu(x)
|
||||
|
||||
|
||||
def test_minimax_h3_swiglu_strict_wrapper_rejects_cpu() -> None:
|
||||
with pytest.raises(ValueError, match="requires a CUDA tensor"):
|
||||
minimax_h3_swiglu(torch.randn(2, 8))
|
||||
|
||||
|
||||
@pytest.mark.gpu
|
||||
def test_minimax_h3_swiglu_multidimensional_bf16_gpu_parity() -> None:
|
||||
_require_bf16_triton_cuda()
|
||||
torch.manual_seed(0)
|
||||
x = torch.randn((2, 3, 5, 66), device="cuda", dtype=torch.bfloat16)
|
||||
|
||||
_assert_bf16_parity(x)
|
||||
|
||||
|
||||
@pytest.mark.gpu
|
||||
def test_minimax_h3_swiglu_noncontiguous_bf16_gpu_parity() -> None:
|
||||
_require_bf16_triton_cuda()
|
||||
torch.manual_seed(1)
|
||||
storage = torch.randn((2, 3, 148), device="cuda", dtype=torch.bfloat16)
|
||||
x = storage[..., ::2]
|
||||
assert not x.is_contiguous()
|
||||
|
||||
_assert_bf16_parity(x)
|
||||
|
||||
|
||||
@pytest.mark.gpu
|
||||
def test_minimax_h3_swiglu_real_ffn_dim_bf16_gpu() -> None:
|
||||
_require_bf16_triton_cuda()
|
||||
torch.manual_seed(2)
|
||||
ffn_dim = 14336
|
||||
x = torch.randn((1, 2 * ffn_dim), device="cuda", dtype=torch.bfloat16)
|
||||
|
||||
_assert_bf16_parity(x)
|
||||
@@ -0,0 +1,314 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""CPU tests for sequence-parallel MiniMax-H3 VAE chunk decode / clip encode.
|
||||
|
||||
The collective transport is simulated with a threaded fake group (one thread
|
||||
per simulated rank, barrier-synchronized slots), so the REAL drivers in
|
||||
``fastvideo.models.vaes.minimax_h3_parallel`` — chunk assignment, placeholder
|
||||
rounds, metadata broadcast, gathered-segment assembly, halo/blend math — run
|
||||
end to end on CPU and are checked bit-exactly against the serial APIs.
|
||||
"""
|
||||
|
||||
import threading
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch.testing import assert_close
|
||||
|
||||
from fastvideo.configs.models.vaes.minimax_h3_video import (
|
||||
MiniMaxH3VideoVAEArchConfig,
|
||||
MiniMaxH3VideoVAEConfig,
|
||||
)
|
||||
from fastvideo.models.vaes.minimax_h3_parallel import (
|
||||
DECODE_GATHER_STRATEGIES,
|
||||
DEFAULT_DECODE_GATHER_STRATEGY,
|
||||
decode_to_pixels_parallel,
|
||||
encode_pixels_parallel,
|
||||
parallel_chunk_indices,
|
||||
)
|
||||
from fastvideo.models.vaes.minimax_h3_video import AutoencoderKLMiniMaxH3
|
||||
|
||||
|
||||
def _tiny_vae(token_drop: int = 3) -> AutoencoderKLMiniMaxH3:
|
||||
arch = MiniMaxH3VideoVAEArchConfig(
|
||||
latent_channels=4,
|
||||
block_out_channels=(32, 32),
|
||||
layers_per_block=1,
|
||||
spatial_downsample_factors=(2, 2),
|
||||
temporal_downsample_factors=(2, 2),
|
||||
decoder_num_layers=1,
|
||||
decoder_num_attention_heads=1,
|
||||
decoder_attention_head_dim=8,
|
||||
decoder_num_register_tokens=2,
|
||||
decoder_ffn_mult=1,
|
||||
token_drop=token_drop,
|
||||
latents_mean=(0.0, ) * 4,
|
||||
latents_std=(1.0, ) * 4,
|
||||
)
|
||||
return AutoencoderKLMiniMaxH3(
|
||||
MiniMaxH3VideoVAEConfig(
|
||||
arch_config=arch,
|
||||
use_tiling=False,
|
||||
use_temporal_tiling=False,
|
||||
use_parallel_tiling=False,
|
||||
)).eval()
|
||||
|
||||
|
||||
class _ThreadedFakeGroup:
|
||||
"""Barrier-synchronized in-process stand-in for a GroupCoordinator.
|
||||
|
||||
One thread per simulated rank runs the SPMD driver; ``gather`` /
|
||||
``all_gather`` / ``broadcast_object`` rendezvous through shared slots
|
||||
with a double barrier (all writes land, everyone reads, then slots are
|
||||
reusable). Matches the GroupCoordinator call signatures the drivers use.
|
||||
"""
|
||||
|
||||
def __init__(self, world_size: int) -> None:
|
||||
self.world_size = world_size
|
||||
self._local = threading.local()
|
||||
self._barrier = threading.Barrier(world_size)
|
||||
self._slots: list = [None] * world_size
|
||||
self._object = None
|
||||
|
||||
@property
|
||||
def rank_in_group(self) -> int:
|
||||
return self._local.rank
|
||||
|
||||
def broadcast_object(self, obj=None, src: int = 0):
|
||||
if self.world_size == 1:
|
||||
return obj
|
||||
if self.rank_in_group == src:
|
||||
self._object = obj
|
||||
self._barrier.wait()
|
||||
received = self._object
|
||||
self._barrier.wait()
|
||||
return received
|
||||
|
||||
def all_gather(self, input_: torch.Tensor, dim: int = -1) -> torch.Tensor:
|
||||
if self.world_size == 1:
|
||||
return input_
|
||||
self._slots[self.rank_in_group] = input_
|
||||
self._barrier.wait()
|
||||
gathered = torch.cat([slot for slot in self._slots], dim=dim)
|
||||
self._barrier.wait()
|
||||
return gathered
|
||||
|
||||
def gather(self, input_: torch.Tensor, dst: int = 0, dim: int = -1):
|
||||
if self.world_size == 1:
|
||||
return input_
|
||||
self._slots[self.rank_in_group] = input_
|
||||
self._barrier.wait()
|
||||
gathered = torch.cat([slot for slot in self._slots], dim=dim) if self.rank_in_group == dst else None
|
||||
self._barrier.wait()
|
||||
return gathered
|
||||
|
||||
def run(self, fn) -> list:
|
||||
"""Run ``fn(rank)`` on one thread per rank; re-raise the first error."""
|
||||
results: list = [None] * self.world_size
|
||||
errors: list = [None] * self.world_size
|
||||
|
||||
def _target(rank: int) -> None:
|
||||
self._local.rank = rank
|
||||
try:
|
||||
# inference_mode is thread-local; the drivers run inference-only.
|
||||
with torch.inference_mode():
|
||||
results[rank] = fn(rank)
|
||||
except BaseException as error: # noqa: BLE001 - propagate to the test
|
||||
errors[rank] = error
|
||||
self._barrier.abort()
|
||||
|
||||
threads = [threading.Thread(target=_target, args=(rank, )) for rank in range(self.world_size)]
|
||||
for thread in threads:
|
||||
thread.start()
|
||||
for thread in threads:
|
||||
thread.join()
|
||||
for error in errors:
|
||||
if error is not None and not isinstance(error, threading.BrokenBarrierError):
|
||||
raise error
|
||||
for error in errors:
|
||||
if error is not None:
|
||||
raise error
|
||||
return results
|
||||
|
||||
|
||||
@pytest.mark.parametrize("num_chunks,world_size", ((0, 4), (1, 4), (7, 4), (8, 4), (20, 4), (5, 3), (2, 5)))
|
||||
def test_parallel_chunk_indices_partition(num_chunks: int, world_size: int) -> None:
|
||||
"""Round-robin ownership covers every chunk exactly once, in order."""
|
||||
owned = [parallel_chunk_indices(num_chunks, world_size, rank) for rank in range(world_size)]
|
||||
flattened = sorted(index for indices in owned for index in indices)
|
||||
assert flattened == list(range(num_chunks))
|
||||
for rank, indices in enumerate(owned):
|
||||
assert indices == sorted(indices)
|
||||
assert all(index % world_size == rank for index in indices)
|
||||
# Round-robin balance: no rank holds more than one extra chunk.
|
||||
assert len(indices) in (num_chunks // world_size, -(-num_chunks // world_size))
|
||||
|
||||
|
||||
def test_parallel_chunk_indices_validates() -> None:
|
||||
with pytest.raises(ValueError, match="world_size"):
|
||||
parallel_chunk_indices(4, 0, 0)
|
||||
with pytest.raises(ValueError, match="rank_in_group"):
|
||||
parallel_chunk_indices(4, 2, 2)
|
||||
with pytest.raises(ValueError, match="num_chunks"):
|
||||
parallel_chunk_indices(-1, 2, 0)
|
||||
|
||||
|
||||
# Latent frames cover: one padded chunk (3), pad on the intra-clip tail (6),
|
||||
# two blended chunks (12), three chunks plus pad trim (13). World sizes cover
|
||||
# fewer chunks than ranks, uneven rounds, and the exact-multiple case.
|
||||
@pytest.mark.parametrize("world_size", (2, 3, 4, 5))
|
||||
@pytest.mark.parametrize("latent_frames", (3, 6, 12, 13))
|
||||
@pytest.mark.parametrize("strategy", DECODE_GATHER_STRATEGIES)
|
||||
@torch.inference_mode()
|
||||
def test_parallel_decode_matches_serial(world_size: int, latent_frames: int, strategy: str) -> None:
|
||||
torch.manual_seed(20260821 + latent_frames)
|
||||
vae = _tiny_vae()
|
||||
latents = torch.randn(1, 4, latent_frames, 4, 4)
|
||||
expected = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32)
|
||||
vae.decode_to_pixels(latents, expected)
|
||||
|
||||
group = _ThreadedFakeGroup(world_size)
|
||||
|
||||
def _rank_main(rank: int):
|
||||
output = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32) if rank == 0 else None
|
||||
return decode_to_pixels_parallel(vae, latents.clone(), output, group, strategy=strategy)
|
||||
|
||||
results = group.run(_rank_main)
|
||||
assert all(result is None for result in results[1:])
|
||||
assert_close(results[0], expected, atol=0.0, rtol=0.0)
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def test_parallel_decode_without_token_drop() -> None:
|
||||
"""token_drop == 0 has no overlap halo; the assembler must skip blending."""
|
||||
torch.manual_seed(20260822)
|
||||
vae = _tiny_vae(token_drop=0)
|
||||
assert vae.frame_overlap == 0
|
||||
latents = torch.randn(1, 4, 10, 4, 4)
|
||||
expected = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32)
|
||||
vae.decode_to_pixels(latents, expected)
|
||||
|
||||
group = _ThreadedFakeGroup(3)
|
||||
|
||||
def _rank_main(rank: int):
|
||||
output = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32) if rank == 0 else None
|
||||
return decode_to_pixels_parallel(vae, latents.clone(), output, group)
|
||||
|
||||
results = group.run(_rank_main)
|
||||
assert_close(results[0], expected, atol=0.0, rtol=0.0)
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def test_parallel_decode_batched_slicing_matches_serial() -> None:
|
||||
torch.manual_seed(20260823)
|
||||
vae = _tiny_vae()
|
||||
vae.enable_slicing()
|
||||
latents = torch.randn(2, 4, 7, 4, 4)
|
||||
expected = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32)
|
||||
vae.decode_to_pixels(latents, expected)
|
||||
|
||||
group = _ThreadedFakeGroup(2)
|
||||
|
||||
def _rank_main(rank: int):
|
||||
output = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32) if rank == 0 else None
|
||||
return decode_to_pixels_parallel(vae, latents.clone(), output, group)
|
||||
|
||||
results = group.run(_rank_main)
|
||||
assert_close(results[0], expected, atol=0.0, rtol=0.0)
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def test_parallel_decode_world_size_one_is_serial() -> None:
|
||||
torch.manual_seed(20260824)
|
||||
vae = _tiny_vae()
|
||||
latents = torch.randn(1, 4, 7, 4, 4)
|
||||
expected = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32)
|
||||
vae.decode_to_pixels(latents, expected)
|
||||
|
||||
group = _ThreadedFakeGroup(1)
|
||||
|
||||
def _rank_main(rank: int):
|
||||
output = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32)
|
||||
return decode_to_pixels_parallel(vae, latents, output, group)
|
||||
|
||||
results = group.run(_rank_main)
|
||||
assert_close(results[0], expected, atol=0.0, rtol=0.0)
|
||||
|
||||
|
||||
def test_parallel_decode_validates_buffers_and_strategy() -> None:
|
||||
vae = _tiny_vae()
|
||||
latents = torch.randn(1, 4, 7, 4, 4)
|
||||
output = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32)
|
||||
group = _ThreadedFakeGroup(1)
|
||||
group._local.rank = 0
|
||||
|
||||
with pytest.raises(ValueError, match="strategy"):
|
||||
decode_to_pixels_parallel(vae, latents, output, group, strategy="scatter")
|
||||
with pytest.raises(ValueError, match="must provide the CPU output buffer"):
|
||||
decode_to_pixels_parallel(vae, latents, None, group)
|
||||
with pytest.raises(ValueError, match="CPU float32 tensor"):
|
||||
decode_to_pixels_parallel(vae, latents, output[:, :, :-1], group)
|
||||
|
||||
group._local.rank = 1 # simulate a non-leader passing a buffer
|
||||
group.world_size = 2
|
||||
with pytest.raises(ValueError, match="Only the first sequence-parallel rank"):
|
||||
decode_to_pixels_parallel(vae, latents, output, group)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("world_size", (2, 4))
|
||||
@pytest.mark.parametrize("num_frames", (16, 22, 40))
|
||||
@torch.inference_mode()
|
||||
def test_parallel_encode_matches_serial(world_size: int, num_frames: int) -> None:
|
||||
"""Every rank must hold the full serial moments, bit for bit."""
|
||||
torch.manual_seed(20260825 + num_frames)
|
||||
vae = _tiny_vae()
|
||||
pixels = torch.randint(0, 256, (1, 3, num_frames, 16, 16), dtype=torch.uint8)
|
||||
expected = vae.encode_pixels(pixels).latent_dist.parameters
|
||||
|
||||
group = _ThreadedFakeGroup(world_size)
|
||||
results = group.run(lambda rank: encode_pixels_parallel(vae, pixels, group).latent_dist.parameters)
|
||||
for moments in results:
|
||||
assert_close(moments, expected, atol=0.0, rtol=0.0)
|
||||
|
||||
|
||||
@torch.inference_mode()
|
||||
def test_parallel_encode_float_and_batched_slicing() -> None:
|
||||
torch.manual_seed(20260826)
|
||||
vae = _tiny_vae()
|
||||
vae.enable_slicing()
|
||||
pixels = torch.rand(2, 3, 22, 16, 16)
|
||||
expected = vae.encode_pixels(pixels).latent_dist.parameters
|
||||
|
||||
group = _ThreadedFakeGroup(3)
|
||||
results = group.run(lambda rank: encode_pixels_parallel(vae, pixels, group).latent_dist.parameters)
|
||||
for moments in results:
|
||||
assert_close(moments, expected, atol=0.0, rtol=0.0)
|
||||
|
||||
|
||||
def test_parallel_encode_validates_input() -> None:
|
||||
vae = _tiny_vae()
|
||||
group = _ThreadedFakeGroup(1)
|
||||
group._local.rank = 0
|
||||
with pytest.raises(ValueError, match="must remain on CPU"):
|
||||
encode_pixels_parallel(vae, torch.empty(1, 3, 4, 16, 16, device="meta"), group)
|
||||
with pytest.raises(TypeError, match="uint8 or a floating-point"):
|
||||
encode_pixels_parallel(vae, torch.zeros(1, 3, 4, 16, 16, dtype=torch.int32), group)
|
||||
with pytest.raises(ValueError, match="must have shape"):
|
||||
encode_pixels_parallel(vae, torch.zeros(1, 4, 4, 16, 16), group)
|
||||
|
||||
|
||||
def test_fastvideo_args_strategy_literals_match_module() -> None:
|
||||
"""fastvideo_args mirrors the strategy literals to avoid importing model
|
||||
modules at args construction; keep the two in sync."""
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
|
||||
args = FastVideoArgs(model_path="test/parallel-vae")
|
||||
assert args.vae_parallel_decode is False
|
||||
assert args.vae_parallel_encode is False
|
||||
assert args.vae_parallel_decode_strategy == DEFAULT_DECODE_GATHER_STRATEGY
|
||||
assert args.vae_parallel_decode_strategy in DECODE_GATHER_STRATEGIES
|
||||
|
||||
for strategy in DECODE_GATHER_STRATEGIES:
|
||||
assert FastVideoArgs(model_path="test/parallel-vae",
|
||||
vae_parallel_decode_strategy=strategy).vae_parallel_decode_strategy == strategy
|
||||
with pytest.raises(ValueError, match="vae_parallel_decode_strategy"):
|
||||
FastVideoArgs(model_path="test/parallel-vae", vae_parallel_decode_strategy="scatter")
|
||||
@@ -0,0 +1,110 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""GPU regression test for sequence-parallel MiniMax-H3 VAE decode/encode.
|
||||
|
||||
Requires a multi-GPU torchrun launch (real NCCL collectives across an SP
|
||||
group); skipped otherwise:
|
||||
|
||||
torchrun --nproc-per-node=4 -m pytest \
|
||||
fastvideo/tests/vaes/test_minimax_h3_parallel_vae_gpu.py -q
|
||||
|
||||
Asserts the parallel drivers are bitwise equal to the serial rank-local
|
||||
decode/encode under the pipeline's fp16 autocast, for both transport
|
||||
strategies, and that repeated parallel runs are deterministic.
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch.testing import assert_close
|
||||
|
||||
from fastvideo.configs.models.vaes.minimax_h3_video import (
|
||||
MiniMaxH3VideoVAEArchConfig,
|
||||
MiniMaxH3VideoVAEConfig,
|
||||
)
|
||||
from fastvideo.models.vaes.minimax_h3_video import AutoencoderKLMiniMaxH3
|
||||
|
||||
_WORLD_SIZE = int(os.environ.get("WORLD_SIZE", "1"))
|
||||
|
||||
|
||||
def _tiny_vae() -> AutoencoderKLMiniMaxH3:
|
||||
"""Same tiny geometry as test_minimax_h3_parallel_vae (test dirs are not packages)."""
|
||||
arch = MiniMaxH3VideoVAEArchConfig(
|
||||
latent_channels=4,
|
||||
block_out_channels=(32, 32),
|
||||
layers_per_block=1,
|
||||
spatial_downsample_factors=(2, 2),
|
||||
temporal_downsample_factors=(2, 2),
|
||||
decoder_num_layers=1,
|
||||
decoder_num_attention_heads=1,
|
||||
decoder_attention_head_dim=8,
|
||||
decoder_num_register_tokens=2,
|
||||
decoder_ffn_mult=1,
|
||||
latents_mean=(0.0, ) * 4,
|
||||
latents_std=(1.0, ) * 4,
|
||||
)
|
||||
return AutoencoderKLMiniMaxH3(
|
||||
MiniMaxH3VideoVAEConfig(
|
||||
arch_config=arch,
|
||||
use_tiling=False,
|
||||
use_temporal_tiling=False,
|
||||
use_parallel_tiling=False,
|
||||
)).eval()
|
||||
|
||||
pytestmark = [
|
||||
pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA"),
|
||||
pytest.mark.skipif(_WORLD_SIZE < 2, reason="requires a torchrun launch with WORLD_SIZE > 1"),
|
||||
]
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def sp_group():
|
||||
from fastvideo.distributed import get_sp_group, maybe_init_distributed_environment_and_model_parallel
|
||||
maybe_init_distributed_environment_and_model_parallel(1, _WORLD_SIZE)
|
||||
return get_sp_group()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("strategy", ("gather", "all_gather"))
|
||||
@pytest.mark.parametrize("latent_frames", (3, 13))
|
||||
@torch.no_grad()
|
||||
def test_parallel_decode_bitwise_matches_serial_on_gpu(sp_group, strategy: str, latent_frames: int) -> None:
|
||||
from fastvideo.models.vaes.minimax_h3_parallel import decode_to_pixels_parallel
|
||||
|
||||
device = torch.device("cuda", torch.cuda.current_device())
|
||||
torch.manual_seed(20260821) # identical weights on every rank
|
||||
vae = _tiny_vae().to(device)
|
||||
latents = torch.randn(1, 4, latent_frames, 4, 4, generator=torch.Generator().manual_seed(7)).to(device)
|
||||
|
||||
with torch.autocast(device_type="cuda", dtype=torch.float16):
|
||||
expected = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32, pin_memory=True)
|
||||
vae.decode_to_pixels(latents, expected)
|
||||
|
||||
outputs = []
|
||||
for _ in range(3): # repeat-determinism
|
||||
output = None
|
||||
if sp_group.is_first_rank:
|
||||
output = torch.empty(vae.decoded_pixel_shape(latents.shape), dtype=torch.float32, pin_memory=True)
|
||||
result = decode_to_pixels_parallel(vae, latents, output, sp_group, strategy=strategy)
|
||||
outputs.append(result.clone() if result is not None else None)
|
||||
|
||||
if sp_group.is_first_rank:
|
||||
for output in outputs:
|
||||
assert_close(output, expected, atol=0.0, rtol=0.0)
|
||||
else:
|
||||
assert all(output is None for output in outputs)
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def test_parallel_encode_bitwise_matches_serial_on_gpu(sp_group) -> None:
|
||||
from fastvideo.models.vaes.minimax_h3_parallel import encode_pixels_parallel
|
||||
|
||||
device = torch.device("cuda", torch.cuda.current_device())
|
||||
torch.manual_seed(20260821)
|
||||
vae = _tiny_vae().to(device)
|
||||
pixels = torch.randint(0, 256, (1, 3, 40, 16, 16), dtype=torch.uint8,
|
||||
generator=torch.Generator().manual_seed(9))
|
||||
|
||||
expected = vae.encode_pixels(pixels).latent_dist.parameters
|
||||
for _ in range(3):
|
||||
moments = encode_pixels_parallel(vae, pixels, sp_group).latent_dist.parameters
|
||||
assert_close(moments, expected, atol=0.0, rtol=0.0)
|
||||
@@ -0,0 +1,428 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""CPU contract tests for MiniMax H3 VAE compilation and profiling ranges.
|
||||
|
||||
The final test is a CUDA regression gate for the reduce-overhead tile path
|
||||
(real tiled decode/encode with an unmocked ``_stitch_tiles``).
|
||||
"""
|
||||
|
||||
from contextlib import contextmanager
|
||||
from types import MethodType, SimpleNamespace
|
||||
from typing import Any
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.models.vaes.minimax_h3_audio import (
|
||||
MiniMaxH3AudioBigVGANDecoder,
|
||||
MiniMaxH3AudioVAE,
|
||||
)
|
||||
from fastvideo.models.vaes.minimax_h3_video import (
|
||||
AutoencoderKLMiniMaxH3,
|
||||
MiniMaxH3VideoAttention,
|
||||
MiniMaxH3VideoViTDecoder3d,
|
||||
)
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
|
||||
|
||||
def _empty_typed_module(module_type: type[nn.Module]) -> nn.Module:
|
||||
"""Create a weightless instance that retains its production module type."""
|
||||
module = object.__new__(module_type)
|
||||
nn.Module.__init__(module)
|
||||
return module
|
||||
|
||||
|
||||
def _assert_dynamic_compile_selects_decoder(
|
||||
vae_type: type[nn.Module],
|
||||
decoder_type: type[nn.Module],
|
||||
) -> None:
|
||||
"""Verify one H3 VAE compiles only its top-level decoder in place."""
|
||||
vae = _empty_typed_module(vae_type)
|
||||
decoder = _empty_typed_module(decoder_type)
|
||||
same_type_under_another_name = _empty_typed_module(decoder_type)
|
||||
unrelated_submodule = nn.Identity()
|
||||
vae.decoder = decoder
|
||||
vae.same_type_under_another_name = same_type_under_another_name
|
||||
vae.unrelated_submodule = unrelated_submodule
|
||||
|
||||
compiled_forward = Mock(name="compiled_forward")
|
||||
compile_kwargs = {"backend": "inductor", "dynamic": False}
|
||||
with patch(
|
||||
"fastvideo.pipelines.composed_pipeline_base.torch.compile",
|
||||
return_value=compiled_forward,
|
||||
) as compile_mock:
|
||||
compiled_count = ComposedPipelineBase._compile_with_conditions(vae, compile_kwargs)
|
||||
|
||||
assert compiled_count == 1
|
||||
compile_mock.assert_called_once()
|
||||
selected_forward = compile_mock.call_args.args[0]
|
||||
assert selected_forward.__self__ is decoder
|
||||
assert selected_forward.__func__ is decoder_type.forward
|
||||
assert compile_mock.call_args.kwargs == compile_kwargs
|
||||
assert decoder.forward is compiled_forward
|
||||
assert "forward" not in same_type_under_another_name.__dict__
|
||||
assert "forward" not in unrelated_submodule.__dict__
|
||||
|
||||
wrong_type_vae = _empty_typed_module(vae_type)
|
||||
wrong_type_vae.decoder = nn.Identity()
|
||||
with patch("fastvideo.pipelines.composed_pipeline_base.torch.compile") as wrong_type_compile:
|
||||
wrong_type_count = ComposedPipelineBase._compile_with_conditions(wrong_type_vae, compile_kwargs)
|
||||
assert wrong_type_count == 0
|
||||
wrong_type_compile.assert_not_called()
|
||||
|
||||
|
||||
def _assert_reduce_overhead_compile(compiled_function: Any) -> None:
|
||||
"""Verify a class-owned compile boundary enables CUDA Graph replay."""
|
||||
assert hasattr(compiled_function, "get_compiler_config")
|
||||
assert compiled_function.get_compiler_config()["triton.cudagraphs"] is True
|
||||
|
||||
|
||||
def test_video_attention_uses_selected_fastvideo_backend() -> None:
|
||||
"""Pass BSHD tensors to the selected dense backend without forward metadata."""
|
||||
backend_call: dict[str, Any] = {}
|
||||
|
||||
class RecordingAttentionImpl:
|
||||
"""Record the backend construction and forward contracts."""
|
||||
|
||||
def __init__(self, **kwargs: Any) -> None:
|
||||
backend_call["init"] = kwargs
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attention_metadata: Any,
|
||||
) -> torch.Tensor:
|
||||
"""Return values unchanged after recording the backend inputs."""
|
||||
backend_call["shapes"] = (query.shape, key.shape, value.shape)
|
||||
backend_call["metadata"] = attention_metadata
|
||||
return value
|
||||
|
||||
class RecordingAttentionBackend:
|
||||
"""Supply the recording implementation through the backend API."""
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> type[RecordingAttentionImpl]:
|
||||
return RecordingAttentionImpl
|
||||
|
||||
with (
|
||||
patch("fastvideo.platforms.current_platform") as current_platform,
|
||||
patch(
|
||||
"fastvideo.models.vaes.minimax_h3_video.get_attn_backend",
|
||||
return_value=RecordingAttentionBackend,
|
||||
) as get_backend,
|
||||
):
|
||||
current_platform.is_cuda_alike.return_value = True
|
||||
attention = MiniMaxH3VideoAttention(dim=8, heads=2, dim_head=4)
|
||||
|
||||
attention.to_q = nn.Identity()
|
||||
attention.to_k = nn.Identity()
|
||||
attention.to_v = nn.Identity()
|
||||
attention.norm_q = nn.Identity()
|
||||
attention.norm_k = nn.Identity()
|
||||
attention.to_out[0] = nn.Identity()
|
||||
output = attention(torch.empty((1, 3, 8), device="meta"))
|
||||
|
||||
get_backend.assert_called_once_with(
|
||||
4,
|
||||
torch.bfloat16,
|
||||
supported_attention_backends=(
|
||||
AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
),
|
||||
)
|
||||
assert backend_call["init"] == {
|
||||
"num_heads": 2,
|
||||
"head_size": 4,
|
||||
"softmax_scale": 0.5,
|
||||
"num_kv_heads": 2,
|
||||
"causal": False,
|
||||
}
|
||||
assert backend_call["shapes"] == ((1, 3, 2, 4), ) * 3
|
||||
assert backend_call["metadata"] is None
|
||||
assert output.shape == (1, 3, 8)
|
||||
|
||||
|
||||
def test_video_attention_cpu_uses_torch_sdpa() -> None:
|
||||
"""Use PyTorch SDPA when H3 VAE attention receives CPU tensors."""
|
||||
with (
|
||||
patch("fastvideo.platforms.current_platform") as current_platform,
|
||||
patch("fastvideo.models.vaes.minimax_h3_video.get_attn_backend") as get_backend,
|
||||
):
|
||||
current_platform.is_cuda_alike.return_value = False
|
||||
attention = MiniMaxH3VideoAttention(dim=8, heads=2, dim_head=4)
|
||||
|
||||
get_backend.assert_not_called()
|
||||
assert attention.attn_impl is None
|
||||
|
||||
attention.to_q = nn.Identity()
|
||||
attention.to_k = nn.Identity()
|
||||
attention.to_v = nn.Identity()
|
||||
attention.norm_q = nn.Identity()
|
||||
attention.norm_k = nn.Identity()
|
||||
attention.to_out[0] = nn.Identity()
|
||||
hidden_states = torch.randn(1, 3, 8)
|
||||
query = hidden_states.unflatten(2, (2, 4)).permute(0, 2, 1, 3)
|
||||
expected = F.scaled_dot_product_attention(query, query, query).permute(0, 2, 1, 3).flatten(2, 3)
|
||||
|
||||
torch.testing.assert_close(attention(hidden_states), expected)
|
||||
|
||||
|
||||
def test_compile_with_conditions_selects_minimax_h3_video_decoder() -> None:
|
||||
"""Compile the registered video decoder with the VAE runtime kwargs."""
|
||||
assert not hasattr(MiniMaxH3VideoViTDecoder3d.forward, "get_compiler_config")
|
||||
_assert_dynamic_compile_selects_decoder(AutoencoderKLMiniMaxH3, MiniMaxH3VideoViTDecoder3d)
|
||||
|
||||
|
||||
def test_project_decoder_tile_uses_reduce_overhead_compile() -> None:
|
||||
"""Compile the per-tile decoder-input projection with CUDA Graph replay."""
|
||||
_assert_reduce_overhead_compile(AutoencoderKLMiniMaxH3._project_decoder_tile)
|
||||
|
||||
|
||||
def test_stitch_tiles_uses_reduce_overhead_compile() -> None:
|
||||
"""Compile spatial tile blending and concatenation with CUDA Graph replay."""
|
||||
_assert_reduce_overhead_compile(AutoencoderKLMiniMaxH3._stitch_tiles)
|
||||
|
||||
|
||||
def test_compile_with_conditions_selects_minimax_h3_audio_decoder() -> None:
|
||||
"""Compile the audio VAE decoder that the H3 waveform decode path calls."""
|
||||
_assert_dynamic_compile_selects_decoder(MiniMaxH3AudioVAE, MiniMaxH3AudioBigVGANDecoder)
|
||||
|
||||
|
||||
def test_decode_emits_indexed_temporal_chunk_ranges() -> None:
|
||||
"""Nest frame-segment ranges under each temporal decoder chunk range."""
|
||||
vae = _empty_typed_module(AutoencoderKLMiniMaxH3)
|
||||
vae.tokens_chunk_size = 1
|
||||
vae.token_overlap = 1
|
||||
vae.temporal_compression_ratio = 1
|
||||
vae.frame_pre_padding = 0
|
||||
vae.frame_overlap = 1
|
||||
vae.config = SimpleNamespace(token_drop=1)
|
||||
vae._decode_clip = Mock(return_value=torch.zeros((1, 1, 2, 1, 1)))
|
||||
range_events = []
|
||||
|
||||
@contextmanager
|
||||
def record_range(name: str):
|
||||
range_events.append(("enter", name))
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
range_events.append(("exit", name))
|
||||
|
||||
with patch("fastvideo.models.vaes.minimax_h3_video.nvtx_range", record_range):
|
||||
decoded = vae._decode(torch.zeros((1, 1, 2, 1, 1)))
|
||||
|
||||
assert decoded.shape == (1, 1, 3, 1, 1)
|
||||
assert vae._decode_clip.call_count == 2
|
||||
assert range_events == [
|
||||
("enter", "minimax_h3.vae.temporal_chunk.0"),
|
||||
("enter", "minimax_h3.vae.temporal_chunk.0.frame_segment.0"),
|
||||
("exit", "minimax_h3.vae.temporal_chunk.0.frame_segment.0"),
|
||||
("enter", "minimax_h3.vae.temporal_chunk.0.frame_segment.1"),
|
||||
("exit", "minimax_h3.vae.temporal_chunk.0.frame_segment.1"),
|
||||
("exit", "minimax_h3.vae.temporal_chunk.0"),
|
||||
("enter", "minimax_h3.vae.temporal_chunk.1"),
|
||||
("enter", "minimax_h3.vae.temporal_chunk.1.frame_segment.0"),
|
||||
("exit", "minimax_h3.vae.temporal_chunk.1.frame_segment.0"),
|
||||
("enter", "minimax_h3.vae.temporal_chunk.1.frame_segment.1"),
|
||||
("exit", "minimax_h3.vae.temporal_chunk.1.frame_segment.1"),
|
||||
("exit", "minimax_h3.vae.temporal_chunk.1"),
|
||||
]
|
||||
|
||||
|
||||
def test_decode_clip_no_spatial_tiling_stage_ranges() -> None:
|
||||
"""Separate untiled latent projection and decoder ranges."""
|
||||
vae = _empty_typed_module(AutoencoderKLMiniMaxH3)
|
||||
vae.use_tiling = False
|
||||
range_events = []
|
||||
vae.post_quant_conv = nn.Identity()
|
||||
vae.post_quant_conv.register_forward_hook(
|
||||
lambda _module, _args, _output: range_events.append(("call", "post_quant_conv")))
|
||||
vae.decoder = nn.Identity()
|
||||
vae.decoder.register_forward_hook(
|
||||
lambda _module, _args, _output: range_events.append(("call", "decoder_forward")))
|
||||
latent_clip = torch.zeros((1, 1, 1, 2, 2))
|
||||
|
||||
@contextmanager
|
||||
def record_range(name: str):
|
||||
range_events.append(("enter", name))
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
range_events.append(("exit", name))
|
||||
|
||||
with patch("fastvideo.models.vaes.minimax_h3_video.nvtx_range", record_range):
|
||||
decoded_clip = vae._decode_clip(latent_clip)
|
||||
|
||||
assert decoded_clip is latent_clip
|
||||
assert range_events == [
|
||||
("enter", "minimax_h3.vae.decode_clip"),
|
||||
("enter", "minimax_h3.vae.decode_clip.no_s_tile.post_quant_conv"),
|
||||
("call", "post_quant_conv"),
|
||||
("exit", "minimax_h3.vae.decode_clip.no_s_tile.post_quant_conv"),
|
||||
("enter", "minimax_h3.vae.decode_clip.no_s_tile.decoder_forward"),
|
||||
("call", "decoder_forward"),
|
||||
("exit", "minimax_h3.vae.decode_clip.no_s_tile.decoder_forward"),
|
||||
("exit", "minimax_h3.vae.decode_clip"),
|
||||
]
|
||||
|
||||
|
||||
def test_decode_clip_emits_tiled_stage_ranges() -> None:
|
||||
"""Nest indexed decoder tiles between tile-splitting and stitching ranges."""
|
||||
vae = _empty_typed_module(AutoencoderKLMiniMaxH3)
|
||||
vae.use_tiling = True
|
||||
vae.spatial_compression_ratio = 1
|
||||
vae.tile_sample_min_height = 1
|
||||
vae.tile_sample_min_width = 1
|
||||
vae.tile_sample_min_overlap_height = 0
|
||||
vae.tile_sample_min_overlap_width = 0
|
||||
vae._split_tiles = Mock(side_effect=[
|
||||
([0, 1], [1, 1], [0]),
|
||||
([0, 1], [1, 1], [0]),
|
||||
])
|
||||
vae.post_quant_conv = nn.Identity()
|
||||
vae._project_decoder_tile = Mock(side_effect=vae.post_quant_conv)
|
||||
vae.decoder = nn.Identity()
|
||||
stitched_clip = torch.zeros((1, 1, 1, 2, 2))
|
||||
vae._stitch_tiles = Mock(return_value=stitched_clip)
|
||||
range_events = []
|
||||
|
||||
@contextmanager
|
||||
def record_range(name: str):
|
||||
range_events.append(("enter", name))
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
range_events.append(("exit", name))
|
||||
|
||||
with patch("fastvideo.models.vaes.minimax_h3_video.nvtx_range", record_range):
|
||||
decoded_clip = vae._decode_clip(torch.zeros((1, 1, 1, 2, 2)))
|
||||
|
||||
# The tile driver must hand back a caller-owned copy: under
|
||||
# mode="reduce-overhead" the stitched canvas is CUDA-graph pooled storage
|
||||
# that the next replay overwrites, so returning it by identity is a bug.
|
||||
assert decoded_clip is not stitched_clip
|
||||
assert torch.equal(decoded_clip, stitched_clip)
|
||||
assert vae._split_tiles.call_count == 2
|
||||
assert vae._project_decoder_tile.call_count == 4
|
||||
assert vae._stitch_tiles.call_count == 1
|
||||
assert range_events == [
|
||||
("enter", "minimax_h3.vae.decode_clip"),
|
||||
("enter", "minimax_h3.vae.decode_clip.split_tiles"),
|
||||
("exit", "minimax_h3.vae.decode_clip.split_tiles"),
|
||||
("enter", "minimax_h3.vae.decode_clip.decode_tiles"),
|
||||
("enter", "minimax_h3.vae.decode_clip.tile.0.0"),
|
||||
("enter", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
|
||||
("exit", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
|
||||
("exit", "minimax_h3.vae.decode_clip.tile.0.0"),
|
||||
("enter", "minimax_h3.vae.decode_clip.tile.0.1"),
|
||||
("enter", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
|
||||
("exit", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
|
||||
("exit", "minimax_h3.vae.decode_clip.tile.0.1"),
|
||||
("enter", "minimax_h3.vae.decode_clip.tile.1.0"),
|
||||
("enter", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
|
||||
("exit", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
|
||||
("exit", "minimax_h3.vae.decode_clip.tile.1.0"),
|
||||
("enter", "minimax_h3.vae.decode_clip.tile.1.1"),
|
||||
("enter", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
|
||||
("exit", "minimax_h3.vae.decode_clip.tile.decoder_forward"),
|
||||
("exit", "minimax_h3.vae.decode_clip.tile.1.1"),
|
||||
("exit", "minimax_h3.vae.decode_clip.decode_tiles"),
|
||||
("enter", "minimax_h3.vae.decode_clip.stitch_tiles"),
|
||||
("exit", "minimax_h3.vae.decode_clip.stitch_tiles"),
|
||||
("exit", "minimax_h3.vae.decode_clip"),
|
||||
]
|
||||
|
||||
|
||||
def _tiny_real_vae() -> AutoencoderKLMiniMaxH3:
|
||||
"""Random-weight VAE small enough for a real tiled decode/encode on GPU."""
|
||||
from fastvideo.configs.models.vaes.minimax_h3_video import (
|
||||
MiniMaxH3VideoVAEArchConfig,
|
||||
MiniMaxH3VideoVAEConfig,
|
||||
)
|
||||
|
||||
arch = MiniMaxH3VideoVAEArchConfig(
|
||||
latent_channels=4,
|
||||
block_out_channels=(32, 32),
|
||||
layers_per_block=1,
|
||||
spatial_downsample_factors=(2, 2),
|
||||
temporal_downsample_factors=(2, 2),
|
||||
decoder_num_layers=1,
|
||||
decoder_num_attention_heads=1,
|
||||
decoder_attention_head_dim=8,
|
||||
decoder_num_register_tokens=2,
|
||||
decoder_ffn_mult=1,
|
||||
latents_mean=(0.0, ) * 4,
|
||||
latents_std=(1.0, ) * 4,
|
||||
)
|
||||
return AutoencoderKLMiniMaxH3(
|
||||
MiniMaxH3VideoVAEConfig(
|
||||
arch_config=arch,
|
||||
use_tiling=False,
|
||||
use_temporal_tiling=False,
|
||||
use_parallel_tiling=False,
|
||||
)).eval()
|
||||
|
||||
|
||||
def _dynamo_original(compiled_function: Any) -> Any:
|
||||
"""Return the eager callable behind a ``torch.compile``-decorated function."""
|
||||
original = getattr(compiled_function, "_torchdynamo_orig_callable", None)
|
||||
if original is None:
|
||||
original = getattr(compiled_function, "__wrapped__", None)
|
||||
assert original is not None, "cannot recover the eager tile helpers"
|
||||
return original
|
||||
|
||||
|
||||
@pytest.mark.skipif(not torch.cuda.is_available(), reason="reduce-overhead tile compile requires CUDA graphs")
|
||||
@torch.inference_mode()
|
||||
def test_tiled_decode_and_encode_survive_cudagraph_buffer_reuse_on_cuda() -> None:
|
||||
"""Real tiled decode()/encode() with an unmocked reduce-overhead ``_stitch_tiles``.
|
||||
|
||||
Regression gate for the CUDA-graph output-clobbering bug: the stitched
|
||||
canvas is a cudagraph static buffer, and the collect-then-``torch.cat``
|
||||
consumers (``_decode``/``_encode``/``_encode_pixels``) hold chunk/clip
|
||||
results across subsequent ``_stitch_tiles`` replays. Without the eager
|
||||
``.clone()`` at the tile-driver returns, the first tiled ``decode()`` with
|
||||
>=2 temporal chunks raises ``accessing tensor output of CUDAGraphs that
|
||||
has been overwritten by a subsequent run``. This test needs >=2 chunks
|
||||
(decode), >=2 clips (encode), and a >=2x2 spatial tile grid.
|
||||
"""
|
||||
torch.manual_seed(20260821)
|
||||
vae = _tiny_real_vae().to("cuda")
|
||||
vae.enable_tiling(16, 16, 4, 4)
|
||||
|
||||
# 8 latent tokens = 2 temporal chunks (tokens_chunk_size 5); 8x8 latents =
|
||||
# 32x32 pixels = a 2x2 grid of 16px tiles.
|
||||
z = torch.randn(1, 4, 8, 8, 8, device="cuda")
|
||||
pad_tokens, num_chunks, _ = vae._temporal_decode_plan(z.shape[2])
|
||||
assert num_chunks >= 2, "decode workload must span multiple stitch replays"
|
||||
|
||||
decoded_first = vae.decode(z).sample
|
||||
decoded_second = vae.decode(z).sample
|
||||
assert torch.equal(decoded_first, decoded_second)
|
||||
|
||||
# 34 frames = 2 encode clips of clip_length 17 -> 2 stitch replays.
|
||||
pixels = torch.rand(1, 3, 34, 32, 32, device="cuda")
|
||||
encoded_first = vae.encode(pixels).latent_dist.parameters
|
||||
encoded_second = vae.encode(pixels).latent_dist.parameters
|
||||
assert torch.equal(encoded_first, encoded_second)
|
||||
|
||||
uint8_pixels = torch.randint(0, 256, (1, 3, 34, 32, 32), dtype=torch.uint8)
|
||||
streamed_first = vae.encode_pixels(uint8_pixels).latent_dist.parameters
|
||||
streamed_second = vae.encode_pixels(uint8_pixels).latent_dist.parameters
|
||||
assert torch.equal(streamed_first, streamed_second)
|
||||
|
||||
# Output parity vs the fully eager tile helpers (same weights, same math;
|
||||
# the tolerance absorbs inductor fusion reassociation only).
|
||||
eager_vae = _tiny_real_vae().to("cuda")
|
||||
eager_vae.load_state_dict(vae.state_dict())
|
||||
eager_vae.enable_tiling(16, 16, 4, 4)
|
||||
eager_vae._stitch_tiles = MethodType(_dynamo_original(AutoencoderKLMiniMaxH3._stitch_tiles), eager_vae)
|
||||
eager_vae._project_decoder_tile = MethodType(_dynamo_original(AutoencoderKLMiniMaxH3._project_decoder_tile),
|
||||
eager_vae)
|
||||
torch.testing.assert_close(decoded_first, eager_vae.decode(z).sample, atol=2e-4, rtol=2e-4)
|
||||
torch.testing.assert_close(encoded_first, eager_vae.encode(pixels).latent_dist.parameters, atol=2e-4, rtol=2e-4)
|
||||
@@ -1,11 +1,40 @@
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines import ForwardBatch
|
||||
from fastvideo.worker.gpu_worker import Worker
|
||||
from fastvideo.worker.gpu_worker import Worker, _log_cuda_device_uuid
|
||||
|
||||
|
||||
def test_cuda_device_uuid_receipt_is_disabled_without_nvtx_profiling(monkeypatch) -> None:
|
||||
"""Avoid NVIDIA property access during ordinary worker initialization."""
|
||||
get_device_properties = Mock()
|
||||
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "0")
|
||||
monkeypatch.setattr(torch.cuda, "get_device_properties", get_device_properties)
|
||||
|
||||
_log_cuda_device_uuid(0, torch.device("cuda:0"))
|
||||
|
||||
get_device_properties.assert_not_called()
|
||||
|
||||
|
||||
def test_cuda_device_uuid_receipt_identifies_profiled_worker(monkeypatch) -> None:
|
||||
"""Bind one profiled worker rank to its NVIDIA device UUID in logs."""
|
||||
log_info = Mock()
|
||||
monkeypatch.setenv("FASTVIDEO_NVTX_PROFILE", "1")
|
||||
monkeypatch.setattr(torch.cuda, "get_device_properties", lambda device: SimpleNamespace(uuid="device-uuid"))
|
||||
monkeypatch.setattr("fastvideo.worker.gpu_worker.logger.info", log_info)
|
||||
|
||||
_log_cuda_device_uuid(2, torch.device("cuda:0"))
|
||||
|
||||
log_info.assert_called_once_with(
|
||||
"Worker %d CUDA device UUID: GPU-%s",
|
||||
2,
|
||||
"device-uuid",
|
||||
local_main_process_only=False,
|
||||
)
|
||||
|
||||
|
||||
def _worker_returning(output_batch: ForwardBatch) -> Worker:
|
||||
|
||||
@@ -4,6 +4,7 @@ from typing import Any, cast
|
||||
|
||||
import torch
|
||||
|
||||
import fastvideo.envs as envs
|
||||
from fastvideo.distributed import (cleanup_dist_env_and_memory, maybe_init_distributed_environment_and_model_parallel)
|
||||
from fastvideo.distributed.parallel_state import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
@@ -13,6 +14,14 @@ from fastvideo.pipelines import ForwardBatch, LoRAPipeline, build_pipeline
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def _log_cuda_device_uuid(rank: int, device: torch.device) -> None:
|
||||
"""Record an NVIDIA worker UUID when external NVTX profiling is enabled."""
|
||||
if not envs.FASTVIDEO_NVTX_PROFILE:
|
||||
return
|
||||
device_uuid = torch.cuda.get_device_properties(device).uuid
|
||||
logger.info("Worker %d CUDA device UUID: GPU-%s", rank, device_uuid, local_main_process_only=False)
|
||||
|
||||
|
||||
class Worker:
|
||||
|
||||
def __init__(self, fastvideo_args: FastVideoArgs, local_rank: int, rank: int, distributed_init_method: str):
|
||||
@@ -61,6 +70,8 @@ class Worker:
|
||||
if current_platform.is_cuda_alike():
|
||||
torch.cuda.set_device(self.device)
|
||||
self.init_gpu_memory = torch.cuda.mem_get_info(self.device)[0]
|
||||
if current_platform.is_cuda():
|
||||
_log_cuda_device_uuid(self.rank, self.device)
|
||||
else:
|
||||
# For MPS, we can't get memory info the same way
|
||||
self.init_gpu_memory = 0
|
||||
|
||||
@@ -6,11 +6,8 @@ pipeline with FastVideo's production ``TextEncoderLoader`` path. It covers
|
||||
the three numerical branches the H3 pipelines exercise: text-only tokens,
|
||||
image features, and video features.
|
||||
|
||||
The production encoder is built only as far as the layer-50 conditioning tap
|
||||
by default (``num_hidden_layers_override``), so it returns fewer hidden states
|
||||
than the official full stack. Every state it does build is compared
|
||||
bit-exactly against the official value at the same index, which pins the tap
|
||||
and would catch a truncated stack that still applied the final norm.
|
||||
The production encoder returns only the selected layer-50 hidden state, which
|
||||
is compared bit-exactly with the same state from the official full stack.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -152,13 +149,13 @@ def _make_cases(root: Path) -> dict[str, dict[str, torch.Tensor]]:
|
||||
return cases
|
||||
|
||||
|
||||
def _run_cases(
|
||||
def _run_reference_cases(
|
||||
model: torch.nn.Module,
|
||||
cases: dict[str, dict[str, torch.Tensor]],
|
||||
device: torch.device,
|
||||
) -> dict[str, tuple[torch.Tensor, ...]]:
|
||||
) -> dict[str, torch.Tensor]:
|
||||
dtype = next(model.parameters()).dtype
|
||||
outputs: dict[str, tuple[torch.Tensor, ...]] = {}
|
||||
outputs: dict[str, torch.Tensor] = {}
|
||||
for name, case in cases.items():
|
||||
inputs = {
|
||||
key: value.to(device=device, dtype=dtype if key.startswith("pixel_values") else value.dtype)
|
||||
@@ -172,7 +169,28 @@ def _run_cases(
|
||||
)
|
||||
assert result.hidden_states is not None
|
||||
assert len(result.hidden_states) > MINIMAX_H3_TEXT_ENCODER_LAYER
|
||||
outputs[name] = tuple(hidden_state.detach().cpu() for hidden_state in result.hidden_states)
|
||||
outputs[name] = result.hidden_states[MINIMAX_H3_TEXT_ENCODER_LAYER][0].detach().cpu()
|
||||
return outputs
|
||||
|
||||
|
||||
def _run_production_cases(
|
||||
model: torch.nn.Module,
|
||||
cases: dict[str, dict[str, torch.Tensor]],
|
||||
device: torch.device,
|
||||
) -> dict[str, torch.Tensor]:
|
||||
dtype = next(model.parameters()).dtype
|
||||
outputs: dict[str, torch.Tensor] = {}
|
||||
for name, case in cases.items():
|
||||
inputs = {
|
||||
key: value.to(device=device, dtype=dtype if key.startswith("pixel_values") else value.dtype)
|
||||
for key, value in case.items()
|
||||
if key not in {"attention_mask", "mm_token_type_ids"}
|
||||
}
|
||||
inputs["input_ids"] = inputs["input_ids"][0]
|
||||
with torch.inference_mode():
|
||||
result = model(**inputs)
|
||||
assert result.ndim == 2
|
||||
outputs[name] = result.detach().cpu()
|
||||
return outputs
|
||||
|
||||
|
||||
@@ -217,32 +235,25 @@ def test_minimax_h3_qwen3_vl_parity() -> None:
|
||||
assert not load_errors, f"Official Qwen3-VL checkpoint did not load strictly: {load_errors}"
|
||||
official = official_full.model.eval().to(device)
|
||||
del official_full
|
||||
expected = _run_cases(official, cases, device)
|
||||
expected = _run_reference_cases(official, cases, device)
|
||||
del official
|
||||
_reclaim_vram()
|
||||
|
||||
production = TextEncoderLoader().load(str(root / "text_encoder"), _production_loader_args())
|
||||
assert getattr(production, "_fastvideo_input_device", device) == device
|
||||
actual = _run_cases(production, cases, device)
|
||||
actual = _run_production_cases(production, cases, device)
|
||||
|
||||
assert actual.keys() == expected.keys()
|
||||
# The production stack is built only as far as the conditioning tap by
|
||||
# default (``num_hidden_layers_override``), so it yields one hidden state
|
||||
# per built layer plus the embeddings, while the official model always
|
||||
# yields the full tuple. Every state the production model produces must be
|
||||
# bit-identical to the official value at the same index; the shared-prefix
|
||||
# comparison would in particular catch a truncated stack that still
|
||||
# applied the final norm, which is the failure mode that silently changes
|
||||
# conditioning. With the override set to None the lengths are equal and
|
||||
# this remains the original full comparison, final normed state included.
|
||||
built_layers = int(production.language_model.num_layers)
|
||||
for name in expected:
|
||||
assert len(actual[name]) == built_layers + 1
|
||||
assert len(actual[name]) <= len(expected[name])
|
||||
for layer, (result, reference) in enumerate(zip(actual[name], expected[name], strict=False)):
|
||||
assert_close(result, reference, atol=0.0, rtol=0.0, msg=lambda message: f"{name} layer {layer}: {message}")
|
||||
result = actual[name][MINIMAX_H3_TEXT_ENCODER_LAYER]
|
||||
reference = expected[name][MINIMAX_H3_TEXT_ENCODER_LAYER]
|
||||
result = actual[name]
|
||||
reference = expected[name]
|
||||
assert_close(
|
||||
result,
|
||||
reference,
|
||||
atol=0.0,
|
||||
rtol=0.0,
|
||||
msg=lambda message: f"{name} layer {MINIMAX_H3_TEXT_ENCODER_LAYER}: {message}",
|
||||
)
|
||||
drift = (result.float() - reference.float()).abs()
|
||||
print(
|
||||
f"{name}: max_abs={drift.max().item():.8f} mean_abs={drift.mean().item():.8f}",
|
||||
|
||||
@@ -57,10 +57,9 @@ pytest \
|
||||
```
|
||||
|
||||
With a gate enabled, missing CUDA, source, or weights is a failure. Recorded component evidence is exact for both DiT
|
||||
partitions, the video VAE, and all Qwen3-VL hidden states; audio decode has maximum absolute drift `2.4e-7`. The
|
||||
production Qwen3-VL stack is now built only to the layer-50 conditioning tap by default
|
||||
(`num_hidden_layers_override`), so the encoder gate compares every hidden state the production model builds
|
||||
bit-exactly against the official full stack at the same index.
|
||||
partitions and the video VAE; audio decode has maximum absolute drift `2.4e-7`. The encoder gate compares the slim
|
||||
forward's selected layer-50 hidden state bit-exactly against the same state from the official full stack across text,
|
||||
image, and video inputs.
|
||||
|
||||
The video VAE test verifies the reference checkout at commit
|
||||
`abc5e9bf71fd38f53cd471bc3acaa84bc5ecbfdc` and compares the production CPU `uint8` `encode_pixels()` path against
|
||||
|
||||
Reference in New Issue
Block a user