Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0a861031c7 | ||
|
|
f08c5ee8af | ||
|
|
755f4a7967 | ||
|
|
d543a67b10 | ||
|
|
741aa8d289 |
@@ -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."
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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",
|
||||
]
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user