From 171c97da152603088500f246e912cd96d9ca55b5 Mon Sep 17 00:00:00 2001 From: Phuong Nguyen Date: Fri, 10 Jul 2026 06:27:24 -0700 Subject: [PATCH] fused prepare and dispatch Signed-off-by: Phuong Nguyen --- 3rdparty/nccl-extensions | 2 +- tests/cpp_distributed/test_ep.cu | 140 ++++++++++++- tests/cpp_distributed/test_ep_common.h | 18 ++ tests/pytorch/distributed/run_ep.py | 29 +++ transformer_engine/common/ep/ep_api.cpp | 25 +++ transformer_engine/common/ep/ep_backend.cpp | 68 ++++++- transformer_engine/common/ep/ep_backend.h | 22 +++ .../common/include/transformer_engine/ep.h | 30 +++ transformer_engine/pytorch/csrc/extensions.h | 5 + .../pytorch/csrc/extensions/ep.cpp | 66 +++++++ transformer_engine/pytorch/ep.py | 185 +++++++++++++++--- 11 files changed, 542 insertions(+), 48 deletions(-) diff --git a/3rdparty/nccl-extensions b/3rdparty/nccl-extensions index 2c6135a721..e7c7d2a83c 160000 --- a/3rdparty/nccl-extensions +++ b/3rdparty/nccl-extensions @@ -1 +1 @@ -Subproject commit 2c6135a721824ff792af7b72900b0ab758fa1f98 +Subproject commit e7c7d2a83ce7945b7e0ff31894ab2a715be85d09 diff --git a/tests/cpp_distributed/test_ep.cu b/tests/cpp_distributed/test_ep.cu index 47732196ed..f5a9282249 100644 --- a/tests/cpp_distributed/test_ep.cu +++ b/tests/cpp_distributed/test_ep.cu @@ -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 @@ -303,17 +304,29 @@ TYPED_TEST(EPDispatchTest, PrepareAndDispatch) { this->template upload_inputs(buf); EPTensors 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 got_counts(num_local_experts_); @@ -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 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 buf; + buf.alloc(num_tokens_, top_k_, hidden_dim_, num_local_experts_, + ep_size_, max_tokens_per_rank_); + this->template upload_inputs(buf); + EPTensors t(buf, num_tokens_, top_k_, hidden_dim_, num_local_experts_); + + // Device scalar for the pre-drop recv total reported by prepare. + DevBuf total_recv_dev(1); + TensorWrapper total_recv(total_recv_dev.get(), std::vector{1}, DType::kInt32); + + auto exp_vals = expected_recv_values_sorted(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 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 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)); } diff --git a/tests/cpp_distributed/test_ep_common.h b/tests/cpp_distributed/test_ep_common.h index 7cf6017090..68354f99aa 100644 --- a/tests/cpp_distributed/test_ep_common.h +++ b/tests/cpp_distributed/test_ep_common.h @@ -184,6 +184,24 @@ static void ep_reinitialize(int zero_copy) { nvte_ep_initialize(static_cast(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(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() { diff --git a/tests/pytorch/distributed/run_ep.py b/tests/pytorch/distributed/run_ep.py index 4778498b7d..9407d6a386 100644 --- a/tests/pytorch/distributed/run_ep.py +++ b/tests/pytorch/distributed/run_ep.py @@ -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 diff --git a/transformer_engine/common/ep/ep_api.cpp b/transformer_engine/common/ep/ep_api.cpp index 0981289ffe..09f33c2713 100644 --- a/transformer_engine/common/ep/ep_api.cpp +++ b/transformer_engine/common/ep/ep_api.cpp @@ -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); @@ -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*/) { diff --git a/transformer_engine/common/ep/ep_backend.cpp b/transformer_engine/common/ep/ep_backend.cpp index dd44773514..41915126f9 100644 --- a/transformer_engine/common/ep/ep_backend.cpp +++ b/transformer_engine/common/ep/ep_backend.cpp @@ -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]"); @@ -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 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 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, diff --git a/transformer_engine/common/ep/ep_backend.h b/transformer_engine/common/ep/ep_backend.h index 80c9b9cea3..26b5a0f50d 100644 --- a/transformer_engine/common/ep/ep_backend.h +++ b/transformer_engine/common/ep/ep_backend.h @@ -57,6 +57,16 @@ class EPBackend { const NVTECommWindow& recv_tokens_win, NVTETensor recv_topk_weights, const NVTECommWindow& recv_topk_weights_win, cudaStream_t stream); + // Fused prepare + dispatch: seeds routing then dispatches in one call. + // Per-expert recv counts are written to recv_tokens_per_expert by the dispatch. + void 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); + void combine(void* handle_mem, const NVTETensor expert_out, const NVTECommWindow& expert_out_win, NVTETensor result, cudaStream_t stream); @@ -109,6 +119,18 @@ class EPBackend { ncclEpHandle_t prepare_handle_locked(void* handle_mem, NVTEEpLayerConfig layer_cfg); ncclEpHandle_t lookup_handle_locked(void* handle_mem); size_t cache_cap_locked(); + + // Build the dispatch in/out structs and issue ncclEpDispatch on the resolved + // handle. When recv_tokens_per_expert != nullptr (count mode), it is wired to + // layout_info.expert_counters so the dispatch writes per-expert recv counts. + // Caller must hold mutex_. + void 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); }; } // namespace ep diff --git a/transformer_engine/common/include/transformer_engine/ep.h b/transformer_engine/common/include/transformer_engine/ep.h index a4800944b5..1299f77b0d 100644 --- a/transformer_engine/common/include/transformer_engine/ep.h +++ b/transformer_engine/common/include/transformer_engine/ep.h @@ -168,6 +168,36 @@ void nvte_ep_dispatch(NVTETensor handle_mem, NVTETensor topk_idx, NVTETensor tok NVTECommWindow recv_tokens_win, NVTETensor recv_topk_weights, NVTECommWindow recv_topk_weights_win, cudaStream_t stream); +/*! \brief Fused prepare + dispatch. + * + * Seeds handle_mem with this step's routing (as nvte_ep_prepare) and then + * dispatches tokens (as nvte_ep_dispatch) in a single call. Per-expert recv + * counts are produced by the dispatch and written to recv_tokens_per_expert; + * they are not available before the dispatch returns. No separate + * nvte_ep_prepare is needed. + * + * \param[in] handle_mem uint8 routing-state buffer. + * \param[in] topk_idx [T, top_k] int64 routing indices. + * \param[in] tokens [T, hidden_dim] input tokens. + * \param[in] tokens_win Optional symmem window for tokens. + * \param[in] topk_weights [T, top_k] float32 weights. + * \param[in] topk_weights_win Optional symmem window for topk_weights. + * \param[out] recv_tokens [recv_T, hidden_dim] received tokens. + * \param[in] recv_tokens_win Optional symmem window for recv_tokens. + * \param[out] recv_topk_weights [recv_T] float32 per-slot weights. + * \param[in] recv_topk_weights_win Optional symmem window for recv_topk_weights. + * \param[out] recv_tokens_per_expert [num_local_experts] int32/int64 counts. + * \param[in] layer_cfg Per-call layer configuration (struct_size set). + * \param[in] stream CUDA 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); + /*! \brief Scatter-sum expert outputs back to originating ranks. * * Inverse of dispatch: the top_k destination slots for token t are summed diff --git a/transformer_engine/pytorch/csrc/extensions.h b/transformer_engine/pytorch/csrc/extensions.h index 95f02c64f0..8861215209 100644 --- a/transformer_engine/pytorch/csrc/extensions.h +++ b/transformer_engine/pytorch/csrc/extensions.h @@ -693,6 +693,11 @@ void ep_prepare(at::Tensor handle_mem, at::Tensor topk_idx, at::Tensor tokens_pe void ep_dispatch(at::Tensor handle_mem, at::Tensor topk_idx, at::Tensor tokens, at::Tensor topk_weights, at::Tensor recv_tokens, at::Tensor recv_topk_weights); +void ep_prepare_and_dispatch(at::Tensor handle_mem, at::Tensor topk_idx, at::Tensor tokens, + at::Tensor topk_weights, at::Tensor recv_tokens, + at::Tensor recv_topk_weights, at::Tensor token_counts, int64_t top_k, + int64_t dispatch_output_per_expert_alignment); + void ep_combine(at::Tensor handle_mem, at::Tensor expert_out, at::Tensor result); void ep_dispatch_bwd(at::Tensor handle_mem, at::Tensor grad, at::Tensor g_recv_topk_weights, diff --git a/transformer_engine/pytorch/csrc/extensions/ep.cpp b/transformer_engine/pytorch/csrc/extensions/ep.cpp index 9bf821721f..2021b16a3e 100644 --- a/transformer_engine/pytorch/csrc/extensions/ep.cpp +++ b/transformer_engine/pytorch/csrc/extensions/ep.cpp @@ -274,6 +274,70 @@ void ep_dispatch(at::Tensor handle_mem, at::Tensor topk_idx, at::Tensor tokens, recv_topk_w_te.data(), recv_topk_w_win, stream); } +void ep_prepare_and_dispatch(at::Tensor handle_mem, at::Tensor topk_idx, at::Tensor tokens, + at::Tensor topk_weights, at::Tensor recv_tokens, + at::Tensor recv_topk_weights, at::Tensor token_counts, int64_t top_k, + int64_t dispatch_output_per_expert_alignment) { + auto stream = at::cuda::getCurrentCUDAStream().stream(); + NVTE_CHECK(tokens.dim() >= 2, "tokens must be at least 2D [..., H]"); + NVTE_CHECK(topk_idx.dim() >= 2, "topk_idx must be at least 2D [..., top_k]"); + NVTE_CHECK(topk_weights.dim() >= 2, "topk_weights must be at least 2D [..., top_k]"); + NVTE_CHECK(recv_tokens.dim() >= 2, "recv_tokens must be at least 2D [..., recv_pr, H]"); + auto idx_dtype = check_topk_idx_dtype(topk_idx); + NVTE_CHECK(tokens.is_contiguous(), "tokens must be contiguous"); + NVTE_CHECK(topk_weights.is_contiguous(), "topk_weights must be contiguous"); + NVTE_CHECK(recv_tokens.is_contiguous(), "recv_tokens must be contiguous"); + NVTE_CHECK(recv_topk_weights.is_contiguous(), "recv_topk_weights must be contiguous"); + NVTE_CHECK(token_counts.is_contiguous(), "token_counts must be contiguous"); + NVTE_CHECK(token_counts.scalar_type() == at::kInt || token_counts.scalar_type() == at::kLong, + "token_counts must be int32 or int64"); + + const size_t H = static_cast(tokens.size(-1)); + const size_t T_flat = tokens.numel() / H; + const size_t topk_n = static_cast(topk_idx.size(-1)); + const size_t recv_pr = recv_tokens.numel() / H; + + NVTE_CHECK(static_cast(topk_weights.size(-1)) == topk_n, + "topk_weights last dim must equal topk_idx last dim"); + NVTE_CHECK(static_cast(topk_idx.numel()) == T_flat * topk_n, + "topk_idx token count must equal tokens token count"); + NVTE_CHECK(static_cast(topk_weights.numel()) == T_flat * topk_n, + "topk_weights token count must equal tokens token count"); + NVTE_CHECK(static_cast(recv_topk_weights.numel()) == recv_pr, + "recv_topk_weights total size must equal recv_tokens recv_pr"); + NVTE_CHECK(recv_tokens.scalar_type() == tokens.scalar_type(), "recv_tokens dtype (", + c10::toString(recv_tokens.scalar_type()), ") must match tokens dtype (", + c10::toString(tokens.scalar_type()), ")"); + check_symm_mem_required(recv_tokens, "recv_tokens"); + check_symm_mem_required(recv_topk_weights, "recv_topk_weights"); + + auto tok_dtype = GetTransformerEngineDType(tokens.scalar_type()); + auto handle_mem_te = makeTransformerEngineTensor( + handle_mem.data_ptr(), Shape{static_cast(handle_mem.numel())}, DType::kByte); + auto topk_idx_te = + makeTransformerEngineTensor(topk_idx.data_ptr(), Shape{T_flat, topk_n}, idx_dtype); + auto tokens_te = makeTransformerEngineTensor(tokens.data_ptr(), Shape{T_flat, H}, tok_dtype); + auto topk_w_te = + makeTransformerEngineTensor(topk_weights.data_ptr(), Shape{T_flat, topk_n}, DType::kFloat32); + auto recv_tokens_te = + makeTransformerEngineTensor(recv_tokens.data_ptr(), Shape{recv_pr, H}, tok_dtype); + auto recv_topk_w_te = + makeTransformerEngineTensor(recv_topk_weights.data_ptr(), Shape{recv_pr}, DType::kFloat32); + auto token_counts_te = makeTransformerEngineTensor( + token_counts.data_ptr(), Shape{static_cast(token_counts.numel())}, + GetTransformerEngineDType(token_counts.scalar_type())); + + NVTECommWindow tokens_win = maybe_make_window(tokens); + NVTECommWindow topk_w_win = maybe_make_window(topk_weights); + NVTECommWindow recv_tokens_win = maybe_make_window(recv_tokens); + NVTECommWindow recv_topk_w_win = maybe_make_window(recv_topk_weights); + auto layer_cfg = make_layer_cfg(top_k, dispatch_output_per_expert_alignment); + nvte_ep_prepare_and_dispatch(handle_mem_te.data(), topk_idx_te.data(), tokens_te.data(), + tokens_win, topk_w_te.data(), topk_w_win, recv_tokens_te.data(), + recv_tokens_win, recv_topk_w_te.data(), recv_topk_w_win, + token_counts_te.data(), &layer_cfg, stream); +} + void ep_combine(at::Tensor handle_mem, at::Tensor expert_out, at::Tensor result) { auto stream = at::cuda::getCurrentCUDAStream().stream(); NVTE_CHECK(expert_out.dim() >= 2, "expert_out must be at least 2D [..., recv_pr, H]"); @@ -397,6 +461,8 @@ void register_ep_bindings(pybind11::module_& m) { py::arg("dispatch_output_per_expert_alignment"), py::arg("total_recv_tokens"), py::call_guard()); m.def("ep_dispatch", &ep_dispatch, "EP dispatch", py::call_guard()); + m.def("ep_prepare_and_dispatch", &ep_prepare_and_dispatch, "EP fused prepare + dispatch", + py::call_guard()); m.def("ep_combine", &ep_combine, "EP combine", py::call_guard()); m.def("ep_dispatch_bwd", &ep_dispatch_bwd, "EP dispatch backward", py::call_guard()); diff --git a/transformer_engine/pytorch/ep.py b/transformer_engine/pytorch/ep.py index 940812b6a4..cd04408388 100644 --- a/transformer_engine/pytorch/ep.py +++ b/transformer_engine/pytorch/ep.py @@ -332,6 +332,40 @@ def _(*_args, **_kw): return None +@torch.library.custom_op( + f"{_LIB}::prepare_and_dispatch", + mutates_args=("recv_tokens", "recv_topk_weights", "token_counts"), + device_types="cuda", +) +def _prepare_and_dispatch_op( + handle_mem: torch.Tensor, + top_k: int, + alignment: int, + topk_idx: torch.Tensor, + tokens: torch.Tensor, + topk_weights: torch.Tensor, + recv_tokens: torch.Tensor, + recv_topk_weights: torch.Tensor, + token_counts: torch.Tensor, +) -> None: + tex.ep_prepare_and_dispatch( + handle_mem, + topk_idx, + tokens, + topk_weights, + recv_tokens, + recv_topk_weights, + token_counts, + top_k, + alignment, + ) + + +@_prepare_and_dispatch_op.register_fake +def _(*_args, **_kw): + return None + + @torch.library.custom_op( f"{_LIB}::combine", mutates_args=("result",), @@ -435,6 +469,26 @@ def _ep_combine_raw(buffer: "EpBuffer", expert_out: torch.Tensor, result: torch. # autograd.Function wrappers +def _dispatch_backward(ctx, g_recv_tokens, g_recv_topk_weights): + """Shared dispatch bwd: run dispatch_bwd and reshape grads to the fwd input layout.""" + (handle_mem,) = ctx.saved_tensors + device = handle_mem.device + g_recv_tokens = g_recv_tokens.contiguous() + g_recv_topk_weights = g_recv_topk_weights.contiguous() + grad_tokens = torch.empty( + ctx.tokens_T_flat, ctx.hidden_dim, dtype=ctx.tokens_dtype, device=device + ) + grad_topk_weights = torch.empty(ctx.topk_T_flat, ctx.top_k, dtype=torch.float32, device=device) + torch.ops.transformer_engine_ep.dispatch_bwd( + handle_mem, + g_recv_tokens, + g_recv_topk_weights, + grad_tokens, + grad_topk_weights, + ) + return grad_tokens.view(ctx.tokens_shape), grad_topk_weights.view(ctx.topk_weights_shape) + + class _EpDispatch(torch.autograd.Function): """Autograd dispatch; caller runs prepare first. bwd uses user-supplied grad inputs as-is.""" @@ -472,30 +526,74 @@ def forward( # type: ignore[override] @staticmethod def backward(ctx, g_recv_tokens, g_recv_topk_weights): # type: ignore[override] """Dispatch bwd; normalizes grad-input layout, otherwise passes through.""" - (handle_mem,) = ctx.saved_tensors - device = handle_mem.device - g_recv_tokens = g_recv_tokens.contiguous() - g_recv_topk_weights = g_recv_topk_weights.contiguous() - grad_tokens = torch.empty( - ctx.tokens_T_flat, ctx.hidden_dim, dtype=ctx.tokens_dtype, device=device - ) - grad_topk_weights = torch.empty( - ctx.topk_T_flat, ctx.top_k, dtype=torch.float32, device=device - ) - torch.ops.transformer_engine_ep.dispatch_bwd( - handle_mem, - g_recv_tokens, - g_recv_topk_weights, + grad_tokens, grad_topk_weights = _dispatch_backward(ctx, g_recv_tokens, g_recv_topk_weights) + return ( + None, # handle_mem + None, # recv_tokens + None, # recv_topk_weights + None, # topk_idx grad_tokens, grad_topk_weights, ) + + +class _EpPrepareAndDispatch(torch.autograd.Function): + """Autograd fused prepare+dispatch. The dispatch seeds routing and + writes the per-expert recv counts, so no separate prepare is needed. bwd uses + user-supplied grad inputs as-is.""" + + @staticmethod + def forward( # type: ignore[override] + ctx, + handle_mem: torch.Tensor, + top_k: int, + alignment: int, + recv_tokens: torch.Tensor, + recv_topk_weights: torch.Tensor, + tokens_per_expert: torch.Tensor, + topk_idx: torch.Tensor, + tokens: torch.Tensor, + topk_weights: torch.Tensor, + ): + """Fused prepare + dispatch fwd; the dispatch writes ``tokens_per_expert``.""" + torch.ops.transformer_engine_ep.prepare_and_dispatch( + handle_mem, + top_k, + alignment, + topk_idx, + tokens, + topk_weights, + recv_tokens, + recv_topk_weights, + tokens_per_expert, + ) + ctx.save_for_backward(handle_mem) + ctx.tokens_shape = tokens.shape + ctx.tokens_dtype = tokens.dtype + ctx.topk_weights_shape = topk_weights.shape + ctx.tokens_T_flat = tokens.numel() // tokens.shape[-1] + ctx.topk_T_flat = topk_weights.numel() // topk_weights.shape[-1] + ctx.top_k = topk_weights.shape[-1] + ctx.hidden_dim = tokens.shape[-1] + ctx.mark_non_differentiable(tokens_per_expert) + # Detach so the long-lived buffers aren't tracked as differentiable outputs; + # autograd re-attaches grad_fn pointing back at this Function. + return recv_tokens.detach(), recv_topk_weights.detach(), tokens_per_expert + + @staticmethod + def backward(ctx, g_recv_tokens, g_recv_topk_weights, _g_tokens_per_expert): # type: ignore[override] + """Dispatch bwd; normalizes grad-input layout, otherwise passes through.""" + grad_tokens, grad_topk_weights = _dispatch_backward(ctx, g_recv_tokens, g_recv_topk_weights) return ( None, # handle_mem + None, # top_k + None, # alignment None, # recv_tokens None, # recv_topk_weights + None, # tokens_per_expert None, # topk_idx - grad_tokens.view(ctx.tokens_shape), - grad_topk_weights.view(ctx.topk_weights_shape), + grad_tokens, + grad_topk_weights, ) @@ -578,6 +676,17 @@ def _alloc_io(shape, dtype: torch.dtype, device, zero_copy: bool) -> torch.Tenso return torch.empty(*shape, dtype=dtype, device=device) +def _size_recv_buffers(buffer, rows, recv_tokens, recv_topk_weights): + """Allocate any dispatch recv output the caller did not supply, sized to ``rows``.""" + if recv_tokens is None: + recv_tokens = _alloc_io( + (rows, buffer.hidden_dim), buffer.payload_dtype, buffer.device, buffer.zero_copy + ) + if recv_topk_weights is None: + recv_topk_weights = _alloc_io((rows,), torch.float32, buffer.device, buffer.zero_copy) + return recv_tokens, recv_topk_weights + + def ep_dispatch( buffer: EpBuffer, tokens: torch.Tensor, @@ -595,6 +704,10 @@ def ep_dispatch( this step's recv-token total and must not be caller-supplied. Returns (recv_tokens, recv_topk_weights, tokens_per_expert); tokens_per_expert is non-diff. See ``buffer.total_recv_tokens`` for the per-step recv total. + + Eager mode host-syncs the recv-token total in a separate prepare so the recv outputs can be + sized. Non-eager mode uses the fused prepare+dispatch: the dispatch seeds routing and writes + tokens_per_expert in one call, and the recv outputs are sized to ``recv_capacity_per_rank``. """ _require_bf16("tokens", tokens) if topk_weights.dtype is not torch.float32: @@ -607,28 +720,40 @@ def ep_dispatch( "eager mode sizes the recv outputs from the per-step recv-token total " "and cannot use caller-supplied recv_tokens / recv_topk_weights" ) - # Prepare (routing AllGather) up front so the recv outputs can be sized; in - # eager mode ep_prepare also host-syncs this step's recv-token total. - tokens_per_expert = ep_prepare(buffer, topk_idx) - rows = buffer._host_total_recv_tokens if buffer.eager else buffer.recv_capacity_per_rank - if recv_tokens is None: - recv_tokens = _alloc_io( - (rows, buffer.hidden_dim), - buffer.payload_dtype, - buffer.device, - buffer.zero_copy, + if buffer.eager: + # Separate prepare + dispatch: eager host-syncs the recv-token total in a + # standalone prepare so the recv outputs can be sized to it. + tokens_per_expert = ep_prepare(buffer, topk_idx) + recv_tokens, recv_topk_weights = _size_recv_buffers( + buffer, buffer._host_total_recv_tokens, recv_tokens, recv_topk_weights ) - if recv_topk_weights is None: - recv_topk_weights = _alloc_io((rows,), torch.float32, buffer.device, buffer.zero_copy) - recv_tokens, recv_topk_weights = _EpDispatch.apply( + recv_tokens, recv_topk_weights = _EpDispatch.apply( + buffer.handle_mem, + recv_tokens, + recv_topk_weights, + topk_idx, + tokens, + topk_weights, + ) + return recv_tokens, recv_topk_weights, tokens_per_expert + + # Non-eager (incl. drop_on_overflow): fused prepare+dispatch, recv outputs sized to + # recv_capacity_per_rank; the fused dispatch drops over-budget tokens when + # drop_on_overflow is enabled, otherwise traps on overflow. + recv_tokens, recv_topk_weights = _size_recv_buffers( + buffer, buffer.recv_capacity_per_rank, recv_tokens, recv_topk_weights + ) + return _EpPrepareAndDispatch.apply( buffer.handle_mem, + buffer.top_k, + buffer.alignment, recv_tokens, recv_topk_weights, + buffer.tokens_per_expert, topk_idx, tokens, topk_weights, ) - return recv_tokens, recv_topk_weights, tokens_per_expert def ep_combine(