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
11 changed files with 1025 additions and 36 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."
+16
View File
@@ -22,6 +22,9 @@ if TYPE_CHECKING:
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
@@ -219,6 +222,19 @@ 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
+44
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).
@@ -287,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``.
@@ -632,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
@@ -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")
@@ -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)