Skip to content

Commit 3452634

Browse files
committed
ep: PyTorch wrapper, autograd ops, symm-mem zero-copy bindings + distributed tests/example
Signed-off-by: Phuong Nguyen <phuonguyen@nvidia.com>
1 parent 44a8a49 commit 3452634

10 files changed

Lines changed: 2241 additions & 0 deletions

File tree

build_tools/pytorch.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -76,6 +76,12 @@ def setup_pytorch_extension(
7676

7777
setup_mpi_flags(include_dirs, cxx_flags)
7878

79+
# Mirror the NCCL EP gate from setup.py / common CMake. When disabled, the
80+
# ep.cpp source no-ops at the #ifdef boundary; without the define it would
81+
# produce undefined references to nvte_ep_*.
82+
if bool(int(os.getenv("NVTE_BUILD_WITH_NCCL_EP", "1"))):
83+
cxx_flags.append("-DNVTE_WITH_NCCL_EP")
84+
7985
library_dirs = []
8086
libraries = []
8187
if bool(int(os.getenv("NVTE_ENABLE_NVSHMEM", 0))):

examples/pytorch/ep/ep_moe.py

Lines changed: 281 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,281 @@
1+
# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2+
#
3+
# See LICENSE for license information.
4+
"""End-to-end MoE example: dispatch -> batched expert linear -> combine, fwd + bwd.
5+
6+
One process per GPU. Launched via run_test_ep.sh (torchrun).
7+
"""
8+
9+
import argparse
10+
import os
11+
import sys
12+
13+
import numpy as np
14+
import torch
15+
import torch.distributed as dist
16+
17+
from transformer_engine.pytorch.ep import (
18+
EpHandle,
19+
ep_bootstrap,
20+
ep_dispatch,
21+
ep_combine,
22+
symm_mem_alloc,
23+
)
24+
25+
26+
def _parse_args():
27+
p = argparse.ArgumentParser(description="TE-PyTorch EP MoE example (fwd + bwd)")
28+
p.add_argument("--num-tokens", type=int, default=8, help="Per-rank token count.")
29+
p.add_argument("--top-k", type=int, default=2)
30+
p.add_argument("--hidden", type=int, default=32)
31+
p.add_argument("--hidden-out", type=int, default=32)
32+
p.add_argument("--num-experts", type=int, default=None)
33+
p.add_argument("--check", action="store_true", default=True)
34+
p.add_argument(
35+
"--benchmark",
36+
action="store_true",
37+
help="Time fwd+bwd over both regular HBM and NCCL symm-mem payload buffers.",
38+
)
39+
p.add_argument("--benchmark-iters", type=int, default=20)
40+
p.add_argument("--benchmark-warmup", type=int, default=5)
41+
return p.parse_args()
42+
43+
44+
def _make_routing(rank, T, K, E, num_local_experts):
45+
"""Deterministic: topk_idx[t, k] = (rank*NLE + t*K + k) % E."""
46+
topk_idx = np.empty((T, K), dtype=np.int32)
47+
for t in range(T):
48+
for k in range(K):
49+
topk_idx[t, k] = (rank * num_local_experts + t * K + k) % E
50+
return topk_idx
51+
52+
53+
def _batched_expert_linear(recv_tokens, kernels, num_local_experts):
54+
"""Per-expert linear over [recv_pr // NLE] tokens per expert."""
55+
recv_pr, _H = recv_tokens.shape
56+
H_out = kernels.shape[-1]
57+
slots_per_expert = recv_pr // num_local_experts
58+
grouped = recv_tokens.view(num_local_experts, slots_per_expert, recv_tokens.shape[-1])
59+
# kernels: [num_local_experts, H, H_out]
60+
out = torch.bmm(grouped, kernels.to(grouped.dtype))
61+
return out.view(recv_pr, H_out)
62+
63+
64+
def _reference_moe(tokens, topk_idx, topk_w, kernels):
65+
T, K = topk_idx.shape
66+
H_out = kernels.shape[-1]
67+
out = np.zeros((T, H_out), dtype=np.float32)
68+
for t in range(T):
69+
tok = tokens[t].astype(np.float32)
70+
for k in range(K):
71+
e = int(topk_idx[t, k])
72+
out[t] += float(topk_w[t, k]) * (tok @ kernels[e].astype(np.float32))
73+
return out
74+
75+
76+
def _reference_grad(tokens, topk_idx, topk_w, kernels):
77+
T, K = topk_idx.shape
78+
H = tokens.shape[-1]
79+
ref_out = _reference_moe(tokens, topk_idx, topk_w, kernels)
80+
grad = np.zeros((T, H), dtype=np.float32)
81+
for t in range(T):
82+
mixed = np.zeros_like(kernels[0])
83+
for k in range(K):
84+
mixed = mixed + float(topk_w[t, k]) * kernels[int(topk_idx[t, k])]
85+
grad[t] = ref_out[t] @ mixed.T
86+
return ref_out, grad
87+
88+
89+
def main():
90+
args = _parse_args()
91+
92+
dist.init_process_group(backend="nccl")
93+
rank = dist.get_rank()
94+
world_size = dist.get_world_size()
95+
torch.cuda.set_device(int(os.environ.get("LOCAL_RANK", rank)))
96+
device = torch.device("cuda", torch.cuda.current_device())
97+
98+
major, minor = torch.cuda.get_device_capability()
99+
if major * 10 + minor < 90:
100+
if rank == 0:
101+
print(f"[ep_moe] SKIPPED: EP requires SM>=90 (got SM{major}{minor})")
102+
dist.destroy_process_group()
103+
return
104+
105+
if world_size < 4:
106+
if rank == 0:
107+
print(f"[ep_moe] SKIPPED: EP requires >= 4 ranks (got {world_size})")
108+
dist.destroy_process_group()
109+
return
110+
111+
ep_size = world_size
112+
num_experts = args.num_experts if args.num_experts is not None else world_size
113+
assert num_experts % ep_size == 0
114+
num_local_experts = num_experts // ep_size
115+
T = args.num_tokens
116+
recv_pr = ep_size * T * args.top_k
117+
118+
ep_group = dist.new_group(ranks=list(range(world_size)), backend="nccl")
119+
ep_bootstrap(
120+
ep_group,
121+
num_experts=num_experts,
122+
max_tokens_per_rank=T,
123+
recv_capacity_per_rank=recv_pr,
124+
hidden_dim=args.hidden,
125+
)
126+
127+
rng = np.random.default_rng(seed=42 + rank)
128+
tokens_np = (rng.standard_normal((T, args.hidden), dtype=np.float32) * 0.5).astype(np.float32)
129+
topk_idx_np = _make_routing(rank, T, args.top_k, num_experts, num_local_experts)
130+
w_np = np.full((T, args.top_k), 1.0 / args.top_k, dtype=np.float32)
131+
# Same seed on every rank → identical kernel array everywhere.
132+
kr = np.random.default_rng(seed=42)
133+
kernels_np = (
134+
kr.standard_normal((num_experts, args.hidden, args.hidden_out), dtype=np.float32)
135+
* (1.0 / np.sqrt(args.hidden))
136+
).astype(np.float32)
137+
138+
tokens = (
139+
torch.from_numpy(tokens_np).to(device=device, dtype=torch.bfloat16).requires_grad_(True)
140+
)
141+
topk_idx = torch.from_numpy(topk_idx_np).to(device)
142+
topk_w = torch.from_numpy(w_np).to(device)
143+
kernels_local = torch.from_numpy(
144+
kernels_np[rank * num_local_experts : (rank + 1) * num_local_experts]
145+
).to(device=device, dtype=torch.bfloat16)
146+
147+
handle = EpHandle(
148+
top_k=args.top_k,
149+
max_tokens_per_rank=T,
150+
recv_capacity_per_rank=recv_pr,
151+
hidden_dim=args.hidden,
152+
num_local_experts=num_local_experts,
153+
)
154+
recv_tokens = torch.empty(recv_pr, args.hidden, dtype=torch.bfloat16, device=device)
155+
recv_w = torch.empty(recv_pr, dtype=torch.float32, device=device)
156+
token_counts = torch.empty(num_local_experts, dtype=torch.int32, device=device)
157+
result = torch.empty(T, args.hidden_out, dtype=torch.bfloat16, device=device)
158+
grad_tokens = torch.empty_like(tokens)
159+
grad_topk_w = torch.empty_like(topk_w)
160+
grad_eo = torch.empty(recv_pr, args.hidden_out, dtype=torch.bfloat16, device=device)
161+
162+
recv_t, recv_w_out, _tc = ep_dispatch(
163+
handle,
164+
topk_idx,
165+
tokens,
166+
topk_w,
167+
recv_tokens,
168+
recv_w,
169+
token_counts,
170+
grad_tokens,
171+
grad_topk_w,
172+
)
173+
expert_out = _batched_expert_linear(recv_t, kernels_local, num_local_experts)
174+
out = ep_combine(handle, expert_out, recv_w_out, result, grad_eo)
175+
176+
loss = 0.5 * (out.float() ** 2).sum()
177+
loss.backward()
178+
torch.cuda.synchronize()
179+
180+
if rank == 0:
181+
print(
182+
f"[ep_moe] loss={float(loss):.4f} grad_tokens.shape={tuple(tokens.grad.shape)} "
183+
f"ep={ep_size} num_experts={num_experts} recv_pr={recv_pr}"
184+
)
185+
186+
if args.benchmark:
187+
# Time forward dispatch + expert + combine over both regular HBM and
188+
# NCCL symm-mem payloads. Recorded only — perf ratio depends heavily on
189+
# topology and is not CI-gated.
190+
import time
191+
192+
def _time(label, tokens_buf, recv_tokens_buf, result_buf):
193+
torch.cuda.synchronize()
194+
dist.barrier()
195+
for _ in range(args.benchmark_warmup):
196+
rt, rw, _tc = ep_dispatch(
197+
handle,
198+
topk_idx,
199+
tokens_buf,
200+
topk_w,
201+
recv_tokens_buf,
202+
recv_w,
203+
token_counts,
204+
grad_tokens,
205+
grad_topk_w,
206+
)
207+
eo = _batched_expert_linear(rt, kernels_local, num_local_experts)
208+
ep_combine(handle, eo, rw, result_buf, grad_eo)
209+
torch.cuda.synchronize()
210+
dist.barrier()
211+
t0 = time.perf_counter()
212+
for _ in range(args.benchmark_iters):
213+
rt, rw, _tc = ep_dispatch(
214+
handle,
215+
topk_idx,
216+
tokens_buf,
217+
topk_w,
218+
recv_tokens_buf,
219+
recv_w,
220+
token_counts,
221+
grad_tokens,
222+
grad_topk_w,
223+
)
224+
eo = _batched_expert_linear(rt, kernels_local, num_local_experts)
225+
ep_combine(handle, eo, rw, result_buf, grad_eo)
226+
torch.cuda.synchronize()
227+
dt_ms = (time.perf_counter() - t0) * 1000.0 / args.benchmark_iters
228+
if rank == 0:
229+
print(
230+
f"[ep_moe --benchmark] {label}: {dt_ms:.3f} ms/iter "
231+
f"(iters={args.benchmark_iters})"
232+
)
233+
return dt_ms
234+
235+
# 1) HBM baseline reuses the existing buffers (already allocated above).
236+
hbm_ms = _time("regular HBM", tokens.detach(), recv_tokens, result)
237+
238+
# 2) Symm-mem variant: reallocate as symm-mem (collective on every rank).
239+
try:
240+
tokens_sm = symm_mem_alloc((T, args.hidden), torch.bfloat16, ep_group, device=device)
241+
recv_tokens_sm = symm_mem_alloc(
242+
(recv_pr, args.hidden), torch.bfloat16, ep_group, device=device
243+
)
244+
result_sm = torch.empty_like(result)
245+
tokens_sm.copy_(tokens.detach())
246+
symm_ms = _time("symm-mem", tokens_sm, recv_tokens_sm, result_sm)
247+
if rank == 0:
248+
print(f"[ep_moe --benchmark] speedup: {hbm_ms / symm_ms:.2f}x")
249+
except RuntimeError as e:
250+
if rank == 0:
251+
print(f"[ep_moe --benchmark] symm-mem path skipped: {e}")
252+
253+
if args.check:
254+
# Gather across ranks for a global reference comparison.
255+
global_tokens = [torch.empty_like(tokens) for _ in range(world_size)]
256+
global_topk_idx = [torch.empty_like(topk_idx) for _ in range(world_size)]
257+
global_topk_w = [torch.empty_like(topk_w) for _ in range(world_size)]
258+
global_out = [torch.empty_like(out) for _ in range(world_size)]
259+
global_grad = [torch.empty_like(tokens.grad) for _ in range(world_size)]
260+
dist.all_gather(global_tokens, tokens.detach())
261+
dist.all_gather(global_topk_idx, topk_idx)
262+
dist.all_gather(global_topk_w, topk_w)
263+
dist.all_gather(global_out, out.detach())
264+
dist.all_gather(global_grad, tokens.grad)
265+
if rank == 0:
266+
all_tokens = torch.cat(global_tokens).float().cpu().numpy()
267+
all_idx = torch.cat(global_topk_idx).cpu().numpy()
268+
all_w = torch.cat(global_topk_w).cpu().numpy()
269+
all_out = torch.cat(global_out).float().cpu().numpy()
270+
all_grad = torch.cat(global_grad).float().cpu().numpy()
271+
ref_out, ref_grad = _reference_grad(all_tokens, all_idx, all_w, kernels_np)
272+
np.testing.assert_allclose(all_out, ref_out, rtol=5e-2, atol=5e-2)
273+
np.testing.assert_allclose(all_grad, ref_grad, rtol=5e-2, atol=5e-2)
274+
print(f"[ep_moe] --check PASSED (ref_out.sum()={float(ref_out.sum()):.4f})")
275+
276+
dist.destroy_process_group()
277+
278+
279+
if __name__ == "__main__":
280+
main()
281+
sys.exit(0)

examples/pytorch/ep/run_test_ep.sh

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,37 @@
1+
#!/bin/bash
2+
# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
3+
#
4+
# See LICENSE for license information.
5+
6+
set -uo pipefail
7+
8+
DETECTED_GPUS=$(nvidia-smi -L 2>/dev/null | wc -l)
9+
NUM_GPUS="${NUM_GPUS:-${DETECTED_GPUS}}"
10+
if [ "${NUM_GPUS}" -lt 4 ]; then
11+
echo "EP requires >= 4 GPUs (found ${NUM_GPUS}); SKIPPING."
12+
exit 0
13+
fi
14+
if [ "${NUM_GPUS}" -gt 8 ]; then NUM_GPUS=8; fi
15+
16+
: ${TE_PATH:=/opt/transformerengine}
17+
: ${TEST_TIMEOUT_S:=120}
18+
19+
SCRIPT="${TE_PATH}/examples/pytorch/ep/ep_moe.py"
20+
export PYTHONPATH="${TE_PATH}${PYTHONPATH:+:${PYTHONPATH}}"
21+
22+
# Stage JIT cubins on tmpfs for fast iteration.
23+
: ${NCCL_EP_JIT_CACHE_DIR:="${TMPDIR:-/tmp}/nccl_ep_jit_cache_$(id -u)"}
24+
export NCCL_EP_JIT_CACHE_DIR
25+
mkdir -p "$NCCL_EP_JIT_CACHE_DIR"
26+
27+
echo "*** Executing ep_moe.py across ${NUM_GPUS} GPUs (timeout=${TEST_TIMEOUT_S}s) ***"
28+
timeout --foreground --signal=KILL "${TEST_TIMEOUT_S}" \
29+
torchrun --standalone --nnodes=1 --nproc-per-node="${NUM_GPUS}" \
30+
"${SCRIPT}" --check 2>&1 | tee stdout_ep_moe.txt
31+
RC=${PIPESTATUS[0]}
32+
33+
RET=0
34+
if [ "${RC}" -ne 0 ]; then RET=1; fi
35+
if grep -qE "FAILED|Traceback" stdout_ep_moe.txt; then RET=1; fi
36+
rm -f stdout_ep_moe.txt
37+
exit $RET

0 commit comments

Comments
 (0)