Compare commits

...
Author SHA1 Message Date
SolitaryThinker 1568d3d1f7 [test]: require actual tuned Ulysses execution in native validation 2026-09-07 01:43:25 +00:00
SolitaryThinker a61f89a194 [docs]: explain bounded Ulysses chunks and long training policy 2026-09-06 21:27:45 +00:00
SolitaryThinker 89836e8a98 [perf]: preserve the full baseline plan for long H3 training 2026-09-06 21:04:57 +00:00
SolitaryThinker 76129f5a89 [perf]: add bounded H3 windows and agreed GB200 launch plans 2026-09-06 20:41:05 +00:00
SolitaryThinker a06e63827d [test] Cover real CUDA capture and recovery after a cold decline 2026-09-06 00:34:17 +00:00
SolitaryThinker e536bb8544 [bugfix] Keep Ulysses rank agreement unconditional and reuse vote buffers 2026-09-06 00:23:28 +00:00
shaoxiongduan 3eabb7b40b [bugfix] gate the fused Ulysses a2a on one host, and stop re-voting per call
Two defects measured on 4x GB200, MiniMax-H3 geometry, per attention layer
(NCCL baseline 2421us at sp=4, 1295us at sp=8 across two trays):

  as shipped   sp=4 2035us (1.19x)   sp=8 3102us (2.4x SLOWER than NCCL)
  with these   sp=4 1590us (1.52x)   sp=8 declines, 1300us

ncclTeamLsa answers "addressable", not "fast". NCCL 2.29 extends the LSA
team across a multi-node NVLink domain, so on a GB200 rack the gate passes
for ranks on different trays and the kernel arms. Its fine-grained 16B
remote stores are far slower there than NCCL's bulk transfers. Require a
single host, which is the regime the slab decomposition was tuned for, and
which matches flashinfer's own gate. torch 2.12 (the pin) bundles NCCL
2.29.7, so this is reachable today.

_can_attempt now computes its local verdict without collectives and then
runs one unconditional all_gather_object carrying (hostname, local_ok).
A collective behind a rank-local early return hangs the group whenever
ranks disagree -- exactly the case this gate exists to detect. Verified:
with one rank reporting the kernel unavailable, the earlier ordering timed
out at 180s while this completes with correct results on every rank.

The per-call agreement cost a flat ~227us regardless of operand size -- a
host-side gloo all_gather before every collective. The signature is
architectural: two distinct values (scatter and gather shapes) across 50
layers x 4 steps. Cache the verdict so the collective runs twice per
generation instead of ~400 times.

The cache trades one property: a rank whose signature diverges mid-run now
misses the cache and calls the collective alone, hanging rather than
falling back. Re-voting every N calls would bound that if wanted.
2026-09-01 09:52:02 +00:00
12 changed files with 625 additions and 38 deletions
+3
View File
@@ -1,5 +1,8 @@
# Profiling FastVideo
For GB200 sequence-parallel transport tuning and its long-training memory
policy, see [Ulysses performance](ulysses_performance.md).
!!! warning
Profiling is only intended for FastVideo developers and maintainers to understand the proportion of time spent in different parts of the codebase. **FastVideo end-users should never turn on profiling** as it will significantly slow down inference.
+59
View File
@@ -0,0 +1,59 @@
# Ulysses performance on GB200
Enable the fused transport with `FASTVIDEO_ULYSSES_A2A=auto` and rebuild
`fastvideo-kernel` from the same checkout. The tuned path applies to contiguous
BF16 operands with 56 global heads, head dimension 128, and sequence parallelism
of four on GB200. It uses 144 CTAs and exchanges one batch plane at a time.
Older kernel builds retain the original 36-CTA path.
Each rank keeps at most 1 GiB of registered window storage. A 250,000-token
sequence needs 896,000,000 bytes per plane with this geometry, so packed QKV or
QKVG can use the fused path without a window large enough for the entire pack.
The full output still needs its own GPU allocation. Chunking reduces registered
storage; it does not remove activation or optimizer memory costs.
Every call collectively agrees on geometry, window capacity, chunking and CTA
count before entering the kernel. Outputs own their storage, and backward uses
the plan saved by its forward. Contiguity alone does not replace these checks.
Existing unsupported-layout, capture, topology and lifecycle fallbacks remain.
## Long training sequences
By default, grad-tracked operands whose per-plane size exceeds 512 MiB keep
the original transport policy, including NCCL fallback for oversized packs.
This avoids enabling a measured allocation-pressure regression in a 250k-token
FSDP4 training recipe. No-grad inference can use tuned chunks directly.
After configuring sufficient activation memory headroom, enable long training
chunks explicitly:
```bash
export FASTVIDEO_ULYSSES_A2A=auto
export FASTVIDEO_ULYSSES_A2A_LONG_TRAINING=chunked
```
An activation offload policy can provide that headroom. Measure the complete
forward, backward, clipping and optimizer step with resident optimizer state;
transport microbenchmarks alone do not predict long-sequence training speed.
Compare allocator retries alongside step latency. Keep the allocator, FSDP
mesh, attention backend, checkpointing and precision identical across routes.
`FASTVIDEO_ULYSSES_A2A_LONG_TRAINING=auto` restores the conservative policy.
The 512 MiB boundary describes a transport operand; it is not a model-wide
activation-memory estimate. Revalidate other training recipes and hardware.
## Validation
CPU policy regressions live in
`fastvideo/tests/distributed/test_ulysses_h3_policy.py`. On an exclusive group
of four GB200 GPUs, run the native transport gate after rebuilding the kernel:
```bash
FASTVIDEO_ULYSSES_A2A=auto torchrun --standalone --nproc_per_node=4 \
fastvideo/tests/distributed/check_ulysses_h3_native.py
```
The gate covers exact transport gradients at 32k, 128k and 250k, the long
training opt-in, retained-output ownership, and recovery after rank capability
disagreement. Full-model throughput and absolute model FLOP utilization (MFU)
are separate measurements; this transport gate reports neither.
@@ -23,6 +23,7 @@
// The per-group context is an ncclDevComm plus a registered symmetric window,
// both created here from the caller's ncclComm_t.
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <c10/cuda/CUDAStream.h>
#include <torch/extension.h>
@@ -162,13 +163,19 @@ bool ulysses_lsa_covers_group(int64_t comm_ptr, int64_t world_size) {
// mode == 0: inp [B, S_local, H, D] -> out [B, S_global, H_local, D]
// mode == 1: inp [B, S_global, H_local, D] -> out [B, S_local, H, D]
// where H is the *global* head count and H_local = H / world_size.
void ulysses_a2a(int64_t handle, torch::Tensor inp, torch::Tensor out, int64_t B, int64_t S_local,
int64_t H, int64_t D, int64_t mode) {
void ulysses_a2a_tuned(int64_t handle, torch::Tensor inp, torch::Tensor out, int64_t B, int64_t S_local,
int64_t H, int64_t D, int64_t mode, int64_t launch_blocks) {
auto* ctx = reinterpret_cast<UlyssesContext*>(handle);
TORCH_CHECK(ctx != nullptr, "handle must come from allocate_ulysses_a2a");
const at::cuda::CUDAGuard device_guard(inp.device());
auto stream = at::cuda::getCurrentCUDAStream();
TORCH_CHECK(launch_blocks == 36 || launch_blocks == 144, "Ulysses launch must use 36 or 144 CTAs");
if (launch_blocks == 144) {
const auto* props = at::cuda::getCurrentDeviceProperties();
TORCH_CHECK(props->major == 10 && props->minor == 0 && props->multiProcessorCount >= 144,
"144-CTA Ulysses launch requires an sm100 device with at least 144 SMs");
}
TORCH_CHECK(inp.is_cuda() && out.is_cuda(), "inp and out must be CUDA tensors");
TORCH_CHECK(inp.is_contiguous() && out.is_contiguous(), "inp and out must be contiguous");
@@ -204,7 +211,7 @@ void ulysses_a2a(int64_t handle, torch::Tensor inp, torch::Tensor out, int64_t B
const int64_t num_rows = B * static_cast<int64_t>(W) * S_local;
const int blocks =
static_cast<int>(std::max<int64_t>(1, std::min<int64_t>(fi::kMaxBlocks, num_rows)));
static_cast<int>(std::max<int64_t>(1, std::min<int64_t>(launch_blocks, num_rows)));
const int threads = fi::kUlyssesThreads;
#define LAUNCH_ULYSSES_A2A(T, NG, MODE) \
@@ -266,6 +273,12 @@ void ulysses_a2a(int64_t handle, torch::Tensor inp, torch::Tensor out, int64_t B
TORCH_CHECK(status == cudaSuccess, "ulysses_a2a copy-out failed: ", cudaGetErrorString(status));
}
// Preserve the original wheel API and its 36-CTA launch policy.
void ulysses_a2a(int64_t handle, torch::Tensor inp, torch::Tensor out, int64_t B, int64_t S_local,
int64_t H, int64_t D, int64_t mode) {
ulysses_a2a_tuned(handle, inp, out, B, S_local, H, D, mode, 36);
}
void register_ulysses_a2a(pybind11::module_& m) {
m.def("allocate_ulysses_a2a", &allocate_ulysses_a2a, "allocate a local ulysses a2a window");
m.def("register_ulysses_a2a_window", &register_ulysses_a2a_window,
@@ -276,4 +289,5 @@ void register_ulysses_a2a(pybind11::module_& m) {
m.def("ulysses_lsa_covers_group", &ulysses_lsa_covers_group,
"whether the whole group is load-store accessible");
m.def("ulysses_a2a", &ulysses_a2a, "fused-transpose Ulysses all-to-all over NVLink");
m.def("ulysses_a2a_tuned", &ulysses_a2a_tuned, "Ulysses all-to-all with a collectively agreed CTA count");
}
@@ -50,7 +50,8 @@ namespace ulysses {
constexpr int kUlyssesThreads = 512;
// Deliberately modest: this is link-bandwidth bound, so a small grid leaves the
// rest of the GPU free without costing throughput.
constexpr int kMaxBlocks = 36;
// One distinct NCCL barrier slot per CTA, including the tuned GB200 launch.
constexpr int kMaxBlocks = 144;
// Shared movement body for the fused-transpose all-to-all (no barriers).
//
@@ -29,6 +29,11 @@ def is_available() -> bool:
return _ops is not None and all(hasattr(_ops, name) for name in _REQUIRED_OPS)
def supports_tuned_launch() -> bool:
"""Whether this build has the expanded barrier capacity and tuned entrypoint."""
return is_available() and hasattr(_ops, "ulysses_a2a_tuned")
def _require() -> None:
if not is_available():
raise RuntimeError(
@@ -73,7 +78,7 @@ def dispose(handle: int) -> None:
def all_to_all(handle: int, inp: torch.Tensor, out: torch.Tensor, B: int, S_local: int, H: int,
D: int, mode: int) -> None:
D: int, mode: int, *, blocks: int = 36) -> None:
"""Run one fused all-to-all on the current stream, writing into ``out``.
``mode == 0``: ``[B, S_local, H, D] -> [B, S_global, H_local, D]``
@@ -83,4 +88,9 @@ def all_to_all(handle: int, inp: torch.Tensor, out: torch.Tensor, B: int, S_loca
geometry in the same order.
"""
_require()
_ops.ulysses_a2a(int(handle), inp, out, int(B), int(S_local), int(H), int(D), int(mode))
if blocks == 36:
_ops.ulysses_a2a(int(handle), inp, out, int(B), int(S_local), int(H), int(D), int(mode))
elif blocks == 144 and supports_tuned_launch():
_ops.ulysses_a2a_tuned(int(handle), inp, out, int(B), int(S_local), int(H), int(D), int(mode), int(blocks))
else:
raise ValueError(f"unsupported Ulysses launch: blocks={blocks}; rebuild the kernel for tuned launches")
@@ -6,6 +6,9 @@ load-store accessible NVLink mesh: same layout, byte-identical results, fewer
passes over local memory. Anything else falls back to the NCCL path.
"""
import socket
from array import array
import torch
import torch.distributed as dist
from torch.distributed import ProcessGroup
@@ -27,6 +30,8 @@ _DTYPE_CODES = {
# Bound persistent registered memory per rank. Larger operands use NCCL instead
# of growing the window without limit.
MAX_WINDOW_BYTES = 1024**3
H3_TRAINING_PLANE_LIMIT_BYTES = 512 * 1024**2
_CONTRACT_SIZE = 12
# (scatter_dim, gather_dim) -> kernel mode.
# 0: [B, S_local, H, D] -> [B, S_global, H_local, D]
@@ -47,17 +52,22 @@ class _FusedUlyssesA2A(torch.autograd.Function):
"""
@staticmethod
def forward(ctx, helper: "UlyssesA2AHelper", x: torch.Tensor, mode: int) -> torch.Tensor: # type: ignore[override]
def forward(ctx, helper: "UlyssesA2AHelper", x: torch.Tensor, mode: int, chunked: bool,
blocks: int) -> torch.Tensor: # type: ignore[override]
ctx.helper = helper
ctx.mode = mode
return helper.run_armed(x, mode)
ctx.chunked = chunked
ctx.blocks = blocks
return helper.run_armed(x, mode, chunked, blocks)
@staticmethod
def backward(ctx, grad_output: torch.Tensor): # type: ignore[override]
# Same numel and dtype as the forward output, so the window is already
# sized for it; only contiguity needs restoring.
grad_input = ctx.helper.run_armed(grad_output.contiguous(), 1 - ctx.mode)
return None, grad_input, None
# Reuse the forward plan, even if later calls chose another plan or
# backward executes with grad tracking disabled.
grad_input = ctx.helper.run_armed(grad_output.contiguous(), 1 - ctx.mode, ctx.chunked, ctx.blocks)
return None, grad_input, None, None, None
class UlyssesA2AHelper:
@@ -76,8 +86,14 @@ class UlyssesA2AHelper:
self.pynccl_comm = pynccl_comm
self._handle: int | None = None
# Reuse storage, but exchange the current contract on every call. A
# rank-local cache hit cannot establish what peers are doing now.
self._local_contract = array("q", [0] * _CONTRACT_SIZE)
self._local_tensor = torch.frombuffer(self._local_contract, dtype=torch.int64)
self._gathered_tensor = torch.empty(world_size * _CONTRACT_SIZE, dtype=torch.int64, device="cpu")
self._nbytes = 0
self._disabled_reason: str | None = None
self._h3_tuning_available: bool | None = None
if world_size not in SUPPORTED_WORLD_SIZES:
self._disabled_reason = (f"world size {world_size} is not one of "
@@ -95,12 +111,12 @@ class UlyssesA2AHelper:
return int(getattr(comm, "value", comm))
def _can_attempt(self) -> tuple[bool, str]:
"""Whether this rank could use the fused path, without allocating anything."""
"""Check local capability only; the caller exchanges every rank's result."""
try:
from fastvideo_kernel import comm_ops
if not comm_ops.is_available():
return False, "fastvideo-kernel was built without the Ulysses a2a kernel"
if not comm_ops.lsa_covers_group(self._comm_ptr(), self.world_size):
elif not comm_ops.lsa_covers_group(self._comm_ptr(), self.world_size):
return False, "the group is not a load-store-accessible (NVLink) mesh"
except Exception as e: # noqa: BLE001
return False, f"backend unavailable ({type(e).__name__}: {e})"
@@ -108,7 +124,7 @@ class UlyssesA2AHelper:
def _agree(self, ok: bool) -> bool:
"""Reduce a local yes/no to a group-wide verdict: True only if all agree."""
vote = torch.tensor([1 if ok else 0], dtype=torch.int32)
vote = torch.tensor([1 if ok else 0], dtype=torch.int32, device="cpu")
dist.all_reduce(vote, op=dist.ReduceOp.MIN, group=self.cpu_group)
return bool(vote.item())
@@ -158,6 +174,20 @@ class UlyssesA2AHelper:
logger.warning("Ulysses partial-context cleanup failed", exc_info=True)
return self._agree(cleanup_ok)
def _execution_plan(self, x: torch.Tensor, mode: int) -> tuple[bool, int]:
"""Use the measured GB200/SP4 launch only for H3's bf16 head geometry."""
global_heads = x.shape[2] if mode == 0 else x.shape[2] * self.world_size
if self.world_size != 4 or x.dtype != torch.bfloat16 or global_heads != 56 or x.shape[3] != 128:
return False, 36
if self._h3_tuning_available is None:
from fastvideo_kernel import comm_ops
props = torch.cuda.get_device_properties(self.device)
self._h3_tuning_available = ("GB200" in props.name and props.major == 10 and props.minor == 0
and props.multi_processor_count >= 144
and getattr(comm_ops, "supports_tuned_launch", lambda: False)())
return (True, 144) if self._h3_tuning_available else (False, 36)
def _call_signature(self, x: torch.Tensor, scatter_dim: int, gather_dim: int) -> tuple[tuple[int, ...], str]:
"""Return a rank-comparable call contract and any local decline reason."""
mode = _MODE_FROM_DIMS.get((scatter_dim, gather_dim))
@@ -188,17 +218,33 @@ class UlyssesA2AHelper:
status, reason = 0, "sequence length is not divisible by the group"
nbytes = int(x.numel() * x.element_size())
chunked, blocks = False, 36
if status == 1 and nbytes:
assert mode is not None
chunked, blocks = self._execution_plan(x, mode)
window_bytes = nbytes // shape[0] if chunked else nbytes
if (status == 1 and chunked and torch.is_grad_enabled() and x.requires_grad
and window_bytes > H3_TRAINING_PLANE_LIMIT_BYTES):
policy = getattr(envs, "FASTVIDEO_ULYSSES_A2A_LONG_TRAINING", "auto")
if policy == "auto":
# Preserve the complete original long-training path, including
# its gather launch. Faster gathers alone did not avoid the
# allocation-pressure regression in the measured FSDP4 recipe.
chunked, blocks, window_bytes = False, 36, nbytes
elif policy != "chunked":
status, reason = 0, "unsupported FASTVIDEO_ULYSSES_A2A_LONG_TRAINING policy"
if status == 1 and nbytes == 0:
status, reason = 0, "input is empty"
elif status == 1 and nbytes > MAX_WINDOW_BYTES:
status, reason = 0, f"operand exceeds the {MAX_WINDOW_BYTES}-byte window cap"
elif status == 1 and window_bytes > MAX_WINDOW_BYTES:
status, reason = 0, f"operand window exceeds the {MAX_WINDOW_BYTES}-byte cap"
# status, armed, mode, dtype, B, S, H, D, bytes, capacity. Comparing the
# status, armed, mode, dtype, B, S, H, D, window bytes, capacity,
# chunked, blocks. Comparing the
# whole vector prevents equal-size but differently-shaped ranks from
# entering the fused kernel with incompatible address math. CUDA device
# ordinals are deliberately absent: rank-local ordinals normally differ.
signature = (status, int(self._handle is not None), -1 if mode is None else mode, dtype_code, *shape, nbytes,
self._nbytes)
signature = (status, int(self._handle is not None), -1 if mode is None else mode, dtype_code, *shape,
window_bytes, self._nbytes, int(chunked), blocks)
return signature, reason
def _agree_call(self, signature: tuple[int, ...]) -> tuple[bool, bool, bool]:
@@ -210,16 +256,14 @@ class UlyssesA2AHelper:
"""
# Host-side Gloo control keeps this agreement outside CUDA graph capture
# and avoids inserting a second NCCL collective ahead of the data path.
local = torch.tensor(signature, dtype=torch.int64)
gathered = torch.empty(self.world_size * local.numel(), dtype=local.dtype)
dist.all_gather_into_tensor(gathered, local, group=self.cpu_group)
contracts = gathered.view(self.world_size, local.numel())
identical = bool(torch.all(contracts == contracts[0]).item())
statuses = contracts[:, 0]
use_fused = identical and bool(torch.all(statuses == 1).item())
permanently_unavailable = bool(torch.any(statuses < 0).item())
lifecycle_consistent = (bool(torch.all(contracts[:, 1] == contracts[0, 1]).item())
and bool(torch.all(contracts[:, -1] == contracts[0, -1]).item()))
self._local_contract[:] = array("q", signature)
dist.all_gather_into_tensor(self._gathered_tensor, self._local_tensor, group=self.cpu_group)
values = self._gathered_tensor.tolist()
contracts = [values[start:start + _CONTRACT_SIZE] for start in range(0, len(values), _CONTRACT_SIZE)]
first = contracts[0]
use_fused = first[0] == 1 and all(contract == first for contract in contracts)
permanently_unavailable = any(contract[0] < 0 for contract in contracts)
lifecycle_consistent = all(contract[1] == first[1] and contract[9] == first[9] for contract in contracts)
return use_fused, permanently_unavailable, lifecycle_consistent
def _build(self, nbytes: int) -> bool:
@@ -327,7 +371,7 @@ class UlyssesA2AHelper:
# -- collective ----------------------------------------------------------
def run_armed(self, x: torch.Tensor, mode: int) -> torch.Tensor:
def run_armed(self, x: torch.Tensor, mode: int, chunked: bool, blocks: int) -> torch.Tensor:
"""Run one collective on an already-armed context."""
assert self._handle is not None, "run_armed called on an unarmed helper"
from fastvideo_kernel import comm_ops
@@ -340,7 +384,22 @@ class UlyssesA2AHelper:
B, S_global, H_local, D = x.shape
S_local, H = S_global // w, H_local * w
out = torch.empty(B, S_local, H, D, dtype=x.dtype, device=x.device)
comm_ops.all_to_all(self._handle, x, out, B, S_local, H, D, mode)
if chunked:
# Each copy completes on this stream before the next plane reuses
# the registered window. The full result owns its storage; saved
# activations never alias the reusable window.
for plane in range(B):
comm_ops.all_to_all(self._handle,
x[plane:plane + 1],
out[plane:plane + 1],
1,
S_local,
H,
D,
mode,
blocks=blocks)
else:
comm_ops.all_to_all(self._handle, x, out, B, S_local, H, D, mode)
return out
def try_all_to_all_4D(self, x: torch.Tensor, scatter_dim: int, gather_dim: int) -> torch.Tensor | None:
@@ -365,7 +424,7 @@ class UlyssesA2AHelper:
return None
mode = signature[2]
nbytes = signature[-2]
nbytes = signature[8]
if self._handle is None:
if not self._build(nbytes):
return None
@@ -376,7 +435,7 @@ class UlyssesA2AHelper:
if not self._build(nbytes):
return None
return _FusedUlyssesA2A.apply(self, x, mode)
return _FusedUlyssesA2A.apply(self, x, mode, bool(signature[10]), signature[11])
def maybe_create_helper(cpu_group: ProcessGroup | None, device_group: ProcessGroup | None, world_size: int,
@@ -404,9 +463,14 @@ def maybe_create_helper(cpu_group: ProcessGroup | None, device_group: ProcessGro
except Exception as e: # noqa: BLE001 - converted to a group verdict below
reason = f"helper construction failed ({type(e).__name__}: {e})"
vote = torch.tensor([int(helper is not None)], dtype=torch.int32)
dist.all_reduce(vote, op=dist.ReduceOp.MIN, group=cpu_group)
if not bool(vote.item()):
# Every rank reaches the same exchange, including configuration, constructor,
# and backend failures. LSA covers addressability, not single-host locality.
gathered: list[tuple[str, bool]] = [("", False)] * world_size
dist.all_gather_object(gathered, (socket.gethostname(), helper is not None), group=cpu_group)
hostnames = {hostname for hostname, _ in gathered}
if len(hostnames) != 1:
reason = f"ranks span multiple hosts: {sorted(hostnames)}"
if len(hostnames) != 1 or not all(ok for _, ok in gathered):
if dist.get_rank(cpu_group) == 0:
logger.info("Ulysses fused all-to-all unavailable: %s", reason or "a peer rank declined")
return None
+5
View File
@@ -28,6 +28,7 @@ if TYPE_CHECKING:
FASTVIDEO_VAE_PARALLEL_ENCODE: bool = False
FASTVIDEO_VAE_PARALLEL_DECODE_STRATEGY: str | None = None
FASTVIDEO_ULYSSES_A2A: str = "off"
FASTVIDEO_ULYSSES_A2A_LONG_TRAINING: str = "auto"
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "spawn"
FASTVIDEO_TARGET_DEVICE: str = "cuda"
MAX_JOBS: str | None = None
@@ -269,6 +270,10 @@ environment_variables: dict[str, Callable[[], Any]] = {
# mesh of 2/4/6/8 ranks in eager execution, else the NCCL path.
"FASTVIDEO_ULYSSES_A2A":
lambda: os.getenv("FASTVIDEO_ULYSSES_A2A", "off").strip().lower(),
# Keep the original long-training transport by default. Opt into bounded
# chunks after validating the training recipe's activation memory budget.
"FASTVIDEO_ULYSSES_A2A_LONG_TRAINING":
lambda: os.getenv("FASTVIDEO_ULYSSES_A2A_LONG_TRAINING", "auto").strip().lower(),
# Use dedicated multiprocess context for workers.
"FASTVIDEO_WORKER_MULTIPROC_METHOD":
@@ -0,0 +1,91 @@
# SPDX-License-Identifier: Apache-2.0
"""Run with torchrun on one NVLink host; compare revisions using identical argv.
Reports host-to-completion latency (including agreement), p50/p95 of rank-max
samples, for bf16 Wan2.1-14B and MiniMax-H3 5s/720p attention operands. No model
weights are required. Set FASTVIDEO_ULYSSES_A2A=off for the NCCL baseline.
"""
import argparse
import json
import os
import statistics
import time
from pathlib import Path
import torch
import torch.distributed as dist
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--iters", type=int, default=60)
parser.add_argument("--warmup", type=int, default=15)
parser.add_argument("--rounds", type=int, default=3)
parser.add_argument("--tag", required=True)
parser.add_argument("--output", type=Path, required=True)
parser.add_argument("--expect-fused", action="store_true")
args = parser.parse_args()
from fastvideo.distributed import cleanup_dist_env_and_memory, maybe_init_distributed_environment_and_model_parallel
from fastvideo.distributed.device_communicators.base_device_communicator import DeviceCommunicatorBase
from fastvideo.distributed.parallel_state import get_sp_group
world = int(os.environ["WORLD_SIZE"])
rank = int(os.environ["RANK"])
torch.cuda.set_device(int(os.environ["LOCAL_RANK"]))
device = torch.device("cuda", torch.cuda.current_device())
torch.manual_seed(2026 + rank)
maybe_init_distributed_environment_and_model_parallel(1, world)
comm = get_sp_group().device_communicator
records = []
try:
for name, sequence, heads in [("Wan2.1-14B", 75600, 40), ("MiniMax-H3", 37296, 56)]:
assert sequence % world == 0
scatter = torch.randn(3, sequence // world, heads, 128, dtype=torch.bfloat16, device=device)
gather = torch.randn(1, sequence, heads // world, 128, dtype=torch.bfloat16, device=device)
for operand, dims in [(scatter, (2, 1)), (gather, (1, 2))]:
actual = comm.all_to_all_4D(operand, *dims)
expected = DeviceCommunicatorBase.all_to_all_4D(comm, operand, *dims)
assert torch.equal(actual, expected), f"{name} {dims}: parity failed"
del actual, expected
armed = comm.ulysses_a2a is not None and comm.ulysses_a2a._handle is not None
if args.expect_fused:
assert armed, "benchmark requires all ranks to engage the fused kernel"
for repeat in range(args.rounds):
for operation in ("scatter", "gather", "layer"):
def run():
if operation in ("scatter", "layer"):
comm.all_to_all_4D(scatter, 2, 1)
if operation in ("gather", "layer"):
comm.all_to_all_4D(gather, 1, 2)
for _ in range(args.warmup):
run()
torch.cuda.synchronize()
samples = []
for _ in range(args.iters):
dist.barrier(group=get_sp_group().cpu_group)
start = time.perf_counter_ns()
run()
torch.cuda.synchronize()
samples.append((time.perf_counter_ns() - start) / 1000)
samples_tensor = torch.tensor(samples, dtype=torch.float64, device=device)
dist.all_reduce(samples_tensor, op=dist.ReduceOp.MAX)
maximums = samples_tensor.cpu().tolist()
record = dict(tag=args.tag, model=name, operation=operation, round=repeat,
world_size=world, fused=armed, p50_us=statistics.median(maximums),
p95_us=sorted(maximums)[int(0.95 * (len(maximums) - 1))],
rank_max_samples_us=maximums)
records.append(record)
if rank == 0:
print(json.dumps({k: v for k, v in record.items() if k != "rank_max_samples_us"}), flush=True)
del scatter, gather
if rank == 0:
args.output.write_text(json.dumps(records, indent=2) + "\n")
finally:
cleanup_dist_env_and_memory()
if __name__ == "__main__":
main()
@@ -0,0 +1,77 @@
# SPDX-License-Identifier: Apache-2.0
"""Explicit torchrun GPU gate for the native H3 transport policy (4x GB200)."""
import os
from unittest.mock import patch
import torch
import torch.distributed as dist
from fastvideo import envs
from fastvideo.distributed import cleanup_dist_env_and_memory, maybe_init_distributed_environment_and_model_parallel
from fastvideo.distributed.device_communicators.base_device_communicator import DeviceCommunicatorBase
from fastvideo.distributed.parallel_state import get_sp_group
from fastvideo.tests.distributed.test_ulysses_a2a_parity import _check_shape
def _check_native_shape(helper, batch, local_sequence, dtype, device, expected_plans):
# Numerical parity also passes when both sides use NCCL. Assert the actual
# helper executions, including the two backward collectives, separately.
with patch.object(helper, 'run_armed', wraps=helper.run_armed) as run:
_check_shape(batch, local_sequence, 56, 128, dtype, 4, device)
plans = [call.args[1:] for call in run.call_args_list]
assert plans == expected_plans, (batch, local_sequence, dtype, plans, expected_plans)
def main():
rank = int(os.environ['RANK'])
assert int(os.environ['WORLD_SIZE']) == 4
torch.cuda.set_device(int(os.environ['LOCAL_RANK']))
device = torch.device('cuda', int(os.environ['LOCAL_RANK']))
torch.manual_seed(20260906 + rank)
envs.FASTVIDEO_ULYSSES_A2A_LONG_TRAINING = 'auto'
maybe_init_distributed_environment_and_model_parallel(1, 4)
comm = get_sp_group().device_communicator
helper = comm.ulysses_a2a
assert helper is not None
for dtype in [torch.bfloat16, torch.float16, torch.float32]:
tuned = dtype == torch.bfloat16
plans = [(mode, tuned, 144 if tuned else 36) for mode in [0, 1, 0, 1, 1, 0]]
_check_native_shape(helper, 3, 64, dtype, device, plans)
for sequence in [32000, 128000, 250000]:
modes = [0, 1] if sequence == 250000 else [0, 1, 0, 1, 1, 0]
_check_native_shape(helper, 4, sequence // 4, torch.bfloat16, device,
[(mode, True, 144) for mode in modes])
assert helper._nbytes <= 1024**3
envs.FASTVIDEO_ULYSSES_A2A_LONG_TRAINING = 'chunked'
_check_native_shape(helper, 4, 250000 // 4, torch.bfloat16, device,
[(mode, True, 144) for mode in [0, 1, 0, 1, 1, 0]])
envs.FASTVIDEO_ULYSSES_A2A_LONG_TRAINING = 'auto'
x = torch.randn(3, 8000, 56, 128, dtype=torch.bfloat16, device=device, requires_grad=True)
first = comm.all_to_all_4D(x, 2, 1)
saved = first.detach().clone()
# A later call and inverse cannot overwrite an earlier autograd output.
second = comm.all_to_all_4D(x.detach() + 1, 2, 1)
comm.all_to_all_4D(second, 1, 2)
assert torch.equal(first, saved)
grad = torch.randn_like(first)
first.backward(grad)
expected = DeviceCommunicatorBase.all_to_all_4D(comm, grad, 1, 2)
assert torch.equal(x.grad, expected)
# A rank with an older launch capability must make the whole group decline.
helper._h3_tuning_available = rank != 0
assert helper.try_all_to_all_4D(x.detach(), 2, 1) is None
helper._h3_tuning_available = True
assert helper.try_all_to_all_4D(x.detach(), 2, 1) is not None
contract, _ = helper._call_signature(x.detach(), 2, 1)
assert contract[10:] == (1, 144)
dist.barrier()
if rank == 0:
print('NATIVE_H3_OK dtypes=3 lengths=32000,128000,250000 exact_gradients=True ownership=True recovery=True '
'execution_plans=True',
flush=True)
cleanup_dist_env_and_memory()
if __name__ == '__main__':
main()
@@ -9,6 +9,7 @@ import socket
import subprocess
import sys
from pathlib import Path
from unittest.mock import patch
import pytest
import torch
@@ -38,6 +39,87 @@ FAULT_CASES = [
(2, 0, "lifecycle"),
(4, 3, "lifecycle"),
]
FAULT_CASES += [(world, fault_rank, stage)
for world, fault_rank in [(2, 1), (4, 0)]
for stage in ("constructor", "backend", "pynccl", "hostname", "late_layout", "late_capture",
"late_configuration", "late_capacity", "late_lifecycle", "cached_shape", "cached_mode",
"decline_then_recover", "late_real_capture", "late_mixed_capture")]
def _check_after_warmup(helper, rank: int, fault_rank: int, stage: str, device: torch.device) -> None:
"""Exercise divergence after an armed signature has already succeeded twice."""
from fastvideo.distributed.device_communicators import ulysses_a2a
assert helper is not None, "this regression must exercise an available fused helper"
x = torch.randn(3, 64, 8, 64, device=device, dtype=torch.bfloat16)
if stage == "decline_then_recover":
# Populate a declined contract before any window or successful call.
operand = torch.stack((x, x), dim=-1)[..., 0] if rank == fault_rank else x
assert helper.try_all_to_all_4D(operand, 2, 1) is None
assert helper._handle is None
for _ in range(3):
y = helper.try_all_to_all_4D(x, 2, 1)
assert y is not None
assert torch.equal(helper.try_all_to_all_4D(y, 1, 2), x)
alt = x.reshape(3, 32, 16, 64)
for _ in range(2):
assert helper.try_all_to_all_4D(alt, 2, 1) is not None
if stage in ("late_real_capture", "late_mixed_capture"):
capture_stream = torch.cuda.Stream(device=device)
graph = torch.cuda.CUDAGraph()
# Synchronize warmup before starting the real, default-global capture.
torch.cuda.synchronize(device)
if stage == "late_real_capture" or rank == fault_rank:
with torch.cuda.graph(graph, stream=capture_stream):
graph_output = x + 1
assert helper.try_all_to_all_4D(x, 2, 1) is None
graph.replay()
torch.cuda.synchronize(device)
assert torch.equal(graph_output, x + 1)
else:
assert helper.try_all_to_all_4D(x, 2, 1) is None
recovered = helper.try_all_to_all_4D(x, 2, 1)
assert recovered is not None
assert torch.equal(helper.try_all_to_all_4D(recovered, 1, 2), x)
print(f"RANK_DONE rank={rank} recovered=True", flush=True)
if rank == 0:
print(f"ALL_RANKS_COMPLETED world={helper.world_size} stage={stage}", flush=True)
return
dims = (2, 1)
operand = x
if rank == fault_rank:
if stage in ("late_layout", "decline_then_recover"):
operand = torch.stack((x, x), dim=-1)[..., 0]
elif stage == "cached_shape":
operand = alt
elif stage == "cached_mode":
operand, dims = y, (1, 2)
elif stage == "late_lifecycle":
helper._nbytes += 1
# Both the rank with the changed input and the ranks with familiar inputs
# must decline together. Calling NCCL with mismatched shapes/modes would
# itself be invalid, so check the helper result before the fallback path.
with patch.object(torch.cuda, "is_current_stream_capturing",
return_value=rank == fault_rank and stage == "late_capture"), \
patch.object(ulysses_a2a, "is_enabled",
return_value=not (rank == fault_rank and stage == "late_configuration")), \
patch.object(ulysses_a2a, "MAX_WINDOW_BYTES",
1 if rank == fault_rank and stage == "late_capacity" else 1024**3):
assert helper.try_all_to_all_4D(operand, *dims) is None, "a peer entered the fused path alone"
if stage == "late_lifecycle":
assert helper._handle is None and helper._disabled_reason is not None
else:
# A transient decline must not poison a signature when peers recover.
recovered = helper.try_all_to_all_4D(x, 2, 1)
assert recovered is not None
assert torch.equal(helper.try_all_to_all_4D(recovered, 1, 2), x)
print(f"RANK_DONE rank={rank} recovered=True", flush=True)
if rank == 0:
print(f"ALL_RANKS_COMPLETED world={helper.world_size} stage={stage}", flush=True)
def _worker() -> None:
@@ -66,6 +148,29 @@ def _worker() -> None:
lambda self: (False, "injected capability failure"))
elif rank == fault_rank and fault_stage == "configuration":
os.environ["FASTVIDEO_ULYSSES_A2A"] = "off"
elif rank == fault_rank and fault_stage == "constructor":
from fastvideo.distributed.device_communicators.ulysses_a2a import UlyssesA2AHelper
def _fail_constructor(*args, **kwargs):
raise RuntimeError("injected helper constructor failure")
UlyssesA2AHelper.__init__ = _fail_constructor
elif rank == fault_rank and fault_stage == "backend":
from fastvideo_kernel import comm_ops
comm_ops.is_available = lambda: False
elif rank == fault_rank and fault_stage == "pynccl":
from fastvideo.distributed.device_communicators.pynccl import PyNcclCommunicator
real_init = PyNcclCommunicator.__init__
def _disable_pynccl(self, *args, **kwargs):
real_init(self, *args, **kwargs)
self.disabled = True
PyNcclCommunicator.__init__ = _disable_pynccl
elif rank == fault_rank and fault_stage == "hostname":
socket.gethostname = lambda: "injected-other-host"
maybe_init_distributed_environment_and_model_parallel(1, world)
communicator = get_sp_group().device_communicator
@@ -73,11 +178,15 @@ def _worker() -> None:
try:
if helper is None:
if fault_stage not in ("capability", "configuration"):
if fault_stage not in ("capability", "configuration", "constructor", "backend", "pynccl", "hostname"):
if rank == 0:
print("UNAVAILABLE helper not created", flush=True)
return
if fault_stage.startswith("late_") or fault_stage in ("cached_shape", "cached_mode", "decline_then_recover"):
_check_after_warmup(helper, rank, fault_rank, fault_stage, device)
return
if helper is not None and rank == fault_rank and fault_stage == "allocation":
helper._allocate = lambda nbytes: (_ for _ in ()).throw(
RuntimeError(f"injected allocation failure for {nbytes} bytes"))
@@ -194,8 +303,11 @@ def test_rank_local_setup_failure_falls_back_group_wide(world: int, fault_rank:
pytest.skip(process.stdout.strip())
assert process.returncode == 0 and "ALL_RANKS_COMPLETED" in process.stdout, (
f"stdout:\n{process.stdout}\nstderr:\n{process.stderr[-6000:]}")
armed = dict(re.findall(r"RANK_DONE rank=(\d+) armed=(True|False)", process.stdout))
assert len(armed) == world and set(armed.values()) == {"False"}
if fault_stage.startswith("late_") or fault_stage in ("cached_shape", "cached_mode", "decline_then_recover"):
assert len(re.findall(r"RANK_DONE rank=\d+ recovered=True", process.stdout)) == world
else:
armed = dict(re.findall(r"RANK_DONE rank=(\d+) armed=(True|False)", process.stdout))
assert len(armed) == world and set(armed.values()) == {"False"}
if __name__ == "__main__" and "--worker" in sys.argv:
@@ -0,0 +1,115 @@
# SPDX-License-Identifier: Apache-2.0
"""Policy, rank agreement and saved backward-plan regressions without a GPU."""
from types import SimpleNamespace
from unittest.mock import Mock, patch
import pytest
import torch
from fastvideo.distributed.device_communicators import ulysses_a2a as ulysses
@pytest.fixture(autouse=True)
def _default_long_training_policy(monkeypatch):
# A caller may run these tests from a shell configured for long training.
monkeypatch.setattr(ulysses.envs, 'FASTVIDEO_ULYSSES_A2A_LONG_TRAINING', 'auto')
def _helper():
return ulysses.UlyssesA2AHelper(object(), object(), 4, torch.device('cuda:0'), object())
def _operand(planes, global_sequence, mode, requires_grad):
shape = (planes, global_sequence // 4, 56, 128) if mode == 0 else (planes, global_sequence, 14, 128)
x = Mock(shape=shape, dtype=torch.bfloat16, device=torch.device('cuda:0'), is_cuda=True,
requires_grad=requires_grad)
x.dim.return_value = 4
x.numel.return_value = planes * global_sequence * 14 * 128
x.element_size.return_value = 2
x.is_contiguous.return_value = True
return x
@pytest.mark.parametrize('sequence,large_pack', [(32000, False), (64000, False), (128000, False),
(149796, False), (149800, True), (250000, True)])
@pytest.mark.parametrize('planes,mode,training', [(3, 0, False), (4, 0, True), (1, 1, True), (4, 1, True)])
def test_h3_policy_caps_planes_and_preserves_long_training_default(sequence, large_pack, planes, mode, training):
helper = _helper()
x = _operand(planes, sequence, mode, training)
with patch.object(helper, '_execution_plan', return_value=(True, 144)), \
patch.object(ulysses, 'is_enabled', return_value=True), \
patch('torch.cuda.is_current_stream_capturing', return_value=False), torch.set_grad_enabled(training):
signature, _ = helper._call_signature(x, *((2, 1) if mode == 0 else (1, 2)))
declined = training and planes > 1 and large_pack
assert signature[0] == (0 if declined else 1)
original_plan = training and large_pack
assert signature[8] == sequence * 14 * 128 * 2 * (planes if original_plan else 1)
assert signature[10:] == ((0, 36) if original_plan else (1, 144))
def test_long_training_chunk_opt_in_keeps_a_bounded_window():
helper = _helper()
x = _operand(4, 250000, 0, True)
with patch.object(helper, '_execution_plan', return_value=(True, 144)), \
patch.object(ulysses, 'is_enabled', return_value=True), \
patch.object(ulysses.envs, 'FASTVIDEO_ULYSSES_A2A_LONG_TRAINING', 'chunked', create=True), \
patch('torch.cuda.is_current_stream_capturing', return_value=False), torch.enable_grad():
signature, _ = helper._call_signature(x, 2, 1)
assert signature[0] == 1
assert signature[8] == 896000000
assert signature[8] <= ulysses.MAX_WINDOW_BYTES
assert signature[10:] == (1, 144)
def test_older_wheel_keeps_original_launch_and_capacity_policy():
helper = _helper()
x = _operand(3, 250000, 0, False)
from fastvideo_kernel import comm_ops
props = SimpleNamespace(name='NVIDIA GB200', major=10, minor=0, multi_processor_count=152)
with patch('torch.cuda.get_device_properties', return_value=props), \
patch.object(comm_ops, 'supports_tuned_launch', return_value=False, create=True), \
patch.object(ulysses, 'is_enabled', return_value=True), \
patch('torch.cuda.is_current_stream_capturing', return_value=False):
signature, _ = helper._call_signature(x, 2, 1)
assert signature[0] == 0
assert signature[8] == x.numel() * 2
assert signature[10:] == (0, 36)
def test_rank_disagreement_on_launch_or_chunking_declines_and_recovers():
helper = _helper()
signature = (1, 0, 0, 2, 4, 32000, 56, 128, 458752000, 0, 1, 144)
for changed_index, value in [(10, 0), (11, 36), (0, 0)]:
peer = list(signature)
peer[changed_index] = value
def gather(output, local, **kwargs):
output.copy_(torch.tensor([*signature, *peer, *signature, *signature]))
with patch.object(ulysses.dist, 'all_gather_into_tensor', side_effect=gather):
assert helper._agree_call(signature) == (False, False, True)
def unanimous(output, local, **kwargs):
output.copy_(local.repeat(4))
with patch.object(ulysses.dist, 'all_gather_into_tensor', side_effect=unanimous):
assert helper._agree_call(signature) == (True, False, True)
def test_backward_uses_its_own_forward_plan_after_another_call():
calls = []
class Helper:
def run_armed(self, x, mode, chunked, blocks):
calls.append((mode, chunked, blocks))
return x.clone()
helper = Helper()
x = torch.randn(4, requires_grad=True)
first = ulysses._FusedUlyssesA2A.apply(helper, x, 0, True, 144)
second = ulysses._FusedUlyssesA2A.apply(helper, x, 1, False, 36)
first.sum().backward()
assert calls == [(0, True, 144), (1, False, 36), (1, True, 144)]
assert torch.equal(first, second)
assert first.data_ptr() != second.data_ptr()
assert torch.equal(x.grad, torch.ones_like(x))
@@ -0,0 +1,36 @@
# SPDX-License-Identifier: Apache-2.0
"""Load a separately built communication extension for development validation."""
import importlib.util
import runpy
import sys
from pathlib import Path
import torch # noqa: F401 - load the extension's libtorch dependencies first
def load_module(name: str, path: Path):
spec = importlib.util.spec_from_file_location(name, path)
if spec is None or spec.loader is None:
raise ImportError(str(path))
module = importlib.util.module_from_spec(spec)
sys.modules[name] = module
spec.loader.exec_module(module)
return module
def install_native_ops(source: Path, extension: Path):
import fastvideo_kernel
native = load_module('ulysses_native', extension)
wrapper = load_module('fastvideo_kernel.comm_ops',
source / 'fastvideo-kernel/python/fastvideo_kernel/comm_ops.py')
wrapper._ops = native
fastvideo_kernel.comm_ops = wrapper
return wrapper
if __name__ == '__main__':
source, extension, target = map(Path, sys.argv[1:4])
install_native_ops(source, extension)
sys.argv = [str(target), *sys.argv[4:]]
runpy.run_path(str(target), run_name='__main__')