Compare commits

...
Author SHA1 Message Date
SolitaryThinker fc51550920 [perf] Probe H3 chunked all-to-all with direct output slices 2026-09-06 01:58:09 +00:00
SolitaryThinker 942e8b8404 [perf] Study H3 long-sequence window limits and bounded chunking 2026-09-06 01:50:22 +00:00
SolitaryThinker c548e5834e [perf] Build the pinned sparse kernel for the MFU experiment 2026-09-06 01:20:21 +00:00
SolitaryThinker 368f6b8891 [perf] Compare launch geometry in dense and sparse training blocks 2026-09-06 01:13:37 +00:00
SolitaryThinker e55934b62e [perf] Add reproducible Ulysses MFU bottleneck experiments 2026-09-06 01:08:01 +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
5 changed files with 782 additions and 19 deletions
@@ -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
@@ -76,6 +79,11 @@ 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] * 10)
self._local_tensor = torch.frombuffer(self._local_contract, dtype=torch.int64)
self._gathered_tensor = torch.empty(world_size * 10, dtype=torch.int64, device="cpu")
self._nbytes = 0
self._disabled_reason: str | None = None
@@ -95,12 +103,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 +116,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())
@@ -210,16 +218,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 + 10] for start in range(0, len(values), 10)]
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[-1] == first[-1] for contract in contracts)
return use_fused, permanently_unavailable, lifecycle_consistent
def _build(self, nbytes: int) -> bool:
@@ -404,9 +410,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
@@ -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,142 @@
# SPDX-License-Identifier: Apache-2.0
"""Build and measure an isolated launch-configuration variant of the real kernel.
Only the CTA/thread counts and optional copy-out suppression differ. Suppressing
copy-out measures the movement kernel, not a usable tensor-returning operation.
The installed kernel and production dispatch are never replaced.
"""
import argparse
import importlib.util
import json
import os
import statistics
import time
from pathlib import Path
import torch
import torch.distributed as dist
def build(directory):
from torch.utils.cpp_extension import load
repo = Path(__file__).resolve().parents[3]
header = (repo / 'fastvideo-kernel/include/comm/ulysses_all_to_all.cuh').read_text()
source = (repo / 'fastvideo-kernel/csrc/comm/ulysses_all_to_all.cu').read_text()
header = header.replace('constexpr int kMaxBlocks = 36;', 'constexpr int kMaxBlocks = 144;')
source = source.replace('int64_t H, int64_t D, int64_t mode) {',
'int64_t H, int64_t D, int64_t mode, int probe_blocks, int probe_threads, bool copy_out) {')
source = source.replace('std::min<int64_t>(fi::kMaxBlocks, num_rows)', 'std::min<int64_t>(probe_blocks, num_rows)')
source = source.replace('const int threads = fi::kUlyssesThreads;',
'TORCH_CHECK(probe_blocks > 0 && probe_blocks <= fi::kMaxBlocks, "bad block count");\n'
' TORCH_CHECK(probe_threads == 128 || probe_threads == 256 || probe_threads == 512, '
'"bad thread count");\n const int threads = probe_threads;')
source = source.replace('// Copy this rank\'s completed result out of the window.',
'if (!copy_out) return;\n // Copy this rank\'s completed result out of the window.')
source += '\nPYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { register_ulysses_a2a(m); }\n'
(directory / 'comm').mkdir(parents=True, exist_ok=True)
(directory / 'comm/ulysses_all_to_all.cuh').write_text(header)
(directory / 'probe.cu').write_text(source)
nccl = Path(next(iter(importlib.util.find_spec('nvidia.nccl').submodule_search_locations)))
os.environ['TORCH_CUDA_ARCH_LIST'] = '10.0a'
os.environ['MAX_JOBS'] = '2'
return load(name='ulysses_launch_probe', sources=[str(directory / 'probe.cu')],
extra_include_paths=[str(directory), str(nccl / 'include')],
extra_cuda_cflags=['-O3', '-std=c++17'],
extra_ldflags=[str(nccl / 'lib/libnccl.so.2')],
build_directory=str(directory), verbose=True)
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument('--build-dir', type=Path, required=True)
parser.add_argument('--build-only', action='store_true')
parser.add_argument('--sparse-build-only', action='store_true')
parser.add_argument('--output', type=Path)
args = parser.parse_args()
args.build_dir.mkdir(parents=True, exist_ok=True)
if args.sparse_build_only:
from torch.utils.cpp_extension import load
repo = Path(__file__).resolve().parents[3]
source_dir = repo / 'fastvideo-kernel/csrc/attention'
source = (source_dir / 'block_sparse_sm100a.cu').read_text()
source += ('\nPYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { '
'm.def("block_sparse_sm100a_fwd", &block_sparse_sm100a_fwd); }\n')
generated = args.build_dir / 'sparse.cu'
generated.write_text(source)
os.environ['TORCH_CUDA_ARCH_LIST'] = '10.0a'
os.environ['MAX_JOBS'] = '2'
load(name='ulysses_sparse_probe', sources=[str(generated)],
extra_include_paths=[str(source_dir)],
extra_cuda_cflags=['-O3', '-std=c++17', '-DVSA_BHSD=true'],
extra_ldflags=['-L/usr/local/cuda/lib64/stubs', '-lcuda'],
build_directory=str(args.build_dir), verbose=True)
return
if args.build_only:
build(args.build_dir)
return
import sys
sys.path.insert(0, str(args.build_dir))
import ulysses_launch_probe as probe
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
rank, world = int(os.environ['RANK']), int(os.environ['WORLD_SIZE'])
torch.cuda.set_device(int(os.environ['LOCAL_RANK']))
torch.manual_seed(20260906 + rank)
maybe_init_distributed_environment_and_model_parallel(1, world)
comm = get_sp_group().device_communicator
records = []
try:
for model, sequence, heads in [('small', 8192, 40), ('Wan', 75600, 40), ('H3', 37296, 56)]:
handle = probe.allocate_ulysses_a2a(3 * sequence // world * heads * 128 * 2,
rank, world, torch.cuda.current_device())
probe.register_ulysses_a2a_window(handle, comm.ulysses_a2a._comm_ptr())
probe.create_ulysses_a2a_dev_comm(handle)
for mode in (0, 1):
shape = (3, sequence // world, heads, 128) if mode == 0 else (1, sequence, heads // world, 128)
x = torch.randn(shape, device='cuda', dtype=torch.bfloat16)
dims = (2, 1) if mode == 0 else (1, 2)
expected = DeviceCommunicatorBase.all_to_all_4D(comm, x, *dims)
out = torch.empty_like(expected)
for blocks in (12, 18, 36, 72, 144):
for threads in (128, 256, 512):
def run(copy_out=True):
probe.ulysses_a2a(handle, x, out, shape[0], sequence // world,
heads, 128, mode, blocks, threads, copy_out)
run()
torch.cuda.synchronize()
assert torch.equal(out, expected), (model, mode, blocks, threads)
for copy_out in (True, False):
for _ in range(5):
run(copy_out)
torch.cuda.synchronize()
samples = []
for _ in range(3):
dist.barrier(group=get_sp_group().cpu_group)
start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True)
start.record()
for _ in range(20):
run(copy_out)
end.record()
torch.cuda.synchronize()
samples.append(start.elapsed_time(end) * 1000 / 20)
tensor = torch.tensor(samples, device='cuda', dtype=torch.float64)
dist.all_reduce(tensor, op=dist.ReduceOp.MAX)
samples = tensor.cpu().tolist()
record = dict(model=model, mode=mode, blocks=blocks, threads=threads,
copy_out=copy_out, p50_us=statistics.median(samples), samples_us=samples,
bytes=x.numel() * x.element_size(), parity=True)
records.append(record)
if rank == 0:
print(json.dumps(record), flush=True)
args.output.write_text(json.dumps(records, indent=2) + '\n')
del x, out, expected
torch.cuda.synchronize()
probe.dispose_ulysses_a2a(handle)
finally:
cleanup_dist_env_and_memory()
if __name__ == '__main__':
main()
@@ -0,0 +1,407 @@
# SPDX-License-Identifier: Apache-2.0
"""Exploratory GB200/SP4 study; synthetic operands and transformer blocks.
Prepared mode is a benchmark-only upper bound: every rank follows this script's
identical fixed contract and warms the largest window before timing. It is not
a production replacement for the helper's dynamic collective fallback protocol.
Run using torchrun; results include all rank-max samples and parity checks.
"""
import argparse
import gc
import json
import os
import statistics
import time
from pathlib import Path
import torch
import torch.distributed as dist
import torch.nn.functional as F
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument('--section', choices=['collectives', 'blocks'], required=True)
parser.add_argument('--output', type=Path, required=True)
parser.add_argument('--iters', type=int, default=20)
parser.add_argument('--rounds', type=int, default=3)
parser.add_argument('--warmup', type=int, default=5)
parser.add_argument('--probe-dir', type=Path)
parser.add_argument('--sparse', action='store_true')
parser.add_argument('--sparse-probe-dir', type=Path)
parser.add_argument('--models', nargs='+', default=['small', 'Wan', 'H3'])
parser.add_argument('--paired', action='store_true')
parser.add_argument('--h3-lengths', nargs='+', type=int)
parser.add_argument('--probe-window-gib', type=int, default=1)
parser.add_argument('--routes', nargs='+')
parser.add_argument('--chunked', action='store_true')
parser.add_argument('--packed-count', type=int, default=3)
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.device_communicators.ulysses_a2a import _FusedUlyssesA2A, UlyssesA2AHelper
import fastvideo.distributed.device_communicators.ulysses_a2a as ulysses_module
from fastvideo.distributed.parallel_state import get_sp_group
rank = int(os.environ['RANK'])
world = int(os.environ['WORLD_SIZE'])
torch.cuda.set_device(int(os.environ['LOCAL_RANK']))
torch.manual_seed(20260906 + rank)
maybe_init_distributed_environment_and_model_parallel(1, world)
group = get_sp_group()
comm = group.device_communicator
helper = comm.ulysses_a2a
assert helper is not None
records = []
engagement = {}
probe_helpers = {}
if args.probe_dir:
import sys
sys.path.insert(0, str(args.probe_dir))
import ulysses_launch_probe as probe_ops
class ProbeHelper(UlyssesA2AHelper):
def _call_signature(self, x, scatter_dim, gather_dim):
# Experimental per-helper cap; this script has one host thread.
original = ulysses_module.MAX_WINDOW_BYTES
ulysses_module.MAX_WINDOW_BYTES = args.probe_window_gib * 1024**3
try:
return super()._call_signature(x, scatter_dim, gather_dim)
finally:
ulysses_module.MAX_WINDOW_BYTES = original
def _allocate(self, nbytes):
return probe_ops.allocate_ulysses_a2a(nbytes, rank, world, torch.cuda.current_device())
def _register_window(self, handle):
probe_ops.register_ulysses_a2a_window(handle, self._comm_ptr())
def _create_dev_comm(self, handle):
probe_ops.create_ulysses_a2a_dev_comm(handle)
def _dispose(self, handle, *, synchronize):
if synchronize:
torch.cuda.synchronize()
probe_ops.dispose_ulysses_a2a(handle)
def run_armed(self, x, mode):
b, s, h, d = x.shape
local_s, global_h = (s, h) if mode == 0 else (s // world, h * world)
shape = (b, s * world, h // world, d) if mode == 0 else (b, s // world, h * world, d)
out = torch.empty(shape, device=x.device, dtype=x.dtype)
probe_ops.ulysses_a2a(self._handle, x, out, b, local_s, global_h, d, mode,
self.probe_blocks, 512, True)
return out
for blocks in (36, 72, 144):
candidate = ProbeHelper(helper.cpu_group, helper.device_group, world, helper.device, helper.pynccl_comm)
candidate.probe_blocks = blocks
probe_helpers[f'probe{blocks}'] = candidate
if args.chunked:
candidate = ProbeHelper(helper.cpu_group, helper.device_group, world, helper.device, helper.pynccl_comm)
candidate.probe_blocks = 144
probe_helpers['chunk144'] = candidate
class ChunkIntoHelper(ProbeHelper):
"""Benchmark-only chunked protocol with one owned final output.
Full-call agreement fixes the chunk count/order on all ranks;
inherited setup votes protect allocation/registration. The real
registered capacity is one batch plane and is advertised as such.
"""
def try_all_to_all_4D(self, x, scatter_dim, gather_dim):
if self._disabled_reason is not None or torch.compiler.is_compiling():
return None
signature, reason = self._call_signature(x, scatter_dim, gather_dim)
fused, permanent, lifecycle = self._agree_call(signature)
if not fused:
if not lifecycle:
self.close()
self._disable('inconsistent chunk-window lifecycle')
if permanent:
self._disable(reason or 'peer declined chunked operation')
return None
slot_bytes = signature[-2] // x.shape[0]
assert slot_bytes <= 1024**3, 'one batch plane must fit the bounded window'
if self._handle is None:
if not self._build(slot_bytes):
return None
elif slot_bytes > self._nbytes:
if not self.close() or not self._build(slot_bytes):
return None
return _FusedUlyssesA2A.apply(self, x, signature[2])
def run_armed(self, x, mode):
b, s, h, d = x.shape
local_s, global_h = (s, h) if mode == 0 else (s // world, h * world)
shape = (b, s * world, h // world, d) if mode == 0 else (b, s // world, h * world, d)
out = torch.empty(shape, device=x.device, dtype=x.dtype)
for plane in range(b):
probe_ops.ulysses_a2a(self._handle, x[plane:plane + 1], out[plane:plane + 1],
1, local_s, global_h, d, mode, 144, 512, True)
return out
candidate = ChunkIntoHelper(helper.cpu_group, helper.device_group, world, helper.device, helper.pynccl_comm)
candidate.probe_blocks = 144
probe_helpers['chunk_into144'] = candidate
routes = ['nccl', 'safe', 'prepared', *probe_helpers]
if args.routes:
assert set(args.routes) <= set(routes)
routes = args.routes
if args.h3_lengths:
assert 'prepared' not in routes, 'legacy prepared control cannot cover the 1 GiB cap transition'
def emit(record):
record.update(sequence=sequence, heads=heads, packed_count=args.packed_count,
probe_window_gib=args.probe_window_gib)
records.append(record)
if rank == 0:
args.output.write_text(json.dumps(records, indent=2) + '\n')
print(json.dumps({k: v for k, v in record.items() if 'samples' not in k}), flush=True)
def measure(fn, metadata, *, iterations=None, repeats=1, round_index=None):
count = iterations or args.iters
for _ in range(args.warmup):
fn()
torch.cuda.synchronize()
for repeat in (range(args.rounds) if round_index is None else [round_index]):
wall, gpu = [], []
before = dict(engagement)
torch.cuda.reset_peak_memory_stats()
for _ in range(count):
dist.barrier(group=group.cpu_group)
start, end = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True)
begin = time.perf_counter_ns()
start.record()
for _ in range(repeats):
fn()
end.record()
torch.cuda.synchronize()
wall.append((time.perf_counter_ns() - begin) / 1000 / repeats)
gpu.append(start.elapsed_time(end) * 1000 / repeats)
values = torch.tensor([wall, gpu], dtype=torch.float64, device='cuda')
dist.all_reduce(values, op=dist.ReduceOp.MAX)
wall, gpu = values.cpu().tolist()
emit(dict(**metadata, round=repeat, repeats=repeats, world=world,
wall_p50_us=statistics.median(wall), gpu_p50_us=statistics.median(gpu),
forward_engagement={k: v - before.get(k, 0) for k, v in engagement.items() if v > before.get(k, 0)},
torch_peak_allocated_gib=torch.cuda.max_memory_allocated() / 1024**3,
registered_windows_gib=sum(h._nbytes for h in [helper, *probe_helpers.values()]) / 1024**3,
wall_rank_max_samples_us=wall, gpu_rank_max_samples_us=gpu))
def a2a(x, mode, route):
dims = (2, 1) if mode == 0 else (1, 2)
if route == 'nccl':
key = f'{route}:{mode}:nccl'
engagement[key] = engagement.get(key, 0) + 1
return DeviceCommunicatorBase.all_to_all_4D(comm, x, *dims)
if route == 'safe':
result = helper.try_all_to_all_4D(x, *dims)
key = f'{route}:{mode}:{"fused" if result is not None else "nccl"}'
engagement[key] = engagement.get(key, 0) + 1
return result if result is not None else DeviceCommunicatorBase.all_to_all_4D(comm, x, *dims)
if route == 'chunk144':
candidate = probe_helpers[route]
# Agree the whole contract before ranks enter a variable number of chunks.
signature, _ = candidate._call_signature(x, *dims)
if not candidate._agree_call(signature)[0]:
return DeviceCommunicatorBase.all_to_all_4D(comm, x, *dims)
results = []
for chunk in x.split(1, dim=0):
assert chunk.numel() * chunk.element_size() <= 1024**3
result = candidate.try_all_to_all_4D(chunk, *dims)
key = f'{route}:{mode}:{"fused" if result is not None else "nccl"}'
engagement[key] = engagement.get(key, 0) + 1
results.append(result if result is not None else DeviceCommunicatorBase.all_to_all_4D(comm, chunk, *dims))
return results[0] if len(results) == 1 else torch.cat(results, dim=0)
if route in probe_helpers:
result = probe_helpers[route].try_all_to_all_4D(x, *dims)
assert result is not None, f'{route} unexpectedly declined'
key = f'{route}:{mode}:fused'
engagement[key] = engagement.get(key, 0) + 1
return result
assert route == 'prepared'
return _FusedUlyssesA2A.apply(helper, x, mode)
def prepare(operands):
# Use the full existing protocol outside the measured region.
for x, mode in sorted(operands, key=lambda item: item[0].numel(), reverse=True):
actual = a2a(x, mode, 'safe')
expected = a2a(x, mode, 'nccl')
assert torch.equal(actual, expected)
for route in [r for r in routes if r in probe_helpers]:
assert torch.equal(a2a(x, mode, route), expected)
if 'prepared' in routes:
assert helper._handle is not None
torch.cuda.synchronize()
try:
workloads = [('small', 8192, 40), ('Wan', 75600, 40), ('H3', 37296, 56)]
workloads = [w for w in workloads if w[0] in args.models]
if args.h3_lengths:
workloads = [(f'H3-{length}', length, 56) for length in args.h3_lengths]
if args.section == 'collectives':
for model, sequence, heads in workloads:
x = torch.randn(args.packed_count, sequence // world, heads, 128, device='cuda', dtype=torch.bfloat16)
y = torch.randn(1, sequence, heads // world, 128, device='cuda', dtype=torch.bfloat16)
prepare([(x, 0), (y, 1)])
signature, _ = helper._call_signature(x, 2, 1)
measure(lambda: helper._agree_call(signature), dict(model=model, operation='agreement', route='safe'))
for mode, operand in [(0, x), (1, y)]:
operation = 'scatter' if mode == 0 else 'gather'
for route in routes:
measure(lambda: a2a(operand, mode, route), dict(model=model, operation=operation, route=route))
dst = torch.empty_like(operand)
measure(lambda: dst.copy_(operand), dict(model=model, operation=operation + '_copy', route='copy'))
for route in routes:
def pair():
a2a(x, 0, route)
a2a(y, 1, route)
measure(pair, dict(model=model, operation='pair_streamed', route=route),
iterations=max(3, args.iters // 4), repeats=20)
x.requires_grad_(True)
y.requires_grad_(True)
dx = torch.randn(args.packed_count, sequence, heads // world, 128, device='cuda', dtype=torch.bfloat16)
dy = torch.randn(1, sequence // world, heads, 128, device='cuda', dtype=torch.bfloat16)
def training_pair():
ox = a2a(x, 0, route)
oy = a2a(y, 1, route)
return torch.autograd.grad((ox, oy), (x, y), (dx, dy))
if args.h3_lengths:
expected_grads = torch.autograd.grad((a2a(x, 0, 'nccl'), a2a(y, 1, 'nccl')),
(x, y), (dx, dy))
actual_grads = training_pair()
assert all(torch.equal(a, b) for a, b in zip(actual_grads, expected_grads))
del expected_grads, actual_grads
measure(training_pair, dict(model=model, operation='pair_fwd_bwd', route=route),
iterations=max(3, args.iters // 4), repeats=5)
x.requires_grad_(False)
y.requires_grad_(False)
del dx, dy
if 'prepared' not in routes:
del x, y, dst
gc.collect()
torch.cuda.empty_cache()
continue
# All ranks collectively capture the same already-prepared calls.
capture_stream = torch.cuda.Stream()
capture_stream.wait_stream(torch.cuda.current_stream())
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph, stream=capture_stream):
captured_x = a2a(x, 0, 'prepared')
captured_y = a2a(y, 1, 'prepared')
torch.cuda.current_stream().wait_stream(capture_stream)
graph.replay()
torch.cuda.synchronize()
assert torch.equal(captured_x, a2a(x, 0, 'nccl'))
assert torch.equal(captured_y, a2a(y, 1, 'nccl'))
measure(graph.replay, dict(model=model, operation='pair_graph', route='prepared'),
iterations=max(3, args.iters // 4), repeats=20)
if model != 'small':
with torch.profiler.profile(activities=[torch.profiler.ProfilerActivity.CPU,
torch.profiler.ProfilerActivity.CUDA]) as prof:
for route in ('nccl', 'safe', 'prepared'):
with torch.profiler.record_function(route):
a2a(x, 0, route)
a2a(y, 1, route)
torch.cuda.synchronize()
if rank == 0:
prof.export_chrome_trace(str(args.output.with_name(model + '-collectives-trace.json')))
del x, y, dst, captured_x, captured_y, graph
gc.collect()
torch.cuda.empty_cache()
else:
from fastvideo.attention.utils.flash_attn_default import fa_version, flash_attn_func_compilable
for model, sequence, heads in workloads:
inner = heads * 128
hidden = 5376 if args.h3_lengths else inner
ffn = 14336 if args.h3_lengths else 4 * hidden
x = torch.randn(1, sequence // world, hidden, device='cuda', dtype=torch.bfloat16,
requires_grad=True)
weights = [torch.nn.Parameter(torch.randn(out_dim, in_dim, device='cuda', dtype=torch.bfloat16)
* in_dim**-0.5)
for in_dim, out_dim in [(hidden, 3 * inner), (inner, hidden),
(hidden, 2 * ffn if args.h3_lengths else ffn), (ffn, hidden)]]
dy = torch.randn_like(x)
probe = torch.empty(3, sequence // world, heads, 128, device='cuda', dtype=torch.bfloat16)
# Initialize before parity; otherwise uninitialized NaNs can fail equality.
probe.zero_()
prepare([(probe, 0)])
del probe
if args.sparse:
from fastvideo_kernel.block_sparse_attn import block_sparse_attn_sm100a_op
if args.sparse_probe_dir:
import sys
sys.path.insert(0, str(args.sparse_probe_dir))
import ulysses_sparse_probe
import fastvideo_kernel.block_sparse_attn_sm100a as sparse_module
sparse_module._FWD_BY_BLOCK[64] = ulysses_sparse_probe.block_sparse_sm100a_fwd
padded_sequence = ((sequence + 127) // 128) * 128
block_count = padded_sequence // 64
topk = max(1, block_count // 10)
ids = ((torch.arange(block_count, device='cuda')[:, None]
+ torch.arange(topk, device='cuda')[None, :]) % block_count).to(torch.int32)
sparse_ids = ids[None, None].expand(1, heads // world, -1, -1).contiguous()
sparse_num = torch.full((1, heads // world, block_count), topk, device='cuda', dtype=torch.int32)
vbs = (sequence - torch.arange(block_count, device='cuda') * 64).clamp(0, 64).to(torch.int32)
def block(route):
projected = F.linear(x, weights[0]).reshape(1, sequence // world, 3, heads, 128)
qkv = torch.cat(projected.unbind(dim=2), dim=0)
q, k, v = a2a(qkv, 0, route).chunk(3, dim=0)
if args.sparse:
padded = [F.pad(t.permute(0, 2, 1, 3), (0, 0, 0, padded_sequence - sequence)).contiguous()
for t in (q, k, v)]
attended = block_sparse_attn_sm100a_op(*padded, sparse_ids, sparse_num, vbs)[0]
attended = attended[:, :, :sequence].permute(0, 2, 1, 3)
else:
attended = flash_attn_func_compilable(q, k, v, causal=False)
local = a2a(attended.contiguous(), 1, route).flatten(2)
residual = x + F.linear(local, weights[1])
activated = F.linear(residual, weights[2])
if args.h3_lengths:
value, gate = activated.chunk(2, dim=-1)
activated = value * F.silu(gate)
else:
activated = F.gelu(activated, approximate='tanh')
return residual + F.linear(activated, weights[3])
# Compare outputs and all input/weight gradients under the identical compute recipe.
reference = block('nccl')
reference_grads = torch.autograd.grad(reference, [x, *weights], dy)
for route in [r for r in routes if r != 'nccl']:
result = block(route)
gradients = torch.autograd.grad(result, [x, *weights], dy)
torch.testing.assert_close(result, reference, rtol=0, atol=0)
# FA backward uses atomic reductions: allow its bf16 summation variation.
for actual, expected in zip(gradients, reference_grads):
torch.testing.assert_close(actual, expected, rtol=0.03, atol=0.03)
del result, gradients
del reference, reference_grads
for training in (False, True):
for round_index in (range(args.rounds) if args.paired else [None]):
ordered = routes if round_index is None else routes[round_index % len(routes):] + routes[:round_index % len(routes)]
for route in ordered:
def step():
with torch.set_grad_enabled(training):
out = block(route)
if training:
torch.autograd.grad(out, [x, *weights], dy)
measure(step, dict(model=model, operation='block_train' if training else 'block_infer',
route=route, flash_attention='sm100a64+triton_bwd' if args.sparse else fa_version,
sparse=args.sparse, sequence=sequence, heads=heads,
note='synthetic block; no FSDP, optimizer, norms, checkpointing, or sparse routing'),
iterations=max(3, args.iters // 4), round_index=round_index)
del x, weights, dy
gc.collect()
torch.cuda.empty_cache()
finally:
for candidate in probe_helpers.values():
candidate.close()
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: