Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
90 changes: 71 additions & 19 deletions experimental/CollectiveX/bench/ep_nccl.py
Original file line number Diff line number Diff line change
Expand Up @@ -114,6 +114,9 @@ def __init__(self, args, rank, world_size, local_rank, device):
self._combine_cfg = CombineConfig(send_only=0)
self._comm = None
self._ep_group = None
# Exactly ONE handle per group, rebound per problem shape — see _ensure_handle.
self._handle = None
self._bound = None

def buffer_cap(self, args):
if self._ll:
Expand Down Expand Up @@ -238,15 +241,34 @@ def create_buffer(self, spec):
self._ht_disp_counts_t = self._t(self._ht_disp_counts)

def _ensure_handle(self, p):
"""Create (once, cached on the problem) the reusable per-step handle for p's routing.

create_handle is collective and, in HT, performs the metadata exchange that fixes the
received-token count — so it must run in the same order on every rank. It first runs
inside Pass 1's untimed warm (ladder order, identical across ranks); the timed passes
then reuse the cached handle and never enter a collective here.
"""Bind the group's single handle to p's routing, creating it on first use.

ONE handle per group, rebound per shape — never one handle per shape. `buffer_idx`, the
LL double-buffer parity selector, is per-HANDLE state, but the buffers it selects are
offsets into the per-GROUP rdma_buffer: two handles built from the same group config
resolve to the SAME parity-0/parity-1 count+flag slots and advance their parity
independently, so one handle's "next buffer, safe to clean" is the other's "current
buffer, in flight". Interleaving handles therefore corrupts the signalling even when it
does not hang outright, which makes every latency drawn from such a run suspect. Filed
upstream as NVIDIA/nccl#2303; reproduced on a stock wheel by ladder [1] (one handle,
clean) vs ladder [1, 2] (two handles, 64 dispatch + 6 combine receive timeouts -> 719).

`ncclEpInitHandle` takes no token count and `ncclEpUpdateHandle` is documented as a
"per-step collective: prepare the handle for the given top-k routing decisions", so
rebinding IS the intended lifecycle. This also matches the other three backends, which
each allocate one object sized to the ladder maximum and vary the token count per call.

Both create_handle and update are collective, and HT additionally performs the metadata
exchange that fixes the received-token count, so they must run in the same order on
every rank. They only ever run on a shape CHANGE, and every timed component is preceded
by an untimed warm() on its own problem — so the collective always lands in warm (or in
the oracle passes), never inside a timed window. Re-entering with the already-bound
problem returns immediately without a collective or a sync.
"""
cached = getattr(p, "_nccl", None)
if cached is not None:
if self._bound is not cached:
self._rebind(cached)
return cached
stream = self._stream()
topk_idx_t = self._t(p.topk_idx)
Expand All @@ -271,18 +293,45 @@ def _ensure_handle(self, p):
expert_counters=self._t(h.recv_experts),
recv_total_counter=self._t(h.recv_total),
)
h.handle = self._ep_group.create_handle(
self._layout,
topk_idx_t,
layout_info=ht_layout_info,
config=HandleConfig(),
stream=stream,
# LL takes layout_info only on dispatch (the API forbids it on create/update); HT needs
# it here so this problem's counters receive its own metadata-exchange results.
h.layout_info = ht_layout_info
if self._handle is None:
self._handle = self._ep_group.create_handle(
self._layout,
topk_idx_t,
layout_info=ht_layout_info,
config=HandleConfig(),
stream=stream,
)
h.handle = self._handle
torch.cuda.synchronize()
if not self._ll:
h.count = int(h.recv_total.item())
self._bound = h
else:
h.handle = self._handle
self._rebind(h)
p._nccl = h
return h

def _rebind(self, h):
"""Point the single handle at h's routing (collective; untimed callers only).

Rebinding does not reallocate: `Handle.update` only swaps in the new top-k indices.
HT re-reads its received-token count because the metadata exchange recomputes it for
this routing; the value is deterministic per problem, so a later rebind to the same
problem reproduces it.
"""
self._handle.update(
h.topk_idx_t,
layout_info=None if self._ll else h.layout_info,
stream=self._stream(),
)
torch.cuda.synchronize()
if not self._ll:
h.count = int(h.recv_total.item())
p._nccl = h
return h
self._bound = h

# ---- transport contract ------------------------------------------------------------------

Expand Down Expand Up @@ -464,8 +513,11 @@ def finalize(self, rc):
return rc

def _destroy_handles(self):
# Per-problem handles are cached on the problem namespaces; the group keeps no registry,
# so there is nothing to walk here. Handles are released when their problems are GC'd
# (Handle.destroy runs in the binding's __del__); the group/comm destroy below reclaims
# the device buffers. Kept as a seam in case bring-up needs explicit handle teardown.
return
# One handle for the whole group, so teardown is a single explicit destroy rather than
# waiting on per-problem GC. The problem namespaces still hold a reference to it for
# their dispatch/combine calls; dropping _bound first keeps a late rebind from touching
# a destroyed handle. The group/comm destroy below reclaims the device buffers.
self._bound = None
if self._handle is not None:
self._handle.destroy()
self._handle = None
171 changes: 171 additions & 0 deletions experimental/CollectiveX/tests/test_ep_nccl_handle.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,171 @@
#!/usr/bin/env python3
"""Torch-free tests for the nccl-ep single-handle contract.

The invariant under test: ONE handle per group, rebound per problem shape. Two handles built
from the same group config resolve to the same LL parity signal slots while advancing their
parity independently, which corrupts signalling (NVIDIA/nccl#2303) -- so a regression that
reintroduces per-shape handles must fail loudly here rather than on a cluster.

torch and nccl are stubbed so this runs without the benchmark image.
"""
from __future__ import annotations

import sys
import types
import unittest
from pathlib import Path
from unittest import mock

ROOT = Path(__file__).resolve().parents[1]


def _stub_modules():
"""Fake torch / nccl modules so `import ep_nccl` succeeds without the benchmark image."""
torch = types.ModuleType("torch")
torch.bfloat16 = "bfloat16"
torch.int32 = "int32"
torch.empty = lambda *a, **k: types.SimpleNamespace(shape=a[0] if a else ())
torch.zeros = lambda *a, **k: types.SimpleNamespace(item=lambda: 7)
torch.cuda = types.SimpleNamespace(synchronize=lambda: None)
dist = types.ModuleType("torch.distributed")
torch.distributed = dist

ep = types.ModuleType("nccl.ep")
for name in (
"Algorithm", "CombineConfig", "CombineInputs", "CombineOutputs", "DispatchConfig",
"DispatchInputs", "DispatchOutputs", "GroupConfig", "HandleConfig", "Layout",
"LayoutInfo", "Tensor",
):
setattr(ep, name, type(name, (), {"__init__": lambda self, *a, **k: None}))
ep.Algorithm = types.SimpleNamespace(LOW_LATENCY="LL", HIGH_THROUGHPUT="HT")
ep.Layout = types.SimpleNamespace(EXPERT_MAJOR="EM", FLAT="FLAT")
core = types.ModuleType("nccl.core")
pkg = types.ModuleType("nccl")
pkg.ep, pkg.core = ep, core
return {
"torch": torch, "torch.distributed": dist,
"nccl": pkg, "nccl.ep": ep, "nccl.core": core,
}


sys.path[:0] = [str(ROOT), str(ROOT / "bench")]

# Import ep_nccl against the stubs, then withdraw them: leaving a fake torch in sys.modules
# makes the genuinely torch-dependent modules in this process (test_runtime, test_ll_oracle)
# error instead of skipping. Dropping ep_nccl too keeps the stub-built module private to us.
with mock.patch.dict(sys.modules, _stub_modules()):
import ep_nccl # noqa: E402

sys.modules.pop("ep_nccl", None)


class FakeHandle:
"""Records every rebind so the tests can assert on the collective call pattern."""

def __init__(self):
self.updates = []
self.destroyed = False

def update(self, topk_idx, *, layout_info=None, stream=None):
self.updates.append((topk_idx, layout_info))

def destroy(self):
self.destroyed = True


class FakeGroup:
def __init__(self):
self.created = 0
self.handle = FakeHandle()

def create_handle(self, layout, topk_idx, *, layout_info=None, config=None, stream=None):
self.created += 1
return self.handle


def backend(ll=True):
"""An NCCLEPBackend with just the fields _ensure_handle touches (no __init__, no GPU)."""
b = object.__new__(ep_nccl.NCCLEPBackend)
b._ll = ll
b._layout = "EM" if ll else "FLAT"
b._handle = None
b._bound = None
b._ep_group = FakeGroup()
b.device = "cuda:0"
b.num_local_experts = 4
b.args = types.SimpleNamespace(hidden=16)
b._t = lambda x: x
b._stream = lambda: 0
return b


def problem(T):
return types.SimpleNamespace(
T=T, dispatch_x=f"x{T}", topk_idx=f"idx{T}", topk_weights=f"w{T}"
)


class TestSingleHandle(unittest.TestCase):
def test_one_handle_across_many_shapes(self):
"""Nine ladder rungs must still produce exactly one create_handle."""
b = backend()
for T in (1, 2, 4, 8, 16, 32, 64, 128, 256):
b._ensure_handle(problem(T))
self.assertEqual(b._ep_group.created, 1)

def test_shape_change_rebinds_and_repeat_does_not(self):
"""update() on a shape switch; no collective when the bound shape is re-entered."""
b = backend()
pa, pb = problem(1), problem(2)
b._ensure_handle(pa)
self.assertEqual(len(b._ep_group.handle.updates), 0) # first bind is the create

b._ensure_handle(pb)
self.assertEqual(len(b._ep_group.handle.updates), 1)

# Re-entering the bound problem repeatedly -- the timed loop's steady state -- must not
# enter a collective, otherwise every iteration gains a rank-synchronising step.
for _ in range(8):
b._ensure_handle(pb)
self.assertEqual(len(b._ep_group.handle.updates), 1)

# Returning to an earlier shape rebinds again (its cached namespace is reused).
b._ensure_handle(pa)
self.assertEqual(len(b._ep_group.handle.updates), 2)

def test_every_problem_shares_the_one_handle(self):
b = backend()
handles = {id(b._ensure_handle(problem(T)).handle) for T in (1, 2, 4)}
self.assertEqual(len(handles), 1)
self.assertIs(b._ensure_handle(problem(1)).handle, b._handle)

def test_ll_never_passes_layout_info_on_rebind(self):
"""The API forbids layout_info on create/update in LL mode."""
b = backend(ll=True)
b._ensure_handle(problem(1))
b._ensure_handle(problem(2))
self.assertEqual([info for _, info in b._ep_group.handle.updates], [None])

def test_ht_rebind_carries_that_problems_counters(self):
"""HT re-runs the metadata exchange into the rebound problem's own counter tensors."""
b = backend(ll=False)
ha = b._ensure_handle(problem(1))
hb = b._ensure_handle(problem(2))
self.assertEqual(len(b._ep_group.handle.updates), 1)
self.assertIs(b._ep_group.handle.updates[0][1], hb.layout_info)
self.assertIsNot(ha.layout_info, hb.layout_info)
self.assertEqual(hb.count, 7) # re-read after the exchange

def test_destroy_releases_the_handle_once(self):
b = backend()
b._ensure_handle(problem(1))
handle = b._handle
b._destroy_handles()
self.assertTrue(handle.destroyed)
self.assertIsNone(b._handle)
self.assertIsNone(b._bound)
b._destroy_handles() # idempotent


if __name__ == "__main__":
unittest.main()