Compare commits

..
Author SHA1 Message Date
Will LinandClaude Fable 5 0a861031c7 [docs] H3 parallel VAE: document the compiled-decoder cross-process determinism caveat
With enable_torch_compile_vae (#1734, opt-in) inductor autotunes kernels per
process, so chunk decodes on other ranks differ from the serial rank's decode
the way two serial processes differ (GB200 @124f: max 63/255 on <0.5% of
pixels, mean ~1e-2/255; audio and chunk 0 bit-identical). Eager decoder (the
default) stays bitwise-equal to serial decode_to_pixels - measured, both
strategies, x3, 124f+345f.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 18:17:28 +00:00
Will LinandClaude Fable 5 f08c5ee8af [perf] H3 parallel VAE: overlap first decodes with the meta rendezvous; assembly on a side stream
Two schedule fixes sized from the first GB200 tray measurements (job 2659,
serial 7.7s/21.2s at 124f/345f):

1. Every rank now decodes its round-0 chunk BEFORE the metadata broadcast.
   Non-leader ranks previously blocked on the broadcast until the leader
   finished chunk 0, serializing a full extra chunk-decode into round 0
   (visible as 1.9x instead of ~2.6x at 7 chunks / 4 ranks). Same reorder
   on the encode path.

2. The leader's per-chunk joining work (blend, denormalize, clamp, output
   copies) moves to a dedicated CUDA side stream. It depends only on
   already-gathered segments, but on the main stream it delayed the
   leader's next-round decode and therefore every rank's next collective
   (~0.1s/chunk on the critical path). Gathered storage is pinned to the
   assembly stream via record_stream; the driver drains the stream in a
   finally so an exception cannot leave an in-flight DMA into the output
   buffer. Stream placement does not change op order or values, so the
   bitwise-parity contract is untouched.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 17:59:18 +00:00
Will LinandClaude Fable 5 755f4a7967 [docs] schema parity inventory: classify vae_parallel_* (+ the stack's unclassified VSA_tile_size)
vae_parallel_decode / vae_parallel_encode / vae_parallel_decode_strategy are
model-specific optimization knobs (compatibility_only, like VSA_sparsity).
VSA_tile_size came in with the merged tile-64 route without an inventory
entry and failed test_fastvideo_args_fields_are_classified on the whole
stack; classify it the same way. The remaining pipeline_config inventory
gaps (image_encoder_precisions, ...) predate this branch and are left for
the owning PRs.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 17:53:41 +00:00
Will LinandClaude Fable 5 d543a67b10 [perf] MiniMax-H3 VAE: SP-rank-parallel chunk decode + reference-clip encode (opt-in)
Under SP>1 the H3 video VAE decoded all temporal chunks serially on the
output rank while the other ranks idled (#1703's gate), and every rank
encoded the full reference video redundantly. Chunk decodes and clip
encodes have no cross-chunk data dependency - only the joining (overlap
blend, trim, denormalize, moment concat) is sequential - so both are
round-robined across the sequence-parallel ranks:

- fastvideo/models/vaes/minimax_h3_parallel.py: decode_to_pixels_parallel
  gathers each round's decoded segments (body+halo tail, one contiguous
  slice per chunk) to the SP group's first rank via NCCL gather (or
  all_gather, FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY), which replays the
  serial blend/trim/denormalize/copy semantics with the same VAE methods -
  bitwise-equal to serial decode_to_pixels by construction. Placeholder
  rounds keep collective participation uniform; the leader decodes chunk 0
  first and broadcasts dtype/shape metadata so placeholders never guess the
  autocast dtype. encode_pixels_parallel all-gathers per-clip moments
  (latent-sized) so every rank keeps the identical full posterior,
  preserving the all-ranks-hold-latents contract.
- decoding stage: output gate moves from world rank 0 to the SP group's
  first rank (identical in the single-group e2e case; correct for the
  trainer validation callback, which consumes each group leader's batch);
  with vae_parallel_decode every rank enters the decode body so no
  rank-dependent branch guards the collectives.
- latent preparation: opt-in clip-parallel reference encode on the same
  seam (vae_parallel_encode).
- knobs: FastVideoArgs.vae_parallel_decode/encode (+ --vae-parallel-decode,
  --vae-parallel-encode, FASTVIDEO_VAE_PARALLEL_DECODE/ENCODE env
  parse-once adapters), default OFF.
- _copy_chunk_pixels factored out of _decode_to_pixels so serial and
  parallel share one output-copy path (behavior unchanged).

Tests: threaded fake-group CPU suite drives the real SPMD functions
end-to-end (world sizes 2-5, both strategies, pad/blend/trim geometries,
token_drop=0, batched slicing, placeholder rounds) bit-exact vs the serial
APIs; GPU regression (torchrun world>1 gated) asserts bitwise parity under
fp16 autocast with real NCCL plus repeat-determinism x3.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 17:38:38 +00:00
Will LinandClaude Fable 5 741aa8d289 [bugfix] logger: info_once crashed on the patched process-aware info (duplicate stacklevel)
_print_info_once passes stacklevel=2 into logger.info, and init_logger's
patched _info passed its own stacklevel=2 positionally into logger.log on
top of the caller's kwarg -> TypeError on every info_once call. Honor an
explicit stacklevel instead of passing the keyword twice.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 17:38:18 +00:00
19 changed files with 1028 additions and 369 deletions
@@ -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."
-25
View File
@@ -337,31 +337,6 @@ Only DiT submodules that declare `_compile_conditions` are compiled
(most shipped models). The text encoder and VAE are not compiled by this
flag.
### Regional fullgraph compile (experimental)
`inference_torch_compile` is a stricter, kwargs-free variant that ports the
training-side regional compile of
[#1718](https://github.com/hao-ai-lab/FastVideo/pull/1718) to inference: the
loader wraps each `_compile_conditions` block in
`torch.compile(fullgraph=True)` with inductor
`options={"emulate_precision_casts": True}` right after the transformer
loads. Attention backends that cannot be traced end-to-end (VSA,
FLASH_ATTN on flash-attn 3, or the `FASTVIDEO_DISABLE_ATTENTION_COMPILE=1`
escape hatch) degrade the transformer to eager with one warning instead of
failing mid-denoise.
```python
generator = VideoGenerator.from_pretrained(
"MiniMaxAI/MiniMax-H3",
inference_torch_compile=True, # or FASTVIDEO_INFERENCE_TORCH_COMPILE=1
)
```
Do not combine it with `torch_compile_kwargs['mode']` (the loader injects
inductor options, and torch.compile forbids mode+options); it is
independent of `enable_torch_compile`, and when both are set the regional
compile wins for the DiT.
### What to expect
| Config | Effect |
-8
View File
@@ -85,12 +85,6 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--compile-mode",
default=None,
help='torch.compile mode, e.g. "reduce-overhead" for CUDA graphs')
parser.add_argument("--inference-torch-compile",
action="store_true",
help="regional fullgraph torch.compile of each DiT block after load. NOTE: this "
"script always runs the VSA-H3 backend, which is not fullgraph-traceable — the "
"loader logs one warning and keeps the transformer eager. The flag is exposed "
"here to exercise exactly that guard")
parser.add_argument("--repeats",
type=int,
default=1,
@@ -125,8 +119,6 @@ def main() -> None:
}
if args.vsa_sparsity > 0.0:
experimental["VSA_sparsity"] = args.vsa_sparsity
if args.inference_torch_compile:
experimental["inference_torch_compile"] = True
generator = VideoGenerator.from_config(
GeneratorConfig(
@@ -15,7 +15,6 @@ from fastvideo.api import (
OffloadConfig,
OutputConfig,
ParallelismConfig,
PipelineSelection,
SamplingConfig,
)
@@ -42,12 +41,6 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--compile-mode",
default=None,
help='torch.compile mode, e.g. "reduce-overhead" for CUDA graphs')
parser.add_argument("--inference-torch-compile",
action="store_true",
help="regional fullgraph torch.compile of each DiT block after load (the #1718 "
"training-port semantics: no kwargs; fullgraph + emulate_precision_casts injected). "
"First generation pays the inductor JIT (~1-2 min); use --repeats >= 2 and time "
"the last repeat. FASTVIDEO_INFERENCE_TORCH_COMPILE=1 is equivalent")
parser.add_argument("--repeats",
type=int,
default=1,
@@ -61,16 +54,9 @@ def main() -> None:
output_dir = Path(args.output)
output_dir.mkdir(parents=True, exist_ok=True)
# Boot-time run configuration folded into FastVideoArgs (the same
# experimental-dict route basic_fasth3.py uses for the VSA knobs).
experimental: dict[str, object] = {}
if args.inference_torch_compile:
experimental["inference_torch_compile"] = True
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,
+4 -4
View File
@@ -18,13 +18,13 @@ from fastvideo.layers.rotary_embedding import _apply_rotary_emb
def _attention_compile_disabled() -> bool:
"""Whether to keep attention ``forward`` out of the torch.compile graph.
Attention backends expose traceable custom-op boundaries, so compilation
is enabled by default. Set ``FASTVIDEO_DISABLE_ATTENTION_COMPILE=1`` for
an explicit eager escape hatch when debugging a backend.
Defaults to ``True`` (the historical behavior: attention runs eager via
``torch.compiler.disable``). Set ``FASTVIDEO_DISABLE_ATTENTION_COMPILE=0``
to let attention be traced/compiled into the surrounding graph.
"""
val = os.environ.get("FASTVIDEO_DISABLE_ATTENTION_COMPILE")
if val is None:
return False
return True
return val.strip().lower() not in ("0", "false", "no", "off", "")
+15 -11
View File
@@ -21,8 +21,10 @@ if TYPE_CHECKING:
FASTVIDEO_TRACE_FUNCTION: int = 0
FASTVIDEO_ATTENTION_BACKEND: str | None = None
FASTVIDEO_FA4: bool = False
FASTVIDEO_INFERENCE_TORCH_COMPILE: 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
@@ -220,16 +222,18 @@ environment_variables: dict[str, Callable[[], Any]] = {
"FASTVIDEO_FA4":
lambda: os.getenv("FASTVIDEO_FA4", "0") != "0",
# If set (=1), enable regional (per-transformer-block) fullgraph
# torch.compile for the DiT at inference — the inference-side counterpart
# of the training regional-compile port of hao-ai-lab/FastVideo#1718.
# Equivalent to FastVideoArgs.inference_torch_compile=True (e.g. via
# PipelineSelection.experimental={"inference_torch_compile": True}). VSA
# and other non-fullgraph-traceable attention backends degrade to eager
# with one warning; see _regional_compile_unsupported_reason in
# fastvideo/models/loader/fsdp_load.py.
"FASTVIDEO_INFERENCE_TORCH_COMPILE":
lambda: os.getenv("FASTVIDEO_INFERENCE_TORCH_COMPILE", "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
+44 -28
View File
@@ -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).
@@ -164,18 +177,6 @@ class FastVideoArgs:
torch_compile_kwargs_text_encoder: dict[str, Any] = field(default_factory=dict)
torch_compile_kwargs_vae: dict[str, Any] = field(default_factory=dict)
torch_compile_kwargs_audio_vae: dict[str, Any] = field(default_factory=dict)
# Regional (per-transformer-block) fullgraph torch.compile of the DiT at
# inference — the inference-side counterpart of the training regional
# compile ported from hao-ai-lab/FastVideo#1718. Applied by the loader
# right after the transformer loads, with fullgraph=True and inductor
# options {emulate_precision_casts: True} injected (no user kwargs
# needed). Attention backends that cannot be fullgraph-traced (VSA,
# FLASH_ATTN on flash-attn 3) degrade the transformer to eager with one
# warning. Opt-in via FASTVIDEO_INFERENCE_TORCH_COMPILE=1 (folded in
# __post_init__) or PipelineSelection.experimental
# {"inference_torch_compile": true}. Distinct from ``enable_torch_compile``,
# which keeps the pipeline-level compile semantics.
inference_torch_compile: bool = False
disable_autocast: bool = False
@@ -284,13 +285,6 @@ class FastVideoArgs:
self._apply_ltx2_vae_overrides()
self._resolve_refine_args()
self._apply_transformer_quant()
if not self.inference_torch_compile:
# Parse-once adapter (same pattern as attention_backend below): the
# environment variable is an input read once here, so the loader
# only ever consults the typed field.
import fastvideo.envs as envs
if envs.FASTVIDEO_INFERENCE_TORCH_COMPILE:
self.inference_torch_compile = True
if self.attention_backend is not None:
# Fail fast on typos instead of silently auto-selecting later.
from fastvideo.attention.selector import coerce_attn_backend
@@ -306,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``.
@@ -611,15 +624,6 @@ class FastVideoArgs:
help=
"JSON string of kwargs to pass to torch.compile. Example: '{\"backend\":\"inductor\",\"mode\":\"reduce-overhead\"}'",
)
parser.add_argument(
"--inference-torch-compile",
action=StoreBoolean,
default=FastVideoArgs.inference_torch_compile,
help="Regional fullgraph torch.compile of each DiT transformer block at inference "
"(port of the #1718 training-side regional compile). The loader injects fullgraph=True "
"and inductor options {emulate_precision_casts: true}; non-traceable attention backends "
"(VSA) degrade to eager with one warning. FASTVIDEO_INFERENCE_TORCH_COMPILE=1 is equivalent.",
)
parser.add_argument(
"--dit-cpu-offload",
@@ -660,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,
+4 -1
View File
@@ -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
@@ -1126,7 +1126,6 @@ class TransformerLoader(ComponentLoader):
training_mode=fastvideo_args.training_mode,
enable_torch_compile=fastvideo_args.enable_torch_compile,
torch_compile_kwargs=fastvideo_args.torch_compile_kwargs,
inference_regional_compile=fastvideo_args.inference_torch_compile,
)
total_params = sum(p.numel() for p in model.parameters())
-117
View File
@@ -136,7 +136,6 @@ def maybe_load_fsdp_model(
pin_cpu_memory: bool = True,
enable_torch_compile: bool = False,
torch_compile_kwargs: dict[str, Any] | None = None,
inference_regional_compile: bool = False,
) -> torch.nn.Module:
"""
Load the model with FSDP if is training, else load the model without FSDP.
@@ -236,125 +235,9 @@ def maybe_load_fsdp_model(
logger.info("Enabling torch.compile for FSDP training module with kwargs=%s", compile_kwargs)
model = torch.compile(model, **compile_kwargs)
logger.info("torch.compile enabled for %s", type(model).__name__)
elif inference_regional_compile and not training_mode:
# Inference-side counterpart of the #1718 training regional compile:
# per-block fullgraph compile right after the transformer loads, no
# user kwargs needed (fullgraph + emulate_precision_casts injected).
unsupported = _regional_compile_unsupported_reason(init_params)
if unsupported is not None:
logger.warning(
"inference_torch_compile requested but disabled: %s. "
"Inference continues in eager mode.", unsupported)
else:
prepare_for_compile = getattr(model, "prepare_for_compile", None)
if callable(prepare_for_compile):
logger.info("Running prepare_for_compile for %s", type(model).__name__)
prepare_for_compile()
_compile_model_regions(model, torch_compile_kwargs or {})
return model
def _regional_compile_unsupported_reason(init_params: dict[str, Any]) -> str | None:
"""Return why regional fullgraph compile cannot run, or None if it can.
FA3's grad-enabled attention path deliberately routes to the raw
autograd.Function at the cost of a dynamo graph break (see
flash_attn_default.py) — under regional ``fullgraph=True`` that break is
a hard RuntimeError at the first compiled forward. FA2 and FA4 route
through traceable custom ops and are compile-safe.
The VSA backends (Triton block-sparse kernels behind sequence-parallel
all-to-alls plus a host-synced metadata guard) are likewise not
fullgraph-traceable; a VSA-backed transformer (e.g. the FastH3 student)
falls back to eager while compile-safe dense loads still compile.
"""
try:
from fastvideo.attention.layer import _attention_compile_disabled
except Exception: # pragma: no cover - attention stack not importable
pass
else:
if _attention_compile_disabled():
# The escape hatch wraps attention forwards in
# torch.compiler.disable, which is a hard dynamo error inside a
# fullgraph region ("Skip inlining `torch.compiler.disable()`d
# function"). Degrade to eager instead, matching the hatch's
# debugging intent.
return ("FASTVIDEO_DISABLE_ATTENTION_COMPILE=1 keeps attention "
"forwards out of compiled graphs via torch.compiler."
"disable, which fullgraph regional compile cannot trace; "
"this model stays eager")
config = init_params.get("config")
resolved = getattr(config, "_resolved_attention_backend", None)
resolved_name = getattr(resolved, "name", "")
if resolved_name in ("VIDEO_SPARSE_ATTN", "VIDEO_SPARSE_ATTN_H3"):
return (f"attention backend resolved to {resolved_name}, whose Triton "
"kernels, sequence-parallel collectives, and sync metadata "
"guard graph-break (incompatible with fullgraph regional "
"compile); this model stays eager")
if resolved is None or resolved_name != "FLASH_ATTN":
return None
try:
from fastvideo.attention.utils.flash_attn_default import fa_version
except Exception: # pragma: no cover - flash-attn stack not importable
return None
if fa_version == "3":
return ("attention backend resolved to FLASH_ATTN with flash-attn 3, "
"whose grad-enabled path graph-breaks (incompatible with "
"fullgraph regional compile); use FA2, FA4 (FASTVIDEO_FA4=1), "
"or TORCH_SDPA for compiled runs")
return None
def _compile_model_regions(model: nn.Module, compile_kwargs: dict[str, Any]) -> int:
"""Compile repeated mathematical regions of a loaded model.
Only the selected module ``forward`` is replaced. This keeps activation
checkpoint wrappers structurally transparent while any module-level hooks
(FSDP pre/post, layerwise offload) execute outside the compiled region.
"""
compile_conditions = getattr(model, "_compile_conditions", None)
if not compile_conditions:
raise ValueError(f"{type(model).__name__} does not declare _compile_conditions")
if compile_kwargs.get("fullgraph", True) is not True:
raise ValueError("Regional compile requires fullgraph=True")
if "mode" in compile_kwargs:
# torch.compile forbids passing both `mode` and `options`, and
# regional compile always injects options (emulate_precision_casts)
# for bf16 numerics parity. Fail here with an actionable message
# instead of letting torch raise a mode/options conflict about an
# `options` key the user never wrote.
raise ValueError("Regional compile sets inductor options "
"(emulate_precision_casts) and cannot be combined "
"with torch_compile_kwargs['mode']. Remove 'mode' or "
"express its effect via torch_compile_kwargs['options'].")
kwargs = {**compile_kwargs, "fullgraph": True}
options = {"emulate_precision_casts": True}
options.update(kwargs.get("options") or {})
kwargs["options"] = options
compiled_count = 0
for name, submodule in list(model.named_modules()):
if not name:
continue
if any(condition(name, submodule) for condition in compile_conditions):
# Activation checkpoint wrappers are control-flow boundaries, not
# mathematical regions. Keep their saved-tensor/recompute logic
# eager and compile only the repeated block they own.
compile_target = getattr(submodule, "_checkpoint_wrapped_module", submodule)
compile_target.forward = torch.compile(compile_target.forward, **kwargs)
compiled_count += 1
if compiled_count == 0:
raise ValueError(f"No submodules in {type(model).__name__} matched _compile_conditions")
logger.info(
"Enabled regional torch.compile for %d submodules in %s with kwargs=%s",
compiled_count,
type(model).__name__,
kwargs,
)
return compiled_count
def shard_model(
model,
*,
@@ -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",
]
+25 -14
View File
@@ -933,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(
@@ -7,9 +7,11 @@ 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 (
@@ -24,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)
@@ -32,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."""
@@ -57,11 +78,13 @@ class MiniMaxH3VideoDecodingStage(PipelineStage):
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""Decode H3 video latents into normalized CPU pixels."""
if model_parallel_is_initialized() and not get_world_group().is_first_rank:
# Distributed executors consume rank 0's ForwardBatch. Keep a
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)
@@ -81,23 +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(),
)
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"),
):
self.vae.decode_to_pixels(latents, output)
batch.output = output
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:
@@ -128,7 +161,9 @@ class MiniMaxH3AudioDecodingStage(PipelineStage):
@torch.no_grad()
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""Decode H3 audio latents into a stereo CPU waveform."""
if model_parallel_is_initialized() and not get_world_group().is_first_rank:
# 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)
@@ -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")
@@ -203,14 +203,6 @@ class ComposedPipelineBase(ABC):
vae_compile_kwargs = (self.fastvideo_args.torch_compile_kwargs_vae or global_compile_kwargs)
audio_vae_compile_kwargs = (self.fastvideo_args.torch_compile_kwargs_audio_vae or global_compile_kwargs)
if compile_transformer and self.fastvideo_args.inference_torch_compile:
# The loader already applied the regional fullgraph
# compile to the DiT blocks (inference_torch_compile);
# wrapping the same forwards again here would stack
# compiled callables.
logger.info("inference_torch_compile already compiled the DiT regions in the "
"loader; skipping the pipeline-level DiT compile")
compile_transformer = False
if compile_transformer:
self._maybe_compile_pipeline_module(
module_name="transformer",
@@ -1,117 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""Contract tests for the inference-side regional torch.compile port.
The loader applies a per-transformer-block fullgraph compile after the
transformer loads (``FastVideoArgs.inference_torch_compile``, env
``FASTVIDEO_INFERENCE_TORCH_COMPILE=1``). These tests pin the two pieces that
must not drift from the #1718 training-port semantics:
- ``_regional_compile_unsupported_reason``: VSA backends (and the attention
eager escape hatch) degrade to eager with a reason instead of hard-failing
fullgraph capture at the first denoising forward.
- ``_compile_model_regions``: exactly the ``_compile_conditions`` blocks are
compiled, fullgraph=True plus inductor ``emulate_precision_casts`` are
injected, and ``mode`` kwargs are rejected (torch.compile forbids
mode+options).
CPU-safe: torch.compile is monkeypatched, no CUDA needed.
"""
from types import SimpleNamespace
import pytest
import torch
from torch import nn
from fastvideo.models.loader import fsdp_load
from fastvideo.models.loader.fsdp_load import (
_compile_model_regions,
_regional_compile_unsupported_reason,
)
def _init_params_for(backend_name: str | None) -> dict:
resolved = None if backend_name is None else SimpleNamespace(name=backend_name)
return {"config": SimpleNamespace(_resolved_attention_backend=resolved)}
@pytest.mark.parametrize("backend_name", ["VIDEO_SPARSE_ATTN", "VIDEO_SPARSE_ATTN_H3"])
def test_vsa_backends_degrade_to_eager(backend_name, monkeypatch) -> None:
monkeypatch.delenv("FASTVIDEO_DISABLE_ATTENTION_COMPILE", raising=False)
reason = _regional_compile_unsupported_reason(_init_params_for(backend_name))
assert reason is not None
assert backend_name in reason
assert "eager" in reason
@pytest.mark.parametrize("backend_name", [None, "TORCH_SDPA"])
def test_dense_backends_allow_compile(backend_name, monkeypatch) -> None:
monkeypatch.delenv("FASTVIDEO_DISABLE_ATTENTION_COMPILE", raising=False)
assert _regional_compile_unsupported_reason(_init_params_for(backend_name)) is None
def test_attention_compile_escape_hatch_degrades_to_eager(monkeypatch) -> None:
monkeypatch.setenv("FASTVIDEO_DISABLE_ATTENTION_COMPILE", "1")
reason = _regional_compile_unsupported_reason(_init_params_for("TORCH_SDPA"))
assert reason is not None
assert "FASTVIDEO_DISABLE_ATTENTION_COMPILE" in reason
class _Block(nn.Module):
def __init__(self) -> None:
super().__init__()
self.linear = nn.Linear(4, 4)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.linear(x)
class _Toy(nn.Module):
_compile_conditions = [lambda name, module: name.startswith("blocks.") and name.count(".") == 1]
def __init__(self) -> None:
super().__init__()
self.blocks = nn.ModuleList([_Block() for _ in range(3)])
self.proj_out = nn.Linear(4, 4)
def forward(self, x: torch.Tensor) -> torch.Tensor:
for block in self.blocks:
x = block(x)
return self.proj_out(x)
def test_compile_model_regions_injects_fullgraph_and_precision_casts(monkeypatch) -> None:
captured: list[dict] = []
def _fake_compile(fn, **kwargs):
captured.append(kwargs)
return fn
monkeypatch.setattr(fsdp_load.torch, "compile", _fake_compile)
model = _Toy()
count = _compile_model_regions(model, {})
# The three repeated blocks compile; proj_out and the root stay eager.
assert count == 3
assert len(captured) == 3
for kwargs in captured:
assert kwargs["fullgraph"] is True
assert kwargs["options"] == {"emulate_precision_casts": True}
def test_compile_model_regions_rejects_mode_kwargs() -> None:
with pytest.raises(ValueError, match="mode"):
_compile_model_regions(_Toy(), {"mode": "reduce-overhead"})
def test_compile_model_regions_requires_conditions_and_matches(monkeypatch) -> None:
monkeypatch.setattr(fsdp_load.torch, "compile", lambda fn, **kwargs: fn)
plain = nn.Linear(4, 4)
with pytest.raises(ValueError, match="_compile_conditions"):
_compile_model_regions(plain, {})
class _NoMatch(_Toy):
_compile_conditions = [lambda name, module: False]
with pytest.raises(ValueError, match="matched"):
_compile_model_regions(_NoMatch(), {})
@@ -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,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)