Skip to content
Draft
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
2 changes: 1 addition & 1 deletion 3rdparty/nccl-extensions
140 changes: 131 additions & 9 deletions tests/cpp_distributed/test_ep.cu
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
* EP pipeline tests: smallest-scope first.
*
* EPDispatchTest/PrepareAndDispatch : exact recv values + per-expert counts
* EPOverflowTest/OverflowDrop : drop_on_overflow clamps recv counts
* EPCombineTest/Combine : round-trip: out == top_k * tokens
* EPCombineBwdTest/CombineBwdCheck : exact grad_expert values
* EPDispatchBwdTest/DispatchBwdCheck : exact grad_tokens
Expand Down Expand Up @@ -303,17 +304,29 @@ TYPED_TEST(EPDispatchTest, PrepareAndDispatch) {
this->template upload_inputs<Tok>(buf);
EPTensors<Tok> t(buf, num_tokens_, top_k_, hidden_dim_, num_local_experts_);

NVTE_CHECK_CUDA(cudaMemset(buf.recv_tokens.get(), 0, buf.recv_tokens.bytes()));

cudaStream_t stream;
NVTE_CHECK_CUDA(cudaStreamCreate(&stream));

ASSERT_NO_THROW(nvte_ep_prepare(t.handle_mem.data(), t.topk_idx.data(), t.recv_tokens_per_expert.data(), nullptr, &t.layer_cfg_, stream));
ASSERT_NO_THROW(nvte_ep_dispatch(t.handle_mem.data(), t.topk_idx.data(),
t.tokens.data(), NVTECommWindow{}, t.topk_weights.data(),
NVTECommWindow{}, t.recv_tokens.data(), NVTECommWindow{},
t.recv_topk_weights.data(), NVTECommWindow{}, stream));
NVTE_CHECK_CUDA(cudaStreamSynchronize(stream));
// Run the same identity check for both dispatch entry points: the separate
// prepare + dispatch, and the fused prepare_and_dispatch (one call seeds
// routing and dispatches; the dispatch writes recv_tokens_per_expert).
for (bool fused : {false, true}) {
SCOPED_TRACE(fused ? "fused prepare_and_dispatch" : "separate prepare + dispatch");
NVTE_CHECK_CUDA(cudaMemset(buf.recv_tokens.get(), 0, buf.recv_tokens.bytes()));
if (fused) {
ASSERT_NO_THROW(nvte_ep_prepare_and_dispatch(
t.handle_mem.data(), t.topk_idx.data(), t.tokens.data(), NVTECommWindow{},
t.topk_weights.data(), NVTECommWindow{}, t.recv_tokens.data(), NVTECommWindow{},
t.recv_topk_weights.data(), NVTECommWindow{}, t.recv_tokens_per_expert.data(),
&t.layer_cfg_, stream));
} else {
ASSERT_NO_THROW(nvte_ep_prepare(t.handle_mem.data(), t.topk_idx.data(), t.recv_tokens_per_expert.data(), nullptr, &t.layer_cfg_, stream));
ASSERT_NO_THROW(nvte_ep_dispatch(t.handle_mem.data(), t.topk_idx.data(),
t.tokens.data(), NVTECommWindow{}, t.topk_weights.data(),
NVTECommWindow{}, t.recv_tokens.data(), NVTECommWindow{},
t.recv_topk_weights.data(), NVTECommWindow{}, stream));
}
NVTE_CHECK_CUDA(cudaStreamSynchronize(stream));

// 1. Per-expert counts.
std::vector<int32_t> got_counts(num_local_experts_);
Expand Down Expand Up @@ -363,7 +376,116 @@ TYPED_TEST(EPDispatchTest, PrepareAndDispatch) {
EXPECT_NEAR(h_w[i], exp_w, 1e-6f) << "recv_topk_weights[" << i << "]";

if (g_process_id == 0)
printf(" PrepareAndDispatch: passed (recv=%d, values + weights exact)\n", total_recv);
printf(" PrepareAndDispatch[%s]: passed (recv=%d, values + weights exact)\n",
fused ? "fused" : "separate", total_recv);
}

NVTE_CHECK_CUDA(cudaStreamDestroy(stream));
}

// =============================================================================
// EPOverflowTest: drop_on_overflow clamps recv counts and completes.
// =============================================================================

template <typename T> class EPOverflowTest : public EpOpTestBase {
protected:
void TearDown() override {
if (g_ep_initialized) ep_reinitialize(/*zero_copy=*/0);
}
};
TYPED_TEST_SUITE(EPOverflowTest, EPBf16Only);

// With drop_on_overflow and a recv budget below the routed total, dispatch keeps
// the first recv_cap tokens (per-expert counts clamped on the cumulative sum)
// and drops the rest instead of trapping. Covers both the separate prepare +
// dispatch and the fused prepare_and_dispatch.
TYPED_TEST(EPOverflowTest, OverflowDrop) {
using Tok = TypeParam;
EP_PULL_FIXTURE();

// Undersize the recv budget so balanced routing overflows it. HT requires the
// recv budget >= send budget, so lower both to the per-rank token count.
const int send_budget = num_tokens_;
const int recv_cap = num_tokens_;
ep_reinitialize_drop(send_budget, recv_cap);

auto exp_counts = expected_recv_tokens_per_expert(g_process_id, g_num_processes, num_tokens_,
top_k_, num_experts_, num_local_experts_);
int pre_drop_total = 0;
for (int c : exp_counts) pre_drop_total += c;
ASSERT_GT(pre_drop_total, recv_cap)
<< "test misconfigured: routing must overflow the recv budget";

EPBuffers<Tok> buf;
buf.alloc(num_tokens_, top_k_, hidden_dim_, num_local_experts_,
ep_size_, max_tokens_per_rank_);
this->template upload_inputs<Tok>(buf);
EPTensors<Tok> t(buf, num_tokens_, top_k_, hidden_dim_, num_local_experts_);

// Device scalar for the pre-drop recv total reported by prepare.
DevBuf<int32_t> total_recv_dev(1);
TensorWrapper total_recv(total_recv_dev.get(), std::vector<size_t>{1}, DType::kInt32);

auto exp_vals = expected_recv_values_sorted<Tok>(g_process_id, g_num_processes, num_tokens_,
top_k_, num_experts_, num_local_experts_);

cudaStream_t stream;
NVTE_CHECK_CUDA(cudaStreamCreate(&stream));

for (bool fused : {false, true}) {
SCOPED_TRACE(fused ? "fused prepare_and_dispatch" : "separate prepare + dispatch");
NVTE_CHECK_CUDA(cudaMemset(buf.recv_tokens.get(), 0, buf.recv_tokens.bytes()));

if (fused) {
ASSERT_NO_THROW(nvte_ep_prepare_and_dispatch(
t.handle_mem.data(), t.topk_idx.data(), t.tokens.data(), NVTECommWindow{},
t.topk_weights.data(), NVTECommWindow{}, t.recv_tokens.data(), NVTECommWindow{},
t.recv_topk_weights.data(), NVTECommWindow{}, t.recv_tokens_per_expert.data(),
&t.layer_cfg_, stream));
} else {
ASSERT_NO_THROW(nvte_ep_prepare(t.handle_mem.data(), t.topk_idx.data(),
t.recv_tokens_per_expert.data(), total_recv.data(),
&t.layer_cfg_, stream));
ASSERT_NO_THROW(nvte_ep_dispatch(t.handle_mem.data(), t.topk_idx.data(),
t.tokens.data(), NVTECommWindow{}, t.topk_weights.data(),
NVTECommWindow{}, t.recv_tokens.data(), NVTECommWindow{},
t.recv_topk_weights.data(), NVTECommWindow{}, stream));
}
NVTE_CHECK_CUDA(cudaStreamSynchronize(stream));

// Per-expert counts are clamped so their sum is exactly the recv budget (the
// cumulative clamp), not the larger pre-drop total.
std::vector<int32_t> got_counts(num_local_experts_);
NVTE_CHECK_CUDA(cudaMemcpy(got_counts.data(), buf.recv_tokens_per_expert.get(),
num_local_experts_ * sizeof(int32_t), cudaMemcpyDeviceToHost));
int kept = 0;
for (int c : got_counts) { EXPECT_GE(c, 0); kept += c; }
EXPECT_EQ(kept, recv_cap) << "kept token count must equal the recv budget";

// The separate path's prepare reports the true (pre-drop) recv total.
if (!fused) {
int32_t pre_drop = 0;
NVTE_CHECK_CUDA(cudaMemcpy(&pre_drop, total_recv_dev.get(), sizeof(int32_t),
cudaMemcpyDeviceToHost));
EXPECT_EQ(pre_drop, pre_drop_total) << "prepare must report the pre-drop recv total";
}

// Every kept recv slot holds a real routed token value: no garbage past the
// clamp. Slabs are packed contiguously (alignment disabled).
std::vector<Tok> h_recv(buf.recv_capacity * hidden_dim_);
NVTE_CHECK_CUDA(cudaMemcpy(h_recv.data(), buf.recv_tokens.get(),
h_recv.size() * sizeof(Tok), cudaMemcpyDeviceToHost));
size_t slot = 0;
for (int e = 0; e < num_local_experts_; ++e)
for (int i = 0; i < got_counts[e]; ++i, ++slot)
EXPECT_TRUE(std::binary_search(exp_vals.begin(), exp_vals.end(),
tok_to_float(h_recv[slot * hidden_dim_])))
<< "kept recv slot " << slot << " is not a routed token value";

if (g_process_id == 0)
printf(" OverflowDrop[%s]: passed (kept=%d, pre-drop=%d)\n",
fused ? "fused" : "separate", kept, pre_drop_total);
}

NVTE_CHECK_CUDA(cudaStreamDestroy(stream));
}
Expand Down
18 changes: 18 additions & 0 deletions tests/cpp_distributed/test_ep_common.h
Original file line number Diff line number Diff line change
Expand Up @@ -184,6 +184,24 @@ static void ep_reinitialize(int zero_copy) {
nvte_ep_initialize(static_cast<void*>(g_ep_comm), &group_config);
}

// Re-initialize the EP backend with a small recv budget and drop-on-overflow,
// to exercise the count-mode overflow-drop path. HT requires the recv budget to
// be at least the send budget, so both are passed in. Restore the default group
// with ep_reinitialize(0).
static void ep_reinitialize_drop(int max_tokens_per_rank, int max_recv_tokens_per_rank) {
if (!g_ep_initialized) return;
nvte_ep_shutdown();
NVTEEpGroupConfig group_config = NVTE_EP_GROUP_CONFIG_INIT;
group_config.ep_size = g_ep_size;
group_config.num_experts = g_num_experts;
group_config.max_tokens_per_rank = max_tokens_per_rank;
group_config.max_recv_tokens_per_rank = max_recv_tokens_per_rank;
group_config.hidden_dim = g_hidden_dim;
group_config.max_token_dtype = g_max_token_dtype;
group_config.drop_on_overflow = 1;
nvte_ep_initialize(static_cast<void*>(g_ep_comm), &group_config);
}

// Tear down in dependency order: backend's ep_group reads from ep_comm,
// so destroy the group first, then the comm.
static void ep_teardown() {
Expand Down
29 changes: 29 additions & 0 deletions tests/pytorch/distributed/run_ep.py
Original file line number Diff line number Diff line change
Expand Up @@ -265,6 +265,35 @@ def test_eager_recv_sizing(self):
self.assertGreaterEqual(total, int(tokens_per_expert.sum().item()))
self.assertLessEqual(total, self.cfg.recv_capacity_per_rank)

def test_fused_prepare_and_dispatch(self):
"""Non-eager ep_dispatch fuses prepare+dispatch: one call seeds routing,
dispatches, and writes the buffer-owned int64 tokens_per_expert. Checks the
per-expert counts, recv sizing, and that grad flows through the fused fwd.
Runs only in the default (count) pass; the mode passes skip it in setUp."""
buf = self._make_buffer()
topk_idx, tokens, w = _make_identity_inputs(self.cfg.rank, self.cfg.ep_size)
tokens_p = tokens.detach().clone().requires_grad_(True)
recv_t, recv_w, tokens_per_expert = ep_dispatch(buf, tokens_p, topk_idx, w)
torch.cuda.synchronize()
# The fused op writes the buffer-owned int64 per-expert counts and returns them.
self.assertEqual(tokens_per_expert.data_ptr(), buf.tokens_per_expert.data_ptr())
self.assertEqual(tokens_per_expert.dtype, torch.int64)
self.assertEqual(tokens_per_expert.shape, (NUM_LOCAL_EXPERTS,))
# Non-eager sizes recv outputs to recv_capacity_per_rank, not the recv total.
self.assertEqual(recv_t.shape[0], self.cfg.recv_capacity_per_rank)
self.assertEqual(recv_w.shape[0], self.cfg.recv_capacity_per_rank)
# Per-expert counts sum, across ranks, to every routed (token, top-k) pair.
local = int(tokens_per_expert.sum().item())
total = torch.tensor([local], dtype=torch.int64, device=self.cfg.device)
dist.all_reduce(total, op=dist.ReduceOp.SUM, group=self.ep_group)
self.assertEqual(int(total.item()), self.cfg.world_size * TOKENS_PER_RANK * TOP_K)
# Round-trip identity and grad flow through the fused forward.
out = ep_combine(buf, self._weighted(recv_t, recv_w))
(0.5 * (out.float() ** 2).sum()).backward()
torch.cuda.synchronize()
torch.testing.assert_close(out.float(), tokens.float(), atol=5e-2, rtol=5e-2)
torch.testing.assert_close(tokens_p.grad.float(), tokens.float(), atol=5e-2, rtol=5e-2)

@_overflow_test_include
def test_overflow_drop(self):
"""drop_on_overflow: recv past capacity is dropped and dispatch continues
Expand Down
25 changes: 25 additions & 0 deletions transformer_engine/common/ep/ep_api.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -90,6 +90,20 @@ void nvte_ep_dispatch(NVTETensor handle_mem, NVTETensor topk_idx, NVTETensor tok
recv_topk_weights_win, stream);
}

void nvte_ep_prepare_and_dispatch(NVTETensor handle_mem, NVTETensor topk_idx, NVTETensor tokens,
NVTECommWindow tokens_win, NVTETensor topk_weights,
NVTECommWindow topk_weights_win, NVTETensor recv_tokens,
NVTECommWindow recv_tokens_win, NVTETensor recv_topk_weights,
NVTECommWindow recv_topk_weights_win,
NVTETensor recv_tokens_per_expert,
const NVTEEpLayerConfig* layer_cfg, cudaStream_t stream) {
NVTEEpLayerConfig cfg = normalize_ep_config(layer_cfg, kLayerConfigMinSize, "layer_cfg");
EPBackend::get().prepare_and_dispatch(handle_mem_ptr(handle_mem), topk_idx, tokens, tokens_win,
topk_weights, topk_weights_win, recv_tokens,
recv_tokens_win, recv_topk_weights, recv_topk_weights_win,
recv_tokens_per_expert, cfg, stream);
}

void nvte_ep_combine(NVTETensor handle_mem, NVTETensor expert_out, NVTECommWindow expert_out_win,
NVTETensor result, cudaStream_t stream) {
EPBackend::get().combine(handle_mem_ptr(handle_mem), expert_out, expert_out_win, result, stream);
Expand Down Expand Up @@ -143,6 +157,17 @@ void nvte_ep_dispatch(NVTETensor /*handle_mem*/, NVTETensor /*topk_idx*/, NVTETe
ep_not_built();
}

void nvte_ep_prepare_and_dispatch(NVTETensor /*handle_mem*/, NVTETensor /*topk_idx*/,
NVTETensor /*tokens*/, NVTECommWindow /*tokens_win*/,
NVTETensor /*topk_weights*/, NVTECommWindow /*topk_weights_win*/,
NVTETensor /*recv_tokens*/, NVTECommWindow /*recv_tokens_win*/,
NVTETensor /*recv_topk_weights*/,
NVTECommWindow /*recv_topk_weights_win*/,
NVTETensor /*recv_tokens_per_expert*/,
const NVTEEpLayerConfig* /*layer_cfg*/, cudaStream_t /*stream*/) {
ep_not_built();
}

void nvte_ep_combine(NVTETensor /*handle_mem*/, NVTETensor /*expert_out*/,
NVTECommWindow /*expert_out_win*/, NVTETensor /*result*/,
cudaStream_t /*stream*/) {
Expand Down
68 changes: 60 additions & 8 deletions transformer_engine/common/ep/ep_backend.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -365,12 +365,14 @@ void EPBackend::prepare(void* handle_mem, const NVTETensor topk_idx,
NVTE_CHECK_NCCL(ncclEpUpdateHandle(h, &nccl_topk_idx, &layout_info, stream));
}

void EPBackend::dispatch(void* handle_mem, const NVTETensor topk_idx, const NVTETensor tokens,
const NVTECommWindow& tokens_win, const NVTETensor topk_weights,
const NVTECommWindow& topk_weights_win, NVTETensor recv_tokens,
const NVTECommWindow& recv_tokens_win, NVTETensor recv_topk_weights,
const NVTECommWindow& recv_topk_weights_win, cudaStream_t stream) {
NVTE_CHECK(handle_mem != nullptr, "handle_mem must not be null");
void EPBackend::issue_dispatch_locked(ncclEpHandle_t handle, const NVTETensor topk_idx,
const NVTETensor tokens, const NVTECommWindow& tokens_win,
const NVTETensor topk_weights,
const NVTECommWindow& topk_weights_win,
NVTETensor recv_tokens, const NVTECommWindow& recv_tokens_win,
NVTETensor recv_topk_weights,
const NVTECommWindow& recv_topk_weights_win,
NVTETensor recv_tokens_per_expert, cudaStream_t stream) {
NVTE_CHECK(nvte_tensor_shape(tokens).ndim == 2, "tokens must be 2D [T, hidden_dim]");
NVTE_CHECK(nvte_tensor_shape(recv_tokens).ndim == 2,
"recv_tokens must be 2D [recv_T, hidden_dim]");
Expand Down Expand Up @@ -422,11 +424,61 @@ void EPBackend::dispatch(void* handle_mem, const NVTETensor topk_idx, const NVTE
ncclEpDispatchConfig_t dispatch_cfg = NCCL_EP_DISPATCH_CONFIG_INIT;
dispatch_cfg.pass_direction = is_forward ? NCCL_EP_FWD_PASS : NCCL_EP_BWD_PASS;

// Count mode: wire the caller's per-expert counts tensor so the dispatch
// writes recv counts. NULL leaves layout_info unset (non-count callers).
ncclEpLayoutInfo_t layout_info = NCCL_EP_LAYOUT_INFO_INIT;
NVTEShape recv_counts_shape;
ncclEpTensor_t recv_counts_desc;
const ncclEpLayoutInfo_t* layout_info_ptr = nullptr;
if (recv_tokens_per_expert != nullptr) {
recv_counts_desc = make_nccl_ep_tensor(recv_tokens_per_expert, recv_counts_shape);
layout_info.expert_counters = &recv_counts_desc;
layout_info_ptr = &layout_info;
}

NVTE_CHECK_NCCL(
ncclEpDispatch(handle, &in_struct, &out_struct, layout_info_ptr, &dispatch_cfg, stream));
}

void EPBackend::dispatch(void* handle_mem, const NVTETensor topk_idx, const NVTETensor tokens,
const NVTECommWindow& tokens_win, const NVTETensor topk_weights,
const NVTECommWindow& topk_weights_win, NVTETensor recv_tokens,
const NVTECommWindow& recv_tokens_win, NVTETensor recv_topk_weights,
const NVTECommWindow& recv_topk_weights_win, cudaStream_t stream) {
NVTE_CHECK(handle_mem != nullptr, "handle_mem must not be null");
std::lock_guard<std::mutex> lock(mutex_);
NVTE_CHECK(initialized_, "EPBackend not initialized");
ncclEpHandle_t h = lookup_handle_locked(handle_mem);
NVTE_CHECK_NCCL(ncclEpDispatch(h, &in_struct, &out_struct,
/*layout_info=*/nullptr, &dispatch_cfg, stream));
issue_dispatch_locked(h, topk_idx, tokens, tokens_win, topk_weights, topk_weights_win,
recv_tokens, recv_tokens_win, recv_topk_weights, recv_topk_weights_win,
/*recv_tokens_per_expert=*/nullptr, stream);
}

void EPBackend::prepare_and_dispatch(void* handle_mem, const NVTETensor topk_idx,
const NVTETensor tokens, const NVTECommWindow& tokens_win,
const NVTETensor topk_weights,
const NVTECommWindow& topk_weights_win, NVTETensor recv_tokens,
const NVTECommWindow& recv_tokens_win,
NVTETensor recv_topk_weights,
const NVTECommWindow& recv_topk_weights_win,
NVTETensor recv_tokens_per_expert, NVTEEpLayerConfig layer_cfg,
cudaStream_t stream) {
NVTE_CHECK(handle_mem != nullptr, "handle_mem must not be null");
NVTE_CHECK(layer_cfg.top_k > 0, "top_k must be > 0, got ", layer_cfg.top_k);
NVTE_CHECK(nvte_tensor_shape(topk_idx).ndim == 2, "topk_idx must be 2D [T, top_k]");

NVTEShape topk_idx_shape;
ncclEpTensor_t nccl_topk_idx = make_nccl_ep_tensor(topk_idx, topk_idx_shape);

std::lock_guard<std::mutex> lock(mutex_);
NVTE_CHECK(initialized_, "EPBackend not initialized");
ncclEpHandle_t h = prepare_handle_locked(handle_mem, layer_cfg);
// Count mode: UpdateHandle only seeds routing; the dispatch produces the
// per-expert recv counts into recv_tokens_per_expert.
NVTE_CHECK_NCCL(ncclEpUpdateHandle(h, &nccl_topk_idx, /*layout_info=*/nullptr, stream));
issue_dispatch_locked(h, topk_idx, tokens, tokens_win, topk_weights, topk_weights_win,
recv_tokens, recv_tokens_win, recv_topk_weights, recv_topk_weights_win,
recv_tokens_per_expert, stream);
}

void EPBackend::combine(void* handle_mem, const NVTETensor expert_out,
Expand Down
Loading
Loading