|
| 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) |
0 commit comments