Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
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
842 changes: 842 additions & 0 deletions tests/pytorch/distributed/moe_ep_reference.py

Large diffs are not rendered by default.

258 changes: 258 additions & 0 deletions tests/pytorch/distributed/run_ep.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,8 @@
import torch
import torch.distributed as dist

from moe_ep_reference import MoeEpReference
from transformer_engine.pytorch import ops as te_ops
from transformer_engine.pytorch.ep import (
EpBuffer,
ep_bootstrap,
Expand Down Expand Up @@ -121,6 +123,30 @@ def _make_identity_inputs(rank, ep_size, device="cuda"):
)


def _make_moe_inputs(rank, ep_size, device="cuda"):
"""Random activations and top-k routing representative of a router output."""
generator = torch.Generator(device=device)
generator.manual_seed(2026 + rank)
num_experts = ep_size * NUM_LOCAL_EXPERTS
tokens = torch.randn(
TOKENS_PER_RANK,
HIDDEN_DIM,
generator=generator,
dtype=torch.float32,
device=device,
).mul_(0.25)
router_logits = torch.randn(
TOKENS_PER_RANK,
num_experts,
generator=generator,
dtype=torch.float32,
device=device,
)
topk_logits, topk_idx = torch.topk(router_logits, TOP_K, dim=-1)
topk_weights = torch.softmax(topk_logits, dim=-1)
return topk_idx, tokens.to(torch.bfloat16), topk_weights


class _Cfg:
rank: int
world_size: int
Expand Down Expand Up @@ -591,6 +617,238 @@ def zero_grads():
rtol=5e-2,
)

@_eager_test_include
def test_fusible_dispatch_combine_moe(self):
"""Fusible NCCL EP MoE matches the PyTorch all-to-all reference."""
if not EAGER:
self.skipTest("variable grouped-linear splits require eager EP output sizing")

topk_idx, tokens, topk_weights = _make_moe_inputs(self.cfg.rank, self.cfg.ep_size)
num_experts = NUM_LOCAL_EXPERTS
intermediate_dim = HIDDEN_DIM

dispatch_buffer = self._make_buffer()
dispatch = te_ops.Dispatch(dispatch_buffer)
fc1 = te_ops.GroupedLinear(
num_experts,
HIDDEN_DIM,
2 * intermediate_dim,
bias=False,
device=self.cfg.device,
dtype=torch.bfloat16,
)
activation = te_ops.ScaledSwiGLU()
fc2 = te_ops.GroupedLinear(
num_experts,
intermediate_dim,
HIDDEN_DIM,
bias=False,
device=self.cfg.device,
dtype=torch.bfloat16,
)
combine = te_ops.Combine(dispatch_buffer, num_local_tokens=TOKENS_PER_RANK)

dispatch.set_extra_output_channel(0, "m_splits")
dispatch.set_extra_output_channel(1, "token_probs")
fc1.set_extra_input_channel(0, "m_splits")
activation.set_extra_input_channel(0, "token_probs")
fc2.set_extra_input_channel(0, "m_splits")
model = te_ops.Sequential(dispatch, fc1, activation, fc2, combine)

generator = torch.Generator(device=self.cfg.device)
generator.manual_seed(1234 + self.cfg.rank)
with torch.no_grad():
for expert_idx in range(num_experts):
getattr(fc1, f"weight{expert_idx}").uniform_(-0.1, 0.1, generator=generator)
getattr(fc2, f"weight{expert_idx}").uniform_(-0.1, 0.1, generator=generator)

ref_fc1_weights = torch.stack(
[getattr(fc1, f"weight{idx}").detach().transpose(0, 1) for idx in range(num_experts)]
)
ref_fc2_weights = torch.stack(
[getattr(fc2, f"weight{idx}").detach().transpose(0, 1) for idx in range(num_experts)]
)
reference = MoeEpReference(
num_experts=self.cfg.num_experts,
hidden_size=HIDDEN_DIM,
intermediate_size=intermediate_dim,
top_k=TOP_K,
ep_group=self.ep_group,
max_tokens_per_rank=TOKENS_PER_RANK,
apply_topk_in_fc1=True,
generate_c=True,
)
ref_output, fc1_c, route_metadata = reference(
tokens,
ref_fc1_weights,
ref_fc2_weights,
topk_idx,
topk_weights,
)

test_tokens = tokens.detach().clone().requires_grad_(True)
test_topk_weights = topk_weights.detach().clone().requires_grad_(True)
test_output = model(test_tokens, topk_idx, test_topk_weights)

grad_output = torch.linspace(
-0.2,
0.2,
TOKENS_PER_RANK * HIDDEN_DIM,
device=self.cfg.device,
dtype=torch.float32,
).reshape(TOKENS_PER_RANK, HIDDEN_DIM)
grad_output = grad_output.to(torch.bfloat16)
ref_grads = reference.backward(
grad_output,
tokens,
ref_fc1_weights,
ref_fc2_weights,
topk_idx,
topk_weights,
fc1_c,
route_metadata,
)
test_output.backward(grad_output)
torch.cuda.synchronize()

torch.testing.assert_close(test_output.float(), ref_output.float(), atol=5e-2, rtol=5e-2)
torch.testing.assert_close(test_tokens.grad.float(), ref_grads[0], atol=5e-2, rtol=5e-2)
torch.testing.assert_close(
test_topk_weights.grad,
ref_grads[3],
atol=5e-2,
rtol=5e-2,
)
for expert_idx in range(num_experts):
torch.testing.assert_close(
getattr(fc1, f"weight{expert_idx}").grad.float(),
ref_grads[1][expert_idx].transpose(0, 1),
atol=5e-2,
rtol=5e-2,
)
torch.testing.assert_close(
getattr(fc2, f"weight{expert_idx}").grad.float(),
ref_grads[2][expert_idx].transpose(0, 1),
atol=5e-2,
rtol=5e-2,
)

@_eager_test_include
def test_fusible_dispatch_combine_moe(self):
"""Fusible NCCL EP MoE matches the PyTorch all-to-all reference."""
if not EAGER:
self.skipTest("variable grouped-linear splits require eager EP output sizing")

topk_idx, tokens, topk_weights = _make_moe_inputs(self.cfg.rank, self.cfg.ep_size)
num_experts = NUM_LOCAL_EXPERTS
intermediate_dim = HIDDEN_DIM

dispatch_buffer = self._make_buffer()
dispatch = te_ops.Dispatch(dispatch_buffer)
fc1 = te_ops.GroupedLinear(
num_experts,
HIDDEN_DIM,
2 * intermediate_dim,
bias=False,
device=self.cfg.device,
dtype=torch.bfloat16,
)
activation = te_ops.ScaledSwiGLU()
fc2 = te_ops.GroupedLinear(
num_experts,
intermediate_dim,
HIDDEN_DIM,
bias=False,
device=self.cfg.device,
dtype=torch.bfloat16,
)
combine = te_ops.Combine(dispatch_buffer, num_local_tokens=TOKENS_PER_RANK)

dispatch.set_extra_output_channel(0, "m_splits")
dispatch.set_extra_output_channel(1, "token_probs")
fc1.set_extra_input_channel(0, "m_splits")
activation.set_extra_input_channel(0, "token_probs")
fc2.set_extra_input_channel(0, "m_splits")
model = te_ops.Sequential(dispatch, fc1, activation, fc2, combine)

generator = torch.Generator(device=self.cfg.device)
generator.manual_seed(1234 + self.cfg.rank)
with torch.no_grad():
for expert_idx in range(num_experts):
getattr(fc1, f"weight{expert_idx}").uniform_(-0.1, 0.1, generator=generator)
getattr(fc2, f"weight{expert_idx}").uniform_(-0.1, 0.1, generator=generator)

ref_fc1_weights = torch.stack(
[getattr(fc1, f"weight{idx}").detach().transpose(0, 1) for idx in range(num_experts)]
)
ref_fc2_weights = torch.stack(
[getattr(fc2, f"weight{idx}").detach().transpose(0, 1) for idx in range(num_experts)]
)
reference = MoeEpReference(
num_experts=self.cfg.num_experts,
hidden_size=HIDDEN_DIM,
intermediate_size=intermediate_dim,
top_k=TOP_K,
ep_group=self.ep_group,
max_tokens_per_rank=TOKENS_PER_RANK,
apply_topk_in_fc1=True,
generate_c=True,
)
ref_output, fc1_c, route_metadata = reference(
tokens,
ref_fc1_weights,
ref_fc2_weights,
topk_idx,
topk_weights,
)

test_tokens = tokens.detach().clone().requires_grad_(True)
test_topk_weights = topk_weights.detach().clone().requires_grad_(True)
test_output = model(test_tokens, topk_idx, test_topk_weights)

grad_output = torch.linspace(
-0.2,
0.2,
TOKENS_PER_RANK * HIDDEN_DIM,
device=self.cfg.device,
dtype=torch.float32,
).reshape(TOKENS_PER_RANK, HIDDEN_DIM)
grad_output = grad_output.to(torch.bfloat16)
ref_grads = reference.backward(
grad_output,
tokens,
ref_fc1_weights,
ref_fc2_weights,
topk_idx,
topk_weights,
fc1_c,
route_metadata,
)
test_output.backward(grad_output)
torch.cuda.synchronize()

torch.testing.assert_close(test_output.float(), ref_output.float(), atol=5e-2, rtol=5e-2)
torch.testing.assert_close(test_tokens.grad.float(), ref_grads[0], atol=5e-2, rtol=5e-2)
torch.testing.assert_close(
test_topk_weights.grad,
ref_grads[3],
atol=5e-2,
rtol=5e-2,
)
for expert_idx in range(num_experts):
torch.testing.assert_close(
getattr(fc1, f"weight{expert_idx}").grad.float(),
ref_grads[1][expert_idx].transpose(0, 1),
atol=5e-2,
rtol=5e-2,
)
torch.testing.assert_close(
getattr(fc2, f"weight{expert_idx}").grad.float(),
ref_grads[2][expert_idx].transpose(0, 1),
atol=5e-2,
rtol=5e-2,
)

@_zero_copy_test_include
@_eager_test_include
def test_combine_autograd(self):
Expand Down
Loading
Loading