Compare commits
8
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
fc51550920 | ||
|
|
942e8b8404 | ||
|
|
c548e5834e | ||
|
|
368f6b8891 | ||
|
|
e55934b62e | ||
|
|
a06e63827d | ||
|
|
e536bb8544 | ||
|
|
3eabb7b40b |
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user