diff --git a/.claude/skills/trtllm-moe-develop/SKILL.md b/.claude/skills/trtllm-moe-develop/SKILL.md index 0b644259e39d..d3a3a3add74e 100644 --- a/.claude/skills/trtllm-moe-develop/SKILL.md +++ b/.claude/skills/trtllm-moe-develop/SKILL.md @@ -673,7 +673,7 @@ Prefer the unified MoE tests: - Communication changes: `pytest tests/unittest/_torch/moe/test_moe_comm.py -k ''`. - Routing changes: `pytest tests/unittest/_torch/moe/test_moe_routing.py -k ''`. - Load balancer changes: `pytest tests/unittest/_torch/moe/test_moe_load_balancer.py -k ''`. -- Multi-GPU EP/all-to-all behavior: `pytest tests/unittest/_torch/moe/multi_gpu/test_moe_a2a.py -k ''`. +- Multi-GPU EP/all-to-all behavior: `pytest tests/unittest/_torch/moe/multi_gpu/test_nvlink_one_sided.py -k ''`. When GPU resources are required, use the TRT-LLM GPU allocation/test-runner skills first and record skipped tests with reasons. diff --git a/.claude/skills/trtllm-moe-develop/references/moe-canonical-code-examples.md b/.claude/skills/trtllm-moe-develop/references/moe-canonical-code-examples.md index 3751ee8c5d41..eff2cf0aedea 100644 --- a/.claude/skills/trtllm-moe-develop/references/moe-canonical-code-examples.md +++ b/.claude/skills/trtllm-moe-develop/references/moe-canonical-code-examples.md @@ -371,7 +371,7 @@ Use these examples when wrapper forward policy grows complicated: should move into the scheduler. - `tests/unittest/_torch/moe/test_moe_module.py` - Module-level multi-GPU, chunking, routing, and EPLB cases. -- `tests/unittest/_torch/moe/multi_gpu/test_moe_a2a.py` +- `tests/unittest/_torch/moe/multi_gpu/test_nvlink_one_sided.py` - Multi-GPU all-to-all behavior when relevant. Good uses: diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index b7a614a6e95a..b0a1fb14b56c 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -273,7 +273,6 @@ common-files: &common_files | tensorrt_llm/_torch/modules/triton_linear.py | tensorrt_llm/_torch/moe/expert_statistic.py | tensorrt_llm/_torch/moe/fused_moe/__init__.py | - tensorrt_llm/_torch/moe/fused_moe/communication/moe_alltoall.py | tensorrt_llm/_torch/moe/fused_moe/create_moe.py | tensorrt_llm/_torch/moe/fused_moe/deep_ep_utils.py | tensorrt_llm/_torch/moe/fused_moe/fused_moe_cute_dsl.py | @@ -582,7 +581,6 @@ common-files: &common_files | tests/unittest/_torch/modeling/test_modeling_vila.py | tests/unittest/_torch/modules/test_group_rmn_norm.py | tests/unittest/_torch/modules/test_triton_linear.py | - tests/unittest/_torch/moe/multi_gpu/test_moe_a2a.py | tests/unittest/_torch/moe/test_fused_moe.py | tests/unittest/_torch/moe/test_moe_host_sharer.py | tests/unittest/_torch/moe/test_moe_load_balancer.py | @@ -1029,7 +1027,6 @@ legacy-files: &legacy_files | tensorrt_llm/_torch/modules/triton_linear.py | tensorrt_llm/_torch/moe/expert_statistic.py | tensorrt_llm/_torch/moe/fused_moe/__init__.py | - tensorrt_llm/_torch/moe/fused_moe/communication/moe_alltoall.py | tensorrt_llm/_torch/moe/fused_moe/create_moe.py | tensorrt_llm/_torch/moe/fused_moe/deep_ep_utils.py | tensorrt_llm/_torch/moe/fused_moe/fused_moe_cute_dsl.py | @@ -1338,7 +1335,6 @@ legacy-files: &legacy_files | tests/unittest/_torch/modeling/test_modeling_vila.py | tests/unittest/_torch/modules/test_group_rmn_norm.py | tests/unittest/_torch/modules/test_triton_linear.py | - tests/unittest/_torch/moe/multi_gpu/test_moe_a2a.py | tests/unittest/_torch/moe/test_fused_moe.py | tests/unittest/_torch/moe/test_moe_host_sharer.py | tests/unittest/_torch/moe/test_moe_load_balancer.py | diff --git a/cpp/tensorrt_llm/common/envUtils.cpp b/cpp/tensorrt_llm/common/envUtils.cpp index aee8dfd2b99a..097225610b76 100644 --- a/cpp/tensorrt_llm/common/envUtils.cpp +++ b/cpp/tensorrt_llm/common/envUtils.cpp @@ -542,50 +542,6 @@ bool getEnvDisableChunkedAttentionInGenPhase() return getBoolEnv("TRTLLM_DISABLE_CHUNKED_ATTENTION_IN_GEN_PHASE"); } -static int sanitizeBlockSize(std::optional const& val) -{ - // Default 256 when not set or invalid - int block = val.value_or(256); - // Clamp to sane CUDA bounds and warp multiples - if (block <= 0) - block = 256; - if (block > 1024) - block = 1024; - // Round to nearest multiple of 32 (warp size) - block = (block + 31) / 32 * 32; - if (block == 0) - block = 256; - return block; -} - -// Read an integer env var and sanitize it as a CUDA block size. Treats malformed -// values (e.g. non-numeric strings that would throw inside std::stoi) as unset and -// falls back to the default, so this debug knob never becomes a hard failure. -static int getSanitizedBlockSizeFromEnv(char const* name) -{ - try - { - return sanitizeBlockSize(getIntEnv(name)); - } - catch (std::exception const&) - { - TLLM_LOG_WARNING("Invalid value for %s. Falling back to default block size.", name); - return sanitizeBlockSize(std::nullopt); - } -} - -int getEnvMoeA2ADispatchBlockSize() -{ - static int const kBlock = getSanitizedBlockSizeFromEnv("TLLM_MOE_A2A_DISPATCH_BLOCK_SIZE"); - return kBlock; -} - -int getEnvMoeA2ACombineBlockSize() -{ - static int const kBlock = getSanitizedBlockSizeFromEnv("TLLM_MOE_A2A_COMBINE_BLOCK_SIZE"); - return kBlock; -} - bool getEnvEplbForceGdrcopy() { return getBoolEnv("TRTLLM_EPLB_FORCE_GDRCOPY"); diff --git a/cpp/tensorrt_llm/common/envUtils.h b/cpp/tensorrt_llm/common/envUtils.h index a81e1b362f82..0c4503629f09 100644 --- a/cpp/tensorrt_llm/common/envUtils.h +++ b/cpp/tensorrt_llm/common/envUtils.h @@ -170,12 +170,6 @@ bool getEnvDisaggBenchmarkGenOnly(); // Whether to disable the chunked-attention in the generation phase. bool getEnvDisableChunkedAttentionInGenPhase(); -// TODO: For DEV purpose temporarily. -// Block size (threads per block) for MoE A2A Dispatch kernels (default 256 if unset or invalid) -int getEnvMoeA2ADispatchBlockSize(); -// Block size (threads per block) for MoE A2A Combine kernels (default 256 if unset or invalid) -int getEnvMoeA2ACombineBlockSize(); - bool getEnvKVCacheTransferAllBlocksForWindow(); bool getEnvEplbForceGdrcopy(); diff --git a/cpp/tensorrt_llm/kernels/moe/communication/moeAlltoAllKernels.cu b/cpp/tensorrt_llm/kernels/moe/communication/moeAlltoAllKernels.cu index 3768d6aa1e43..2438f0bafeae 100644 --- a/cpp/tensorrt_llm/kernels/moe/communication/moeAlltoAllKernels.cu +++ b/cpp/tensorrt_llm/kernels/moe/communication/moeAlltoAllKernels.cu @@ -23,7 +23,6 @@ #include "tensorrt_llm/kernels/moe/communication/moeAlltoAllCftSupport.h" #include "tensorrt_llm/kernels/moe/communication/moeAlltoAllKernels.h" #include "tensorrt_llm/kernels/quantization.cuh" -#include #include #include #include @@ -57,8 +56,6 @@ #define DISABLE_TIMEOUT 0 #endif -#define TLLM_MOE_A2A_COMPILE_CFT_DISPATCH TLLM_MOE_A2A_COMPILE_SM100 - #if TLLM_MOE_A2A_COMPILE_SM90 #include #include @@ -71,54 +68,6 @@ namespace kernels::moe_comm using tensorrt_llm::common::launchWithPdlWhenEnabled; -// Resolve the completion-flag wait budget; see the header. Seconds are converted at -// an assumed 2 GHz SM clock, so they are nominal rather than wall-clock. -int64_t moeA2AGetTimeoutCycles(bool is_warmup) -{ - static constexpr int64_t kAssumedClockHz = 2000ll * 1000ll * 1000ll; - static constexpr int64_t kDefaultTimeoutSec = 300; - // Warmup contains one-time per-rank costs (JIT compilation, autotuning, module - // loading) that can run for minutes and are not synchronized against this - // collective, so it needs a larger budget than steady state. - static constexpr int64_t kDefaultWarmupTimeoutSec = 1800; - - // Reject trailing garbage, out-of-range values and anything that would overflow - // the cycle multiplication. - auto const readEnv = [](char const* name, int64_t fallback) -> int64_t - { - static constexpr int64_t kMaxSec = 24 * 60 * 60; // 1 day; * 2e9 stays well inside int64 - char const* v = std::getenv(name); - if (v == nullptr || *v == '\0') - { - return fallback; - } - errno = 0; - char* end = nullptr; - int64_t parsed = std::strtoll(v, &end, 10); - bool const trailingGarbage = (end == v) || (*end != '\0'); - if (trailingGarbage || errno == ERANGE || parsed <= 0 || parsed > kMaxSec) - { - TLLM_LOG_WARNING("Ignoring invalid %s=\"%s\" (expected 1..%ld seconds); using %ld s", name, v, - static_cast(kMaxSec), static_cast(fallback)); - return fallback; - } - return parsed; - }; - - static int64_t const sSteadySec = readEnv("TRTLLM_MOE_A2A_TIMEOUT_SEC", kDefaultTimeoutSec); - static int64_t const sWarmupSec = readEnv("TRTLLM_MOE_A2A_WARMUP_TIMEOUT_SEC", kDefaultWarmupTimeoutSec); - static bool const sLogged = []() - { - TLLM_LOG_INFO( - "MoE all-to-all completion-flag budget: steady=%ld s, warmup=%ld s (nominal, at an " - "assumed 2 GHz clock64 rate)", - static_cast(sSteadySec), static_cast(sWarmupSec)); - return true; - }(); - (void) sLogged; - return (is_warmup ? sWarmupSec : sSteadySec) * kAssumedClockHz; -} - #define ENABLE_DEBUG_PRINT 0 #define DISABLE_SYNC_FOR_PROFILING 0 @@ -231,16 +180,10 @@ int64_t moeA2AGetTimeoutCycles(bool is_warmup) } \ } -#ifndef TLLM_MOE_A2A_TIMEOUT_SECONDS -#define TLLM_MOE_A2A_TIMEOUT_SECONDS 300 -#endif #if DISABLE_TIMEOUT #define check_timeout(s, budget) false #else -// `budget` is in clock64() cycles, resolved on the host by moeA2AGetTimeoutCycles(). -// Multi-rank warmup can enter these kernels with large rank skew while CuTeDSL -// kernels are still being JIT/autotuned on peer ranks; host budgets (incl. warmup) -// cover that skew via moeA2AGetTimeoutCycles(). +// The host supplies the wait budget in clock64() cycles as a launch argument. #define check_timeout(s, budget) ((clock64() - (s)) > (budget)) #endif @@ -284,34 +227,100 @@ __device__ __forceinline__ uint32_t round_parity(uint32_t flag_val) return ((flag_val - 1U) >> 1U) & 1U; } -template +// Round flags use system-scope relaxed accesses. Required payload visibility +// fences stay at the call sites; fence and CFT use different memory proxies. +__device__ __forceinline__ void publish_round_flag(uint32_t* address, uint32_t value) +{ + asm volatile("st.relaxed.sys.u32 [%0], %1;" ::"l"(address), "r"(value) : "memory"); +} + +__device__ __forceinline__ bool wait_round_flag( + uint32_t const* address, uint32_t expected, int64_t timeout_cycles, int rank_id, int peer_rank, char const* phase) +{ + auto const start = clock64(); + uint32_t observed; + do + { + asm volatile("ld.relaxed.sys.u32 %0, [%1];" : "=r"(observed) : "l"(address) : "memory"); +#if ENABLE_DEBUG_PRINT + printf("%s: rank %d waiting for rank %d flag=%u expected=%u address=%p\n", phase, rank_id, peer_rank, observed, + expected, address); +#endif + if (observed == expected) + { + return true; + } + } while (!check_timeout(start, timeout_cycles)); + + printf("%s: rank %d timed out waiting for rank %d flag=%u expected=%u\n", phase, rank_id, peer_rank, observed, + expected); + asm volatile("trap;" ::: "memory"); + return false; +} + +// Mask the first NUM_LANES lanes of a warp. +template +__host__ __device__ constexpr uint32_t make_warp_lane_mask() +{ + static_assert(NUM_LANES > 0 && NUM_LANES <= 32, "warp lane count must be in [1, 32]"); + return (NUM_LANES == 32) ? ~0U : ((1U << NUM_LANES) - 1U); +} + +template __device__ __forceinline__ void route_dispatch_token(int32_t const* token_selected_experts, DispatchKernelPointers const& ptrs, int local_token_idx, int ep_size, int num_experts, int* topk_target_ranks, - int* topk_send_indices) + int* topk_target_indices, uint32_t lane_mask) { static_assert(TOP_K <= 32, "warp-parallel routing requires TOP_K <= warpSize"); - uint32_t const lane_mask = (TOP_K == 32) ? ~0U : ((1U << TOP_K) - 1U); int const k = threadIdx.x; + if constexpr (USE_RANK_COMPACT_ROUTING) + { + // EP < TOP_K, so the routing lanes initialize every destination slot. + topk_target_ranks[k] = -1; + topk_target_indices[k] = -1; + __syncwarp(lane_mask); + } + int const ep_base = num_experts / ep_size; int const ep_remainder = num_experts - ep_base * ep_size; int const expert_id = token_selected_experts[local_token_idx * TOP_K + k]; - int const target_rank = compute_target_rank_id(expert_id, ep_base, ep_remainder); + // Invalid experts have no destination, but their lanes still participate in warp collectives. + bool const valid_expert = expert_id >= 0 && expert_id < num_experts; + int const target_rank = valid_expert ? compute_target_rank_id(expert_id, ep_base, ep_remainder) : -1; + // Experts on the same rank share one token transfer; only their first lane allocates a send slot. uint32_t const same_target = __match_any_sync(lane_mask, target_rank); - bool keep = (__ffs(same_target) - 1) == k; + bool should_send = valid_expert && ((__ffs(same_target) - 1) == k); if constexpr (ENABLE_RANK_MASK) { - keep = keep && is_rank_active(ptrs.active_rank_mask, target_rank); + should_send = should_send && is_rank_active(ptrs.active_rank_mask, target_rank); } - int const target_rank_to_store = keep ? target_rank : -1; - int const send_index_to_store = keep ? atomicAdd(&ptrs.send_counters[target_rank], 1) : -1; + int const target_rank_to_store = should_send ? target_rank : -1; + int const target_index_to_store = should_send ? atomicAdd(&ptrs.send_counters[target_rank], 1) : -1; + // Keep persistent routing metadata indexed by top-k slot for combine. ptrs.topk_target_ranks[local_token_idx * TOP_K + k] = target_rank_to_store; - ptrs.topk_send_indices[local_token_idx * TOP_K + k] = send_index_to_store; - topk_target_ranks[k] = target_rank_to_store; - topk_send_indices[k] = send_index_to_store; + ptrs.topk_target_indices[local_token_idx * TOP_K + k] = target_index_to_store; + if constexpr (USE_RANK_COMPACT_ROUTING) + { + // Pack unique target ranks into consecutive shared-memory slots instead of leaving + // holes at duplicate expert lanes. Dispatch then scans a rank-bounded list, not TOP_K slots. + uint32_t const sending_lanes = __ballot_sync(lane_mask, should_send); + if (should_send) + { + // Count sending lanes before this lane to obtain its packed slot. + int const compact_index = __popc(sending_lanes & ((1U << k) - 1U)); + topk_target_ranks[compact_index] = target_rank; + topk_target_indices[compact_index] = target_index_to_store; + } + } + else + { + topk_target_ranks[k] = target_rank_to_store; + topk_target_indices[k] = target_index_to_store; + } } // ============================================================================ @@ -360,80 +369,53 @@ __device__ void vectorized_copy(void* dst, void const* src, int size) } } -// Vectorized dispatch: load one vec from source and write to up to TOP_K destinations -template -__device__ void vectorized_dispatch_impl(uint8_t const* src_ptr, int bytes_per_token, int rank_id, - int max_tokens_per_rank, int payload_idx, DispatchKernelPointers const& ptrs, int const* topk_target_ranks, - int const* topk_send_indices) +// Fan each source vector out through the CTA's shared destination table. +template +__device__ void vectorized_dispatch_impl( + uint8_t const* src_ptr, int bytes_per_token, int ep_size, uint8_t* const* destinations) { using flashinfer::vec_t; - // Precompute destination base pointers per k - uint8_t* dst_base_k[TOP_K]; -#pragma unroll - for (int k = 0; k < TOP_K; ++k) + int const num_destinations = USE_RANK_COMPACT_ROUTING ? ep_size : TOP_K; + for (int offset = threadIdx.x * VEC_SIZE; offset < bytes_per_token; offset += blockDim.x * VEC_SIZE) { - int dst_idx_k = topk_send_indices[k]; - int target_rank_k = topk_target_ranks[k]; - if (dst_idx_k < 0) + vec_t value; + value.load(src_ptr + offset); +#pragma unroll(USE_RANK_COMPACT_ROUTING ? 2 : TOP_K) + for (int k = 0; k < num_destinations; ++k) { - dst_base_k[k] = nullptr; - continue; - } - uint8_t* dst_data = static_cast(ptrs.recv_buffers[target_rank_k][payload_idx]); - size_t base_source_rank - = static_cast(rank_id) * static_cast(max_tokens_per_rank) + static_cast(dst_idx_k); - size_t base_token = base_source_rank * static_cast(bytes_per_token); - dst_base_k[k] = dst_data + base_token; - } - - int const stride = blockDim.x * VEC_SIZE; - for (int offset = threadIdx.x * VEC_SIZE; offset < bytes_per_token; offset += stride) - { - vec_t v; - v.load(src_ptr + offset); - -#pragma unroll - for (int k = 0; k < TOP_K; ++k) - { - uint8_t* dst_base = dst_base_k[k]; - if (dst_base == nullptr) + uint8_t* const destination = destinations[k]; + if (destination != nullptr) { - continue; + value.store(destination + offset); } - v.store(dst_base + offset); } } } -template -__device__ void vectorized_dispatch(uint8_t const* src_ptr, int bytes_per_token, int rank_id, int max_tokens_per_rank, - int payload_idx, DispatchKernelPointers const& ptrs, int const* topk_target_ranks, int const* topk_send_indices) +template +__device__ void vectorized_dispatch( + uint8_t const* src_ptr, int bytes_per_token, int ep_size, uint8_t* const* destinations) { if (bytes_per_token % 16 == 0) { - vectorized_dispatch_impl<16, TOP_K>(src_ptr, bytes_per_token, rank_id, max_tokens_per_rank, payload_idx, ptrs, - topk_target_ranks, topk_send_indices); + vectorized_dispatch_impl<16, TOP_K, USE_RANK_COMPACT_ROUTING>(src_ptr, bytes_per_token, ep_size, destinations); } else if (bytes_per_token % 8 == 0) { - vectorized_dispatch_impl<8, TOP_K>(src_ptr, bytes_per_token, rank_id, max_tokens_per_rank, payload_idx, ptrs, - topk_target_ranks, topk_send_indices); + vectorized_dispatch_impl<8, TOP_K, USE_RANK_COMPACT_ROUTING>(src_ptr, bytes_per_token, ep_size, destinations); } else if (bytes_per_token % 4 == 0) { - vectorized_dispatch_impl<4, TOP_K>(src_ptr, bytes_per_token, rank_id, max_tokens_per_rank, payload_idx, ptrs, - topk_target_ranks, topk_send_indices); + vectorized_dispatch_impl<4, TOP_K, USE_RANK_COMPACT_ROUTING>(src_ptr, bytes_per_token, ep_size, destinations); } else if (bytes_per_token % 2 == 0) { - vectorized_dispatch_impl<2, TOP_K>(src_ptr, bytes_per_token, rank_id, max_tokens_per_rank, payload_idx, ptrs, - topk_target_ranks, topk_send_indices); + vectorized_dispatch_impl<2, TOP_K, USE_RANK_COMPACT_ROUTING>(src_ptr, bytes_per_token, ep_size, destinations); } else { - vectorized_dispatch_impl<1, TOP_K>(src_ptr, bytes_per_token, rank_id, max_tokens_per_rank, payload_idx, ptrs, - topk_target_ranks, topk_send_indices); + vectorized_dispatch_impl<1, TOP_K, USE_RANK_COMPACT_ROUTING>(src_ptr, bytes_per_token, ep_size, destinations); } } @@ -460,16 +442,13 @@ __global__ void moeA2APrepareDispatchKernel( uint32_t const next_parity = current_parity ^ 1U; recv_counters[next_parity * ep_size + idx] = -1; } - // NOTE: LE-backed counters use cumulative baselines and are deliberately not zeroed - // here, so that the kernel never issues SM stores to LE-backed memory (historically - // broke fabric.try_put.counted with PDL). } // ============================================================================ // Dispatch Kernels // ============================================================================ -template +template __global__ void moeA2ADispatchKernel(int32_t const* token_selected_experts, // [local_num_tokens, TOP_K] const DispatchKernelPointers ptrs, // Struct containing all kernel pointers int num_payloads, // Number of payloads @@ -496,41 +475,52 @@ __global__ void moeA2ADispatchKernel(int32_t const* token_selected_experts, // [ if (local_token_idx >= local_num_tokens) return; - // One block per token: a single shared-memory tile is reused by the entire CTA. + // Compact tiles keep valid destinations in top-k order and pad with -1 send indices. + // Global routing metadata always retains all top-k slots. extern __shared__ int smem[]; int* smem_topk_target_ranks = smem; - int* smem_topk_send_indices = smem + TOP_K; + int* smem_topk_target_indices = smem + TOP_K; #if TLLM_MOE_A2A_COMPILE_SM90 cudaGridDependencySynchronize(); #endif + static_assert(TOP_K > 0 && TOP_K <= 32, "routing and destination setup require one warp"); + constexpr uint32_t kRoutingLaneMask = make_warp_lane_mask(); + __shared__ uint8_t* destinations[kMaxPayloads][TOP_K]; if (thread_idx < TOP_K) { - route_dispatch_token(token_selected_experts, ptrs, local_token_idx, ep_size, - num_experts, smem_topk_target_ranks, smem_topk_send_indices); + route_dispatch_token(token_selected_experts, ptrs, + local_token_idx, ep_size, num_experts, smem_topk_target_ranks, smem_topk_target_indices, + kRoutingLaneMask); + // Compact routing slots may be written by a different lane. + __syncwarp(kRoutingLaneMask); + + // Resolve each payload's destinations once per CTA; unused routes remain null. + int const target_index = smem_topk_target_indices[thread_idx]; + int const peer = smem_topk_target_ranks[thread_idx]; + for (int p = 0; p < num_payloads; ++p) + { + uint8_t* destination = nullptr; + if (target_index >= 0) + { + size_t const token = static_cast(rank_id) * max_tokens_per_rank + target_index; + destination + = static_cast(ptrs.recv_buffers[peer][p]) + token * ptrs.payload_bytes_per_token[p]; + } + destinations[p][thread_idx] = destination; + } } - // Sync before dispatching data + // Publish the destination table to all payload-copying warps. __syncthreads(); - // Read staged routing once into registers per thread - int topk_target_ranks[TOP_K]; - int topk_send_indices[TOP_K]; -#pragma unroll - for (int k = 0; k < TOP_K; ++k) - { - topk_target_ranks[k] = smem_topk_target_ranks[k]; - topk_send_indices[k] = smem_topk_send_indices[k]; - } - // Perform a single source load and TOP_K fanout per payload for (int payload_idx = 0; payload_idx < num_payloads; payload_idx++) { uint8_t const* src_data = static_cast(ptrs.src_data_ptrs[payload_idx]); int bytes_per_token = ptrs.payload_bytes_per_token[payload_idx]; uint8_t const* src_ptr = src_data + local_token_idx * bytes_per_token; - - vectorized_dispatch(src_ptr, bytes_per_token, rank_id, max_tokens_per_rank, payload_idx, ptrs, - topk_target_ranks, topk_send_indices); + vectorized_dispatch( + src_ptr, bytes_per_token, ep_size, destinations[payload_idx]); } __syncthreads(); @@ -618,7 +608,7 @@ __global__ void moeA2ADispatchKernel(int32_t const* token_selected_experts, // [ continue; } uint32_t* flag_addr = &ptrs.completion_flags[target_rank][rank_id]; - asm volatile("st.relaxed.sys.u32 [%0], %1;" ::"l"(flag_addr), "r"(expected_value)); + publish_round_flag(flag_addr, expected_value); #if ENABLE_DEBUG_PRINT printf("dispatch: +++Rank %d setting completion flag to %d for rank %d\n", rank_id, expected_value, @@ -639,28 +629,9 @@ __global__ void moeA2ADispatchKernel(int32_t const* token_selected_experts, // [ continue; } } - bool flag_set = false; - auto s = clock64(); - do - { - uint32_t* flag_ptr = &ptrs.completion_flags[rank_id][peer_rank]; - uint32_t flag_value; - // Acquire load to ensure visibility of peer's release-store - asm volatile("ld.relaxed.sys.u32 %0, [%1];" : "=r"(flag_value) : "l"(flag_ptr)); -#if ENABLE_DEBUG_PRINT - printf( - "combine: ---Rank %d received completion flag from rank %d, flag_value: %d, expected_value: " - "%d, address: %p\n", - rank_id, peer_rank, flag_value, expected_value, flag_ptr); -#endif - flag_set = flag_value == expected_value; - } while (!flag_set && !check_timeout(s, ptrs.timeout_cycles)); - - if (__builtin_expect(!flag_set, 0)) + if (!wait_round_flag(&ptrs.completion_flags[rank_id][peer_rank], expected_value, ptrs.timeout_cycles, + rank_id, peer_rank, "dispatch")) { - printf("dispatch: ---Rank %d timed out waiting for completion flag from rank %d\n", rank_id, - peer_rank); - asm volatile("trap;"); return; } } @@ -679,7 +650,7 @@ __global__ void moeA2ADispatchKernel(int32_t const* token_selected_experts, // [ // 3. Last block sends recv_counters through symmetric memory using the current round parity // 4. Poll metadata + data counters from all peers (no fence.sys needed) // ============================================================================ -#if TLLM_MOE_A2A_COMPILE_CFT_DISPATCH +#if TLLM_MOE_A2A_COMPILE_SM100 __device__ __forceinline__ void cft_barrier_wait_parity(__mbarrier_t* barrier, int parity) { while (!::cuda::ptx::mbarrier_try_wait_parity(::cuda::ptx::sem_relaxed, ::cuda::ptx::scope_cta, @@ -791,11 +762,15 @@ __device__ __forceinline__ void cft_publish_recv_counters(DispatchKernelPointers } // Elect the CTA that finishes routing last and have it publish the counters. -// Should be run by only 1 warp. +// Called by the CTA; only warp 0 participates. template __device__ __forceinline__ void cft_elect_and_publish(DispatchKernelPointers const& ptrs, int rank_id, int ep_size, uint32_t parity, int eplb_stats_num_experts, int local_num_tokens, int& is_last_token_cta) { + if (threadIdx.x >= warpSize) + { + return; + } int const lane_id = threadIdx.x % warpSize; bool is_last_token = false; if (lane_id == 0) @@ -823,10 +798,10 @@ __device__ __forceinline__ void cft_elect_and_publish(DispatchKernelPointers con } } -template -__global__ void moeA2ADispatchCountedWriteKernel(int32_t const* token_selected_experts, - DispatchKernelPointers const ptrs, int num_payloads, int max_tokens_per_rank, int local_num_tokens, int rank_id, - int ep_size, int num_experts, int eplb_stats_num_experts) +template +__global__ void moeA2ADispatchKernel_Cft(int32_t const* token_selected_experts, DispatchKernelPointers const ptrs, + int num_payloads, int max_tokens_per_rank, int local_num_tokens, int rank_id, int ep_size, int num_experts, + int eplb_stats_num_experts) { int local_token_idx = blockIdx.x; uint32_t parity = 0; @@ -842,12 +817,8 @@ __global__ void moeA2ADispatchCountedWriteKernel(int32_t const* token_selected_e return; cudaGridDependencySynchronize(); parity = round_parity(*ptrs.flag_val); - __syncthreads(); - if (threadIdx.x < warpSize) - { - cft_elect_and_publish( - ptrs, rank_id, ep_size, parity, eplb_stats_num_experts, local_num_tokens, is_last_token_cta); - } + cft_elect_and_publish( + ptrs, rank_id, ep_size, parity, eplb_stats_num_experts, local_num_tokens, is_last_token_cta); } else { @@ -856,7 +827,7 @@ __global__ void moeA2ADispatchCountedWriteKernel(int32_t const* token_selected_e extern __shared__ int smem[]; int* smem_topk_target_ranks = smem; - int* smem_topk_send_indices = smem + TOP_K; + int* smem_topk_target_indices = smem + TOP_K; // CFT smem layout (disjoint regions, kept stable across phases): // [0 .. kRoutingBytes) routing indices (above) @@ -866,6 +837,7 @@ __global__ void moeA2ADispatchCountedWriteKernel(int32_t const* token_selected_e // runs on the TMA engine in parallel with routing + self-send. constexpr int kRoutingBytes = 2 * TOP_K * static_cast(sizeof(int)); uint8_t* smem_bytes = reinterpret_cast(smem); + // Wait for TMA staging, then reuse this barrier as the fabric put report target. __mbarrier_t* tma_bar = reinterpret_cast<__mbarrier_t*>(smem_bytes + kRoutingBytes); uint8_t* smem_staging = smem_bytes + kRoutingBytes + kCftMbarrierSlotBytes; @@ -897,45 +869,33 @@ __global__ void moeA2ADispatchCountedWriteKernel(int32_t const* token_selected_e } // ---- Routing: map tokens to target ranks ---- + constexpr uint32_t kRoutingLaneMask = make_warp_lane_mask(); if (threadIdx.x < TOP_K) { - route_dispatch_token(token_selected_experts, ptrs, local_token_idx, ep_size, - num_experts, smem_topk_target_ranks, smem_topk_send_indices); + route_dispatch_token(token_selected_experts, ptrs, + local_token_idx, ep_size, num_experts, smem_topk_target_ranks, smem_topk_target_indices, + kRoutingLaneMask); } __syncthreads(); // Routing is done, so send_counters is final. Publish it here rather than after the data // issue, so peers see the counts as early as possible. - if (threadIdx.x < warpSize) - { - cft_elect_and_publish( - ptrs, rank_id, ep_size, parity, eplb_stats_num_experts, local_num_tokens, is_last_token_cta); - } - - int topk_target_ranks[TOP_K]; - int topk_send_indices[TOP_K]; -#pragma unroll - for (int k = 0; k < TOP_K; ++k) - { - topk_target_ranks[k] = smem_topk_target_ranks[k]; - topk_send_indices[k] = smem_topk_send_indices[k]; - } + cft_elect_and_publish( + ptrs, rank_id, ep_size, parity, eplb_stats_num_experts, local_num_tokens, is_last_token_cta); // ---- Data dispatch: self via TMA s2g, remote via fabric.try_put.counted ---- - // Both are issued by thread 0 fire-and-forget. They run in parallel on - // different HW units (TMA engine for s2g, fabric engine for puts), with a - // single combined wait phase at the end. - // - // Self-send needs smem_staging populated, so it must come AFTER the TMA g2s - // wait. Remote-send also reads smem_staging — both share the same source. + // Separate issuing warps overlap self and remote transfers. Both consume + // smem_staging only after the TMA g2s wait below. bool has_remote = false; bool has_self = false; -#pragma unroll - for (int k = 0; k < TOP_K; ++k) +#pragma unroll(USE_RANK_COMPACT_ROUTING ? 2 : TOP_K) + for (int k = 0; k < (USE_RANK_COMPACT_ROUTING ? ep_size : TOP_K); ++k) { - if (topk_send_indices[k] < 0) + int const dst_idx = smem_topk_target_indices[k]; + int const target_rank = smem_topk_target_ranks[k]; + if (dst_idx < 0) continue; - if (topk_target_ranks[k] == rank_id) + if (target_rank == rank_id) has_self = true; else has_remote = true; @@ -959,11 +919,11 @@ __global__ void moeA2ADispatchCountedWriteKernel(int32_t const* token_selected_e for (int payload_idx = 0; payload_idx < num_payloads; payload_idx++) { int bytes_per_token = ptrs.payload_bytes_per_token[payload_idx]; -#pragma unroll - for (int k = 0; k < TOP_K; ++k) +#pragma unroll(USE_RANK_COMPACT_ROUTING ? 2 : TOP_K) + for (int k = 0; k < (USE_RANK_COMPACT_ROUTING ? ep_size : TOP_K); ++k) { - int dst_idx_k = topk_send_indices[k]; - int target_rank_k = topk_target_ranks[k]; + int const dst_idx_k = smem_topk_target_indices[k]; + int const target_rank_k = smem_topk_target_ranks[k]; if (dst_idx_k < 0 || target_rank_k != rank_id) continue; uint8_t* dst = static_cast(ptrs.recv_buffers[rank_id][payload_idx]) @@ -987,11 +947,11 @@ __global__ void moeA2ADispatchCountedWriteKernel(int32_t const* token_selected_e for (int payload_idx = 0; payload_idx < num_payloads; payload_idx++) { int bytes_per_token = ptrs.payload_bytes_per_token[payload_idx]; -#pragma unroll - for (int k = 0; k < TOP_K; ++k) +#pragma unroll(USE_RANK_COMPACT_ROUTING ? 2 : TOP_K) + for (int k = 0; k < (USE_RANK_COMPACT_ROUTING ? ep_size : TOP_K); ++k) { - int dst_idx_k = topk_send_indices[k]; - int target_rank_k = topk_target_ranks[k]; + int const dst_idx_k = smem_topk_target_indices[k]; + int const target_rank_k = smem_topk_target_ranks[k]; if (dst_idx_k < 0 || target_rank_k == rank_id) continue; uint64_t base_le_offset = ptrs.le_payload_offsets[payload_idx] @@ -1012,7 +972,6 @@ __global__ void moeA2ADispatchCountedWriteKernel(int32_t const* token_selected_e if (threadIdx.x == 0 && has_remote) cft_fabric_wait_reads(); - __syncthreads(); } cudaTriggerProgrammaticLaunchCompletion(); @@ -1130,11 +1089,11 @@ __global__ void moeA2ADispatchCountedWriteKernel(int32_t const* token_selected_e } #endif // !DISABLE_SYNC_FOR_PROFILING } -#else // TLLM_MOE_A2A_COMPILE_CFT_DISPATCH +#else // TLLM_MOE_A2A_COMPILE_SM100 template -__global__ void moeA2ADispatchCountedWriteKernel(int32_t const* token_selected_experts, - DispatchKernelPointers const ptrs, int num_payloads, int max_tokens_per_rank, int local_num_tokens, int rank_id, - int ep_size, int num_experts, int eplb_stats_num_experts) +__global__ void moeA2ADispatchKernel_Cft(int32_t const* token_selected_experts, DispatchKernelPointers const ptrs, + int num_payloads, int max_tokens_per_rank, int local_num_tokens, int rank_id, int ep_size, int num_experts, + int eplb_stats_num_experts) { (void) token_selected_experts; (void) ptrs; @@ -1147,13 +1106,10 @@ __global__ void moeA2ADispatchCountedWriteKernel(int32_t const* token_selected_e (void) eplb_stats_num_experts; asm volatile("trap;" ::: "memory"); } -#endif // TLLM_MOE_A2A_COMPILE_CFT_DISPATCH +#endif // TLLM_MOE_A2A_COMPILE_SM100 void moe_a2a_prepare_dispatch_launch(MoeA2ADispatchParams const& params) { - // NOTE: LE counters are NOT zeroed between iterations. They grow monotonically. - // Cumulative baselines in regular device memory track the expected value. - launchWithPdlWhenEnabled("moeA2APrepareDispatchKernel", moeA2APrepareDispatchKernel, 1, params.ep_size, 0, params.stream, params.send_counters, params.recv_counters[params.ep_rank], params.local_token_counter, params.ep_size, params.flag_val); @@ -1165,6 +1121,8 @@ void moe_a2a_prepare_dispatch_launch(MoeA2ADispatchParams const& params) void moe_a2a_dispatch_launch(MoeA2ADispatchParams const& params) { + constexpr int kBlockSize = 256; + // Validate parameters TLLM_CHECK(params.top_k > 0 && params.top_k <= kMaxTopK); TLLM_CHECK(params.ep_size > 0 && params.ep_size <= kMaxRanks); @@ -1214,7 +1172,7 @@ void moe_a2a_dispatch_launch(MoeA2ADispatchParams const& params) kernel_ptrs.send_counters = params.send_counters; kernel_ptrs.local_token_counter = params.local_token_counter; kernel_ptrs.topk_target_ranks = params.topk_target_ranks; - kernel_ptrs.topk_send_indices = params.topk_send_indices; + kernel_ptrs.topk_target_indices = params.topk_target_indices; kernel_ptrs.eplb_local_stats = params.eplb_local_stats; // CFT handle-based counted writes fields @@ -1241,8 +1199,6 @@ void moe_a2a_dispatch_launch(MoeA2ADispatchParams const& params) kernel_ptrs.active_rank_mask[w] = params.active_rank_mask[w]; } - int const kBlockSize = tensorrt_llm::common::getEnvMoeA2ADispatchBlockSize(); - int grid_size = params.local_num_tokens; if (grid_size == 0) { @@ -1280,30 +1236,42 @@ void moe_a2a_dispatch_launch(MoeA2ADispatchParams const& params) shared_bytes, maxOptinBytes); } - SWITCH_BOOL(params.enable_rank_mask, ENABLE_RANK_MASK, - {SWITCH_BOOL(params.enable_eplb, EPLB_STATS, SWITCH_TOP_K(params.top_k, TOP_K, { - auto kernel_fn = moeA2ADispatchCountedWriteKernel; - if (shared_bytes > kDefaultDynamicSmemBytes) - { - TLLM_CUDA_CHECK( - cudaFuncSetAttribute(kernel_fn, cudaFuncAttributeMaxDynamicSharedMemorySize, shared_bytes)); - } - launchWithPdlWhenEnabled("moeA2ADispatchCountedWriteKernel", kernel_fn, grid_size, kBlockSize, - shared_bytes, params.stream, params.token_selected_experts, kernel_ptrs, params.num_payloads, - params.max_tokens_per_rank, params.local_num_tokens, params.ep_rank, params.ep_size, - params.num_experts, params.eplb_stats_num_experts); - }))}) + SWITCH_BOOL(params.enable_rank_mask, ENABLE_RANK_MASK, { + SWITCH_BOOL(params.enable_eplb, EPLB_STATS, { + SWITCH_TOP_K(params.top_k, TOP_K, { + SWITCH_BOOL(params.ep_size < TOP_K, ENABLE_RANK_COMPACT_ROUTING, { + auto kernel_fn = moeA2ADispatchKernel_Cft; + if (shared_bytes > kDefaultDynamicSmemBytes) + { + TLLM_CUDA_CHECK(cudaFuncSetAttribute( + kernel_fn, cudaFuncAttributeMaxDynamicSharedMemorySize, shared_bytes)); + } + launchWithPdlWhenEnabled("moeA2ADispatchKernel_Cft", kernel_fn, grid_size, kBlockSize, + shared_bytes, params.stream, params.token_selected_experts, kernel_ptrs, + params.num_payloads, params.max_tokens_per_rank, params.local_num_tokens, params.ep_rank, + params.ep_size, params.num_experts, params.eplb_stats_num_experts); + }); + }); + }); + }); } else { - SWITCH_BOOL(params.enable_rank_mask, ENABLE_RANK_MASK, - {SWITCH_BOOL(params.enable_eplb, EPLB_STATS, SWITCH_TOP_K(params.top_k, TOP_K, { - auto kernel_fn = moeA2ADispatchKernel; - launchWithPdlWhenEnabled("moeA2ADispatchKernel", kernel_fn, grid_size, kBlockSize, shared_bytes, - params.stream, params.token_selected_experts, kernel_ptrs, params.num_payloads, - params.max_tokens_per_rank, params.local_num_tokens, params.ep_rank, params.ep_size, - params.num_experts, params.eplb_stats_num_experts); - }))}) + SWITCH_BOOL(params.enable_rank_mask, ENABLE_RANK_MASK, { + SWITCH_BOOL(params.enable_eplb, EPLB_STATS, { + SWITCH_TOP_K(params.top_k, TOP_K, { + SWITCH_BOOL(params.ep_size < TOP_K, ENABLE_RANK_COMPACT_ROUTING, { + auto kernel_fn + = moeA2ADispatchKernel; + launchWithPdlWhenEnabled("moeA2ADispatchKernel", kernel_fn, grid_size, kBlockSize, shared_bytes, + params.stream, params.token_selected_experts, kernel_ptrs, params.num_payloads, + params.max_tokens_per_rank, params.local_num_tokens, params.ep_rank, params.ep_size, + params.num_experts, params.eplb_stats_num_experts); + }); + }); + }); + }); } } @@ -1311,298 +1279,175 @@ void moe_a2a_dispatch_launch(MoeA2ADispatchParams const& params) // Combine kernels // ============================================================================ -// Accumulate across all valid ranks into float32 registers, then store as OutputT. -// InputT is the wire element type in the receive buffer. -// -// Unified path: load VEC_SIZE bytes, reinterpret as InputT[elems_per_vec], accumulate as float32, -// store as OutputT. Works for same-type and FP8-to-payload-type accumulation. -// sizeof(InputT) must divide VEC_SIZE. -template -__device__ void vectorized_combine_impl(OutputT* dst_typed_base, int size_per_token, int stride_per_token, int rank_id, - int max_tokens_per_rank, CombineKernelPointers const& ptrs) +// Load source vectors in their wire dtype and accumulate in FP32. +// Only the first source_count entries are valid, one per contributing rank. +template +__device__ void vectorized_combine_impl( + OutputT* output, int size_per_token, uint8_t const* const* sources, int source_count) { using flashinfer::vec_t; - - // elems_per_vec is the number of InputT elements per VEC_SIZE-byte load. - constexpr int elems_per_vec = VEC_SIZE / static_cast(sizeof(InputT)); - - int const stride = blockDim.x * VEC_SIZE; - int const local_token_idx = blockIdx.x; - - // offset is a byte offset into the recv buffer, stepping by VEC_SIZE bytes. - for (int offset = threadIdx.x * VEC_SIZE; offset < size_per_token; offset += stride) + constexpr int kElements = VEC_SIZE / static_cast(sizeof(InputT)); + if (source_count <= 2) { - // Per-k vec_t accumulators, zero-initialised via fill(). - // Using vec_t enables cast_store() for the output, emitting a vectorized int4 write. - vec_t acc[TOP_K]; - - // Pass 1: issue all TOP_K loads back-to-back without any type conversion. - // Raw InputT bytes are loaded directly into acc[k]'s register storage, reinterpreted as - // vec_t (VEC_SIZE bytes, fitting in the low end of acc[k]'s - // sizeof(float)*elems_per_vec allocation). Separating load from cast lets the compiler - // schedule all VEC_SIZE-byte global loads consecutively, hiding memory latency across k. -#pragma unroll - for (int k = 0; k < TOP_K; ++k) - { - int target_rank = ptrs.topk_target_ranks[local_token_idx * TOP_K + k]; - int dst_idx = ptrs.topk_send_indices[local_token_idx * TOP_K + k]; - if (dst_idx < 0) - { - acc[k].fill(0.0f); - continue; - } - - // Every contribution uses the same compact receive-buffer layout. - uint8_t const* recv_buffer = static_cast(ptrs.recv_buffers[target_rank][0]); - size_t base_source_rank = static_cast(rank_id) * static_cast(max_tokens_per_rank) - + static_cast(dst_idx); - size_t base_token = base_source_rank * static_cast(stride_per_token); - - reinterpret_cast&>(acc[k]).load( - reinterpret_cast(recv_buffer + base_token + offset)); - } - - // Pass 2: in-place cast InputT to float, iterating j in descending order. - // float[j] occupies bytes [j*4, j*4+3]; InputT[j] occupies - // [j*sizeof(InputT), ...). For narrow inputs, high-j float writes land above all - // remaining InputT bytes, so descending order is write-after-read safe. -#pragma unroll - for (int k = 0; k < TOP_K; ++k) + for (int offset = threadIdx.x * VEC_SIZE; offset < size_per_token; offset += blockDim.x * VEC_SIZE) { - int target_rank = ptrs.topk_target_ranks[local_token_idx * TOP_K + k]; - int dst_idx = ptrs.topk_send_indices[local_token_idx * TOP_K + k]; - if (dst_idx < 0) + vec_t left; + vec_t right; + left.fill(InputT(0.0f)); + right.fill(InputT(0.0f)); + if (source_count > 0) { - continue; // acc[k] already holds 0.0f from fill() above + left.load(reinterpret_cast(sources[0] + offset)); } -#pragma unroll - for (int j = elems_per_vec - 1; j >= 0; --j) - acc[k][j] = static_cast(reinterpret_cast(&acc[k])[j]); - } - // Reduce acc[TOP_K] into acc[0] via unrolled tree-reduction. - // acc[k][j] uses vec_t::operator[] which returns float& — no indirection overhead. - if constexpr (TOP_K == 22) - { -#pragma unroll - for (int j = 0; j < elems_per_vec; ++j) - { - acc[0][j] += acc[1][j]; - acc[2][j] += acc[3][j]; - acc[4][j] += acc[5][j]; - acc[6][j] += acc[7][j]; - acc[8][j] += acc[9][j]; - acc[10][j] += acc[11][j]; - acc[12][j] += acc[13][j]; - acc[14][j] += acc[15][j]; - acc[16][j] += acc[17][j]; - acc[18][j] += acc[19][j]; - acc[20][j] += acc[21][j]; - } -#pragma unroll - for (int j = 0; j < elems_per_vec; ++j) + if (source_count > 1) { - acc[0][j] += acc[2][j]; - acc[4][j] += acc[6][j]; - acc[8][j] += acc[10][j]; - acc[12][j] += acc[14][j]; - acc[16][j] += acc[18][j]; + right.load(reinterpret_cast(sources[1] + offset)); } + vec_t value; #pragma unroll - for (int j = 0; j < elems_per_vec; ++j) + for (int j = 0; j < kElements; ++j) { - acc[0][j] += acc[4][j]; - acc[8][j] += acc[12][j]; - acc[16][j] += acc[20][j]; - } -#pragma unroll - for (int j = 0; j < elems_per_vec; ++j) - { - acc[0][j] += acc[8][j]; - acc[0][j] += acc[16][j]; + value[j] = static_cast(left[j]) + static_cast(right[j]); } + value.cast_store(output + offset / static_cast(sizeof(InputT))); } - else if constexpr (TOP_K == 16) - { -#pragma unroll - for (int j = 0; j < elems_per_vec; ++j) - { - acc[0][j] += acc[1][j]; - acc[2][j] += acc[3][j]; - acc[4][j] += acc[5][j]; - acc[6][j] += acc[7][j]; - acc[8][j] += acc[9][j]; - acc[10][j] += acc[11][j]; - acc[12][j] += acc[13][j]; - acc[14][j] += acc[15][j]; - } -#pragma unroll - for (int j = 0; j < elems_per_vec; ++j) - { - acc[0][j] += acc[2][j]; - acc[4][j] += acc[6][j]; - acc[8][j] += acc[10][j]; - acc[12][j] += acc[14][j]; - } -#pragma unroll - for (int j = 0; j < elems_per_vec; ++j) - { - acc[0][j] += acc[4][j]; - acc[8][j] += acc[12][j]; - } + return; + } + + // TOP_K <= 32: eight sources per lane require at most four cooperating lanes. + int const rank_bits = (source_count > 8) + (source_count > 16); + int const rank_lanes = 1 << rank_bits; + int const lane = threadIdx.x & 31; + int const rank_lane = lane & (rank_lanes - 1); + int const output_lane = lane >> rank_bits; + int const warp = threadIdx.x >> 5; + int const vectors_per_warp = warpSize >> rank_bits; + int const stride = (blockDim.x >> rank_bits) * VEC_SIZE; + // Prefetch low-precision vectors, then combine adjacent sources in tree order. + for (int base = warp * vectors_per_warp * VEC_SIZE; base < size_per_token; base += stride) + { + int const offset = base + output_lane * VEC_SIZE; + bool const valid_output = offset < size_per_token; + vec_t packed[8]; #pragma unroll - for (int j = 0; j < elems_per_vec; ++j) - { - acc[0][j] += acc[8][j]; - } - } - else if constexpr (TOP_K == 10) + for (int k = 0; k < 8; ++k) { -#pragma unroll - for (int j = 0; j < elems_per_vec; ++j) + packed[k].fill(InputT(0.0f)); + int const source = 8 * rank_lane + k; + if (valid_output && source < source_count) { - acc[0][j] += acc[1][j]; - acc[2][j] += acc[3][j]; - acc[4][j] += acc[5][j]; - acc[6][j] += acc[7][j]; - acc[8][j] += acc[9][j]; - } -#pragma unroll - for (int j = 0; j < elems_per_vec; ++j) - { - acc[0][j] += acc[2][j]; - acc[4][j] += acc[6][j]; - } -#pragma unroll - for (int j = 0; j < elems_per_vec; ++j) - { - acc[0][j] += acc[4][j]; - acc[0][j] += acc[8][j]; + packed[k].load(reinterpret_cast(sources[source] + offset)); } } - else if constexpr (TOP_K == 8) + vec_t value; + if constexpr (std::is_same_v && kElements % 2 == 0) { #pragma unroll - for (int j = 0; j < elems_per_vec; ++j) + for (int j = 0; j < kElements; j += 2) { - acc[0][j] += acc[1][j]; - acc[2][j] += acc[3][j]; - acc[4][j] += acc[5][j]; - acc[6][j] += acc[7][j]; - } + float2 pair[8]; #pragma unroll - for (int j = 0; j < elems_per_vec; ++j) - { - acc[0][j] += acc[2][j]; - acc[4][j] += acc[6][j]; - } -#pragma unroll - for (int j = 0; j < elems_per_vec; ++j) - { - acc[0][j] += acc[4][j]; + for (int k = 0; k < 8; ++k) + { + __nv_fp8x2_e4m3 input; + input.__x = reinterpret_cast(&packed[k])[j / 2]; + pair[k] = static_cast(input); + } + float const x0 = pair[0].x + pair[1].x; + float const x1 = pair[2].x + pair[3].x; + float const x2 = pair[4].x + pair[5].x; + float const x3 = pair[6].x + pair[7].x; + float const y0 = pair[0].y + pair[1].y; + float const y1 = pair[2].y + pair[3].y; + float const y2 = pair[4].y + pair[5].y; + float const y3 = pair[6].y + pair[7].y; + value[j] = (x0 + x1) + (x2 + x3); + value[j + 1] = (y0 + y1) + (y2 + y3); } } - else if constexpr (TOP_K == 6) + else { #pragma unroll - for (int j = 0; j < elems_per_vec; ++j) - { - acc[0][j] += acc[1][j]; - acc[2][j] += acc[3][j]; - acc[4][j] += acc[5][j]; - } -#pragma unroll - for (int j = 0; j < elems_per_vec; ++j) + for (int j = 0; j < kElements; ++j) { - acc[0][j] += acc[2][j]; - acc[0][j] += acc[4][j]; + float const p0 = static_cast(packed[0][j]) + static_cast(packed[1][j]); + float const p1 = static_cast(packed[2][j]) + static_cast(packed[3][j]); + float const p2 = static_cast(packed[4][j]) + static_cast(packed[5][j]); + float const p3 = static_cast(packed[6][j]) + static_cast(packed[7][j]); + value[j] = (p0 + p1) + (p2 + p3); } } - else if constexpr (TOP_K == 4) - { #pragma unroll - for (int j = 0; j < elems_per_vec; ++j) + for (int step = 1; step < 4; step *= 2) + { + if (step < rank_lanes) { - acc[0][j] += acc[1][j]; - acc[2][j] += acc[3][j]; - } #pragma unroll - for (int j = 0; j < elems_per_vec; ++j) - { - acc[0][j] += acc[2][j]; + for (int j = 0; j < kElements; ++j) + { + float const next = __shfl_down_sync(~0U, value[j], step, rank_lanes); + if ((rank_lane & (2 * step - 1)) == 0) + { + value[j] += next; + } + } } } - else if constexpr (TOP_K == 2) + if (rank_lane == 0 && valid_output) { -#pragma unroll - for (int j = 0; j < elems_per_vec; ++j) - { - acc[0][j] += acc[1][j]; - } + value.cast_store(output + offset / static_cast(sizeof(InputT))); } - else if constexpr (TOP_K == 1) + } +} + +// Pack valid source pointers in routing order and count the contributing ranks. +// stride_per_token can exceed the wire size for in-place low-precision combine. +template +__device__ void vectorized_combine( + OutputT* output, int size_per_token, int stride_per_token, CombineKernelPointers const& ptrs) +{ + static_assert(TOP_K > 0 && TOP_K <= 32, "combine routing requires TOP_K <= warpSize"); + constexpr uint32_t kRoutingLaneMask = make_warp_lane_mask(); + __shared__ uint8_t const* sources[TOP_K]; + __shared__ int source_count; + if (threadIdx.x < TOP_K) + { + int const k = threadIdx.x; + int const index = blockIdx.x * TOP_K + k; + int const target_index = ptrs.topk_target_indices[index]; + uint32_t const valid = __ballot_sync(kRoutingLaneMask, target_index >= 0); + if (k == 0) { - // nothing to do + source_count = __popc(valid); } - else + if (target_index >= 0) { - // Generic fallback: accumulate all into acc[0] -#pragma unroll - for (int k = 1; k < TOP_K; ++k) - { -#pragma unroll - for (int j = 0; j < elems_per_vec; ++j) - { - acc[0][j] += acc[k][j]; - } - } + int const peer = ptrs.topk_target_ranks[index]; + int const compact_index = __popc(valid & ((1U << k) - 1U)); + sources[compact_index] = ptrs.source_buffers[peer] + static_cast(target_index) * stride_per_token; } - - // cast_store converts each accumulated element to OutputT before the vectorized store. - acc[0].cast_store(dst_typed_base + offset / static_cast(sizeof(InputT))); } -} + __syncthreads(); -// Wrapper that selects vector width based on size_per_token alignment. -// stride_per_token: byte distance between tokens in the recv buffer (may differ from -// size_per_token when low-precision in-place data retains its payload-dtype workspace stride). -// InputT is the input element type in the receive buffer. -template -__device__ void vectorized_combine(OutputT* dst_typed_base, int size_per_token, int stride_per_token, int rank_id, - int max_tokens_per_rank, CombineKernelPointers const& ptrs) -{ - // Each branch is guarded by if constexpr (sizeof(InputT) <= VEC_SIZE) so that the compiler - // never instantiates vectorized_combine_impl with elems_per_vec=0. - // Branches where VEC_SIZE < sizeof(InputT) are unreachable at runtime because size_per_token - // is always a multiple of sizeof(InputT), so a larger alignment branch is taken first. if (size_per_token % 16 == 0) { - if constexpr (static_cast(sizeof(InputT)) <= 16) - vectorized_combine_impl<16, TOP_K, OutputT, InputT>( - dst_typed_base, size_per_token, stride_per_token, rank_id, max_tokens_per_rank, ptrs); + vectorized_combine_impl<16, OutputT, InputT>(output, size_per_token, sources, source_count); } else if (size_per_token % 8 == 0) { - if constexpr (static_cast(sizeof(InputT)) <= 8) - vectorized_combine_impl<8, TOP_K, OutputT, InputT>( - dst_typed_base, size_per_token, stride_per_token, rank_id, max_tokens_per_rank, ptrs); + vectorized_combine_impl<8, OutputT, InputT>(output, size_per_token, sources, source_count); } else if (size_per_token % 4 == 0) { - if constexpr (static_cast(sizeof(InputT)) <= 4) - vectorized_combine_impl<4, TOP_K, OutputT, InputT>( - dst_typed_base, size_per_token, stride_per_token, rank_id, max_tokens_per_rank, ptrs); + vectorized_combine_impl<4, OutputT, InputT>(output, size_per_token, sources, source_count); } else if (size_per_token % 2 == 0) { - if constexpr (static_cast(sizeof(InputT)) <= 2) - vectorized_combine_impl<2, TOP_K, OutputT, InputT>( - dst_typed_base, size_per_token, stride_per_token, rank_id, max_tokens_per_rank, ptrs); + vectorized_combine_impl<2, OutputT, InputT>(output, size_per_token, sources, source_count); } - else + else if constexpr (sizeof(InputT) == 1) { - if constexpr (static_cast(sizeof(InputT)) <= 1) - vectorized_combine_impl<1, TOP_K, OutputT, InputT>( - dst_typed_base, size_per_token, stride_per_token, rank_id, max_tokens_per_rank, ptrs); + vectorized_combine_impl<1, OutputT, InputT>(output, size_per_token, sources, source_count); } } @@ -1695,13 +1540,15 @@ __device__ void vectorized_quant(DstT* dst, SrcT const* src, int num_elements) vectorized_quant_impl<1, SrcT, DstT>(dst, src, num_elements); } -// LOW_PRECISION=false: vectorized byte-copy (SrcT = payload dtype). -// LOW_PRECISION=true: vectorized SrcT→FP8 quantization via vectorized_quant. +// Advance flag_val to the combine phase and prepare valid tokens in the requested range. +// Copy SrcT payloads, or quantize them to FP8 when LOW_PRECISION is enabled. +// CFT self contributions go to combine_recv_base; other prepared tokens go to combine_input_base. +// This kernel performs no remote transfers or reduction. template -__global__ void moeA2APrepareCombineKernel(uint8_t* recv_buffer_bytes, void const* source_payload, +__global__ void moeA2APrepareCombineKernel(uint8_t* combine_input_base, void const* source_payload, int elements_per_token, int ep_size, int max_tokens_per_rank, uint32_t* flag_val_ptr, int const* recv_counters, int source_stride_per_token, int workspace_stride_per_token, int prepare_first_token, int prepare_num_tokens, - uint8_t* region_c_base, int ep_rank) + uint8_t* combine_recv_base, int ep_rank) { #if TLLM_MOE_A2A_COMPILE_SM90 cudaGridDependencySynchronize(); @@ -1714,7 +1561,6 @@ __global__ void moeA2APrepareCombineKernel(uint8_t* recv_buffer_bytes, void cons { *flag_val_ptr = *flag_val_ptr + 1; } - // NOTE: LE counters are NOT zeroed. They grow monotonically with cumulative baselines. if (blockIdx.x >= prepare_num_tokens) return; @@ -1730,7 +1576,7 @@ __global__ void moeA2APrepareCombineKernel(uint8_t* recv_buffer_bytes, void cons // CFT combine stages local tokens compactly into the dedicated receive region. This keeps // local and peer contributions in one uniform layout without an in-place write-after-read hazard. - bool const stage_self = (region_c_base != nullptr && rank_idx == ep_rank); + bool const stage_self = (combine_recv_base != nullptr && rank_idx == ep_rank); size_t const source_offset = static_cast(global_token_idx) * source_stride_per_token; size_t const workspace_offset = static_cast(global_token_idx) * workspace_stride_per_token; @@ -1742,16 +1588,16 @@ __global__ void moeA2APrepareCombineKernel(uint8_t* recv_buffer_bytes, void cons { SrcT const* src_ptr = reinterpret_cast(static_cast(source_payload) + source_offset); - // Self -> region C (compact, separate buffer). Peer -> in-place workspace (push reads it). - __nv_fp8_e4m3* dst_ptr = stage_self ? reinterpret_cast<__nv_fp8_e4m3*>(region_c_base + self_slot_offset) - : reinterpret_cast<__nv_fp8_e4m3*>(recv_buffer_bytes + workspace_offset); + // Self contributions go directly to the receive inbox; peer contributions are staged for the push. + __nv_fp8_e4m3* dst_ptr = stage_self ? reinterpret_cast<__nv_fp8_e4m3*>(combine_recv_base + self_slot_offset) + : reinterpret_cast<__nv_fp8_e4m3*>(combine_input_base + workspace_offset); vectorized_quant(dst_ptr, src_ptr, elements_per_token); } else { // Same-type byte copy. CFT self tokens use the receive region; fence combine uses the workspace. uint8_t const* src = static_cast(source_payload) + source_offset; - uint8_t* dst = stage_self ? (region_c_base + self_slot_offset) : (recv_buffer_bytes + workspace_offset); + uint8_t* dst = stage_self ? (combine_recv_base + self_slot_offset) : (combine_input_base + workspace_offset); vectorized_copy(dst, src, elements_per_token * static_cast(sizeof(SrcT))); } } @@ -1761,10 +1607,8 @@ __global__ void moeA2APrepareCombineKernel(uint8_t* recv_buffer_bytes, void cons // ============================================================================ template -__global__ void moeA2ACombineKernel( - const CombineKernelPointers ptrs, // Combine-specific struct, src_data_ptrs[0] is output - int max_tokens_per_rank, int elements_per_token, int local_num_tokens, int rank_id, int ep_size, - int stride_per_token) +__global__ void moeA2ACombineKernel(const CombineKernelPointers ptrs, int max_tokens_per_rank, int elements_per_token, + int local_num_tokens, int rank_id, int ep_size, int stride_per_token) { using InputT = std::conditional_t; @@ -1813,7 +1657,7 @@ __global__ void moeA2ACombineKernel( continue; } uint32_t* flag_addr = &ptrs.completion_flags[peer_rank][rank_id]; - asm volatile("st.relaxed.sys.u32 [%0], %1;" ::"l"(flag_addr), "r"(expected_value)); + publish_round_flag(flag_addr, expected_value); #if ENABLE_DEBUG_PRINT printf("combine: +++Rank %d setting completion flag to %d for rank %d\n", rank_id, expected_value, peer_rank); @@ -1831,28 +1675,9 @@ __global__ void moeA2ACombineKernel( if (!is_rank_active(ptrs.active_rank_mask, peer_rank)) continue; } - bool flag_set = false; - auto s = clock64(); - do - { - uint32_t* flag_ptr = &ptrs.completion_flags[rank_id][peer_rank]; - uint32_t flag_value; - // Acquire load to ensure visibility of peer's release-store - asm volatile("ld.relaxed.sys.u32 %0, [%1];" : "=r"(flag_value) : "l"(flag_ptr)); -#if ENABLE_DEBUG_PRINT - printf( - "combine: ---Rank %d received completion flag from rank %d, flag_value: %d, expected_value: " - "%d, " - "address: %p\n", - rank_id, peer_rank, flag_value, expected_value, flag_ptr); -#endif - flag_set = flag_value == expected_value; - } while (!flag_set && !check_timeout(s, ptrs.timeout_cycles)); - - if (__builtin_expect(!flag_set, 0)) + if (!wait_round_flag(&ptrs.completion_flags[rank_id][peer_rank], expected_value, ptrs.timeout_cycles, + rank_id, peer_rank, "combine")) { - printf("combine: ---Rank %d timed out waiting for completion flag from rank %d\n", rank_id, peer_rank); - asm volatile("trap;"); return; } } @@ -1869,9 +1694,8 @@ __global__ void moeA2ACombineKernel( if (local_num_tokens == 0) return; - T* token_output = static_cast(ptrs.src_data_ptrs[0]) + local_token_idx * elements_per_token; - vectorized_combine( - token_output, size_per_token, stride_per_token, rank_id, max_tokens_per_rank, ptrs); + T* token_output = static_cast(ptrs.output) + local_token_idx * elements_per_token; + vectorized_combine(token_output, size_per_token, stride_per_token, ptrs); #if TLLM_MOE_A2A_COMPILE_SM90 cudaTriggerProgrammaticLaunchCompletion(); #endif @@ -1891,31 +1715,42 @@ __global__ void moeA2ACombineKernel( static constexpr int kCombinePushWarpsPerBlock = 4; template -__global__ void moeA2ACftCombinePushKernel( +__global__ void moeA2ACombinePushKernel_Cft( uint8_t const* local_payload, // Expert output (combine payload or dispatch recv_buffer) int const* recv_counters, // [2, ep_size] tokens received from each source rank uint32_t const* flag_val, - CftPeerLeIds peer_le_ids, // LE IDs passed by value (no device pointer needed) + CftCombinePeerInfo peer_info, // Peer LE IDs and readiness flag pointers passed by value. int rank_id, int ep_size, int max_tokens_per_rank, int bytes_per_token, uint64_t combine_payload_base, uint64_t combine_counter_base, int combine_counter_ep_stride, int local_stride_per_token) { -#if (defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 1000) || CLANGD_HOST_PASS +#if TLLM_MOE_A2A_COMPILE_SM100 // Wait for prepareCombine to finish writing the workspace we read from, then immediately - // signal the next kernel (combineCountedWrite) that it can start. combineCountedWrite + // signal moeA2ACombineKernel_Cft that it can start. moeA2ACombineKernel_Cft // polls for incoming counter writes from peers — it touches disjoint memory from our // local pushes, so it can run concurrently with the rest of this kernel. cudaGridDependencySynchronize(); cudaTriggerProgrammaticLaunchCompletion(); int source_rank = blockIdx.x; - if (source_rank == rank_id) - return; if constexpr (ENABLE_RANK_MASK) { - if (!is_rank_active(peer_le_ids.active_rank_mask, source_rank)) + if (!is_rank_active(peer_info.active_rank_mask, source_rank)) return; } +#if !DISABLE_SYNC_FOR_PROFILING + // The dependency wait has completed upstream MoE reads of the dispatch region. + // Publish readiness even when this peer has no contribution to receive. + if (blockIdx.y == 0 && threadIdx.x == 0) + { + uint32_t* flag_addr = &peer_info.completion_flags[source_rank][rank_id]; + uint32_t const expected_value = *flag_val; + publish_round_flag(flag_addr, expected_value); + } +#endif + if (source_rank == rank_id) + return; + uint32_t const parity = round_parity(*flag_val); int num_tokens = recv_counters[parity * ep_size + source_rank]; // Nothing to push (0 tokens) or an out-of-range count (corrupt recv_counters): skip. @@ -1928,18 +1763,16 @@ __global__ void moeA2ACftCombinePushKernel( return; // only lane 0 of each warp drives TMA + fabric // Per-warp smem layout: [mbarrier slot (kCftMbarrierSlotBytes) | staging (bytes_per_token)] - // repeated kCombinePushWarpsPerBlock times. Each warp drives its own slot independently - // and tracks its own put completion via put_bar — no __syncthreads, no CTA-scope drain - // serialization. submit + wait_reads are issued per-warp; the fabric engine pipelines - // drains across warps. + // repeated kCombinePushWarpsPerBlock times. Each warp drives its own slot and + // issues submit + wait_reads independently, without a CTA barrier between tokens. extern __shared__ uint8_t smem_push[]; int per_warp_bytes = kCftMbarrierSlotBytes + bytes_per_token; uint8_t* warp_smem = smem_push + warp_id * per_warp_bytes; - __mbarrier_t* tma_bar = reinterpret_cast<__mbarrier_t*>(warp_smem); - __mbarrier_t* put_bar = tma_bar + 1; + __mbarrier_t* tma_bar = reinterpret_cast<__mbarrier_t*>(warp_smem); // Tracks global-to-shared staging. + __mbarrier_t* put_bar = tma_bar + 1; // Fabric put report target; not polled by this kernel. uint8_t* staging = warp_smem + kCftMbarrierSlotBytes; - uint32_t le_id = peer_le_ids.ids[source_rank]; // push back to source rank's LE + uint32_t le_id = peer_info.ids[source_rank]; // push back to source rank's LE // Tokens are fanned across all warps of all blocks for this source rank; blockIdx.y selects // the token-chunk. Each warp uses its own per-warp smem slot. @@ -1960,15 +1793,12 @@ __global__ void moeA2ACftCombinePushKernel( cft_barrier_wait_parity(tma_bar, tma_phase & 1); tma_phase++; - // Issue the fabric put with put_bar tracking; arm put_bar to expect bytes_per_token - // bytes of fabric.report::fabric.counted::bytes events. + // Push the staged token to its source rank and increment that receive slot's byte counter. uint64_t data_offset = combine_payload_base + (static_cast(rank_id) * max_tokens_per_rank + t) * bytes_per_token; uint64_t counter_offset = combine_counter_base + (static_cast(rank_id) * combine_counter_ep_stride + t) * kCftCounterStride; - // put_bar is a required mbarrier::report destination for the PTX but is not waited on: - // mbarrier::report::fabric does not deliver reports on this Rubin/driver combo, so - // smem-reuse completion is enforced via CTA-scope fabric.wait.sync_restrict::reads below. + // wait_reads protects staging reuse; the receiver polls its byte counter for arrival. cft_fabric_try_put_counted(le_id, data_offset, counter_offset, staging, bytes_per_token, put_bar); cft_fabric_submit(); cft_fabric_wait_reads(); @@ -1981,89 +1811,114 @@ __global__ void moeA2ACftCombinePushKernel( } template -__global__ void moeA2ACombineCountedWriteKernel(const CombineKernelPointers ptrs, int max_tokens_per_rank, - int elements_per_token, int local_num_tokens, int rank_id) +__global__ void moeA2ACombineKernel_Cft(const CombineKernelPointers ptrs, int max_tokens_per_rank, + int elements_per_token, int local_num_tokens, int rank_id, int ep_size) { using InputT = std::conditional_t; -#if (defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 1000) || CLANGD_HOST_PASS +#if TLLM_MOE_A2A_COMPILE_SM100 int local_token_idx = blockIdx.x; int const size_per_token = elements_per_token * sizeof(InputT); cudaGridDependencySynchronize(); cudaTriggerProgrammaticLaunchCompletion(); - // Empty rank: this block exists only for the PDL handshake above; no local token to reduce. - if (local_num_tokens == 0) - return; + // Empty ranks skip reduction but still participate in the readiness wait below. + if (local_num_tokens > 0) + { #if !DISABLE_SYNC_FOR_PROFILING - // Per-token readiness: warp 0 polls ONLY the k receive-slots its local token needs - // (slot = target_rank*max + dst_idx), so a token reduces as soon as its own pieces land - // (overlapping the still-running push under PDL) rather than waiting for every peer's counter. - int lane_id = threadIdx.x % warpSize; - // One token per block -> only warp 0 polls its receive-slots. - if (threadIdx.x / warpSize == 0) - { - int const my_token = local_token_idx; -#pragma unroll 1 - for (int kk = lane_id; kk < TOP_K; kk += warpSize) + // Per-token readiness: warp 0 polls ONLY the k receive-slots its local token needs + // (slot = target_rank*max + dst_idx), so a token reduces as soon as its own pieces land + // (overlapping the still-running push under PDL) rather than waiting for every peer's counter. + int lane_id = threadIdx.x % warpSize; + // One token per block -> only warp 0 polls its receive-slots. + if (threadIdx.x / warpSize == 0) { - int tr = ptrs.topk_target_ranks[my_token * TOP_K + kk]; - int di = ptrs.topk_send_indices[my_token * TOP_K + kk]; - if (tr < 0 || di < 0) - continue; // duplicate / invalid routing slot - if constexpr (ENABLE_RANK_MASK) - { - if (!is_rank_active(ptrs.active_rank_mask, tr)) - continue; - } - if (tr == rank_id) - continue; // self contribution: not fabric-pushed - - int slot = tr * ptrs.combine_counter_ep_stride + di; - uint64_t combine_base = ptrs.combine_counter_baseline[slot]; - uint64_t combine_target = combine_base + static_cast(size_per_token); - - uint64_t* combineCounterPtr = &ptrs.combine_counters[static_cast(slot) * kCftCounterStrideU64]; - uint64_t current_combine_counter = 0; - auto s = clock64(); - while (true) + int const my_token = local_token_idx; +#pragma unroll 1 + for (int kk = lane_id; kk < TOP_K; kk += warpSize) { - asm volatile("ld.relaxed.sys.u64 %0, [%1];" : "=l"(current_combine_counter) : "l"(combineCounterPtr)); - if (current_combine_counter >= combine_target) + int tr = ptrs.topk_target_ranks[my_token * TOP_K + kk]; + int di = ptrs.topk_target_indices[my_token * TOP_K + kk]; + if (tr < 0 || di < 0) + continue; // duplicate / invalid routing slot + if constexpr (ENABLE_RANK_MASK) { - break; + if (!is_rank_active(ptrs.active_rank_mask, tr)) + continue; } - if (check_timeout(s, ptrs.timeout_cycles)) + if (tr == rank_id) + continue; // self contribution: not fabric-pushed + + int slot = tr * ptrs.combine_counter_ep_stride + di; + uint64_t combine_base = ptrs.combine_counter_baseline[slot]; + uint64_t combine_target = combine_base + static_cast(size_per_token); + + uint64_t* combineCounterPtr = &ptrs.combine_counters[static_cast(slot) * kCftCounterStrideU64]; + uint64_t current_combine_counter = 0; + auto s = clock64(); + while (true) { - printf( - "combine(cft): ---Rank %d tok %d k %d slot %d timed out counter=%llu base=%llu target=%llu\n", - rank_id, my_token, kk, slot, (unsigned long long) current_combine_counter, - (unsigned long long) combine_base, (unsigned long long) combine_target); - asm volatile("trap;"); - return; + asm volatile("ld.relaxed.sys.u64 %0, [%1];" + : "=l"(current_combine_counter) + : "l"(combineCounterPtr)); + if (current_combine_counter >= combine_target) + { + break; + } + if (check_timeout(s, ptrs.timeout_cycles)) + { + printf( + "combine(cft): ---Rank %d tok %d k %d slot %d timed out counter=%llu base=%llu " + "target=%llu\n", + rank_id, my_token, kk, slot, (unsigned long long) current_combine_counter, + (unsigned long long) combine_base, (unsigned long long) combine_target); + asm volatile("trap;"); + return; + } } + ptrs.combine_counter_baseline[slot] = combine_target; } - ptrs.combine_counter_baseline[slot] = combine_target; - } #if TLLM_CFT_HAS_CUDA_13_4_SUPPORT - asm volatile("fence.proxy.generic::fabric.alias.acquire.sys;" ::: "memory"); + asm volatile("fence.proxy.generic::fabric.alias.acquire.sys;" ::: "memory"); #endif - } - __syncthreads(); + } + __syncthreads(); #endif - __threadfence_system(); // system-scope fence before the gather (dispatch gets this from its kernel boundary) - T* token_output = static_cast(ptrs.src_data_ptrs[0]) + local_token_idx * elements_per_token; - vectorized_combine( - token_output, size_per_token, size_per_token, rank_id, max_tokens_per_rank, ptrs); + T* token_output = static_cast(ptrs.output) + local_token_idx * elements_per_token; + vectorized_combine(token_output, size_per_token, size_per_token, ptrs); + } + +#if !DISABLE_SYNC_FOR_PROFILING + // Keep one CTA alive until every peer has consumed its old dispatch inputs. + // Other token CTAs reduce independently; the next dispatch's dependency wait + // prevents it from overwriting peer inputs before this grid completes. + if (blockIdx.x == 0 && threadIdx.x < warpSize) + { + uint32_t const expected_value = *ptrs.flag_val; + for (int peer_rank = threadIdx.x; peer_rank < ep_size; peer_rank += warpSize) + { + if constexpr (ENABLE_RANK_MASK) + { + if (!is_rank_active(ptrs.active_rank_mask, peer_rank)) + continue; + } + if (!wait_round_flag(&ptrs.completion_flags[rank_id][peer_rank], expected_value, ptrs.timeout_cycles, + rank_id, peer_rank, "combine(cft)")) + { + return; + } + } + } +#endif cudaTriggerProgrammaticLaunchCompletion(); #else // Launched only on the CFT path, which requires sm_100+; fail loudly rather than // completing with an empty body. asm volatile("trap;" ::: "memory"); -#endif // __CUDA_ARCH__ >= 1000 +#endif // TLLM_MOE_A2A_COMPILE_SM100 } void moe_a2a_cft_combine_push_launch(MoeA2ACombineParams const& params) @@ -2074,17 +1929,20 @@ void moe_a2a_cft_combine_push_launch(MoeA2ACombineParams const& params) uint8_t const* local_payload = static_cast(params.cft_push_payload); // Pass peer metadata by value as a kernel argument. - CftPeerLeIds le_ids = {}; + CftCombinePeerInfo peer_info = {}; for (int i = 0; i < params.ep_size; i++) - le_ids.ids[i] = params.cft_peer_le_ids[i]; + { + peer_info.ids[i] = params.cft_peer_le_ids[i]; + peer_info.completion_flags[i] = params.completion_flags[i]; + } for (int w = 0; w < kRankMaskWords; ++w) - le_ids.active_rank_mask[w] = params.active_rank_mask[w]; + peer_info.active_rank_mask[w] = params.active_rank_mask[w]; // Push parallelism is env-overridable for tuning: - // TRTLLM_CFT_PUSH_WARPS : warps per block (default kCombinePushWarpsPerBlock) - // TRTLLM_CFT_PUSH_BLOCKS_PER_RANK : blocks per source rank == grid.y + // TRTLLM_NVLINK_ONE_SIDED_A2A_CFT_PUSH_WARPS : warps per block (default kCombinePushWarpsPerBlock) + // TRTLLM_NVLINK_ONE_SIDED_A2A_CFT_PUSH_BLOCKS_PER_RANK : blocks per source rank == grid.y int push_warps = kCombinePushWarpsPerBlock; - if (char const* e = std::getenv("TRTLLM_CFT_PUSH_WARPS")) + if (char const* e = std::getenv("TRTLLM_NVLINK_ONE_SIDED_A2A_CFT_PUSH_WARPS")) { int v = std::atoi(e); if (v >= 1) @@ -2096,7 +1954,7 @@ void moe_a2a_cft_combine_push_launch(MoeA2ACombineParams const& params) blocks_per_rank = 1; if (blocks_per_rank > 32) blocks_per_rank = 32; - if (char const* e = std::getenv("TRTLLM_CFT_PUSH_BLOCKS_PER_RANK")) + if (char const* e = std::getenv("TRTLLM_NVLINK_ONE_SIDED_A2A_CFT_PUSH_BLOCKS_PER_RANK")) { int v = std::atoi(e); if (v >= 1) @@ -2113,17 +1971,17 @@ void moe_a2a_cft_combine_push_launch(MoeA2ACombineParams const& params) auto set_attr = [&](auto* kernel_fn) { TLLM_CUDA_CHECK(cudaFuncSetAttribute(kernel_fn, cudaFuncAttributeMaxDynamicSharedMemorySize, smem_size)); }; if (params.enable_rank_mask) - set_attr(moeA2ACftCombinePushKernel); + set_attr(moeA2ACombinePushKernel_Cft); else - set_attr(moeA2ACftCombinePushKernel); + set_attr(moeA2ACombinePushKernel_Cft); } SWITCH_BOOL(params.enable_rank_mask, ENABLE_RANK_MASK, { - auto kernel_fn = moeA2ACftCombinePushKernel; - launchWithPdlWhenEnabled("moeA2ACftCombinePushKernel", kernel_fn, dim3(params.ep_size, blocks_per_rank), - dim3(blockThreads), smem_size, params.stream, local_payload, params.recv_counters, params.flag_val, le_ids, - params.ep_rank, params.ep_size, params.max_tokens_per_rank, bytes_per_token, - params.cft_le_combine_payload_base, params.cft_le_combine_counter_base, params.combine_counter_ep_stride, + auto kernel_fn = moeA2ACombinePushKernel_Cft; + launchWithPdlWhenEnabled("moeA2ACombinePushKernel_Cft", kernel_fn, dim3(params.ep_size, blocks_per_rank), + dim3(blockThreads), smem_size, params.stream, local_payload, params.recv_counters, params.flag_val, + peer_info, params.ep_rank, params.ep_size, params.max_tokens_per_rank, bytes_per_token, + params.cft_combine_recv_offset, params.cft_le_combine_counter_base, params.combine_counter_ep_stride, local_stride_per_token); }); } @@ -2133,25 +1991,22 @@ void moe_a2a_prepare_combine_launch(MoeA2ACombineParams const& params) constexpr int kBlockSize = 256; TLLM_CHECK(params.max_tokens_per_rank > 0); - uint8_t* recv_buffer_bytes = static_cast(const_cast(params.recv_buffers[params.ep_rank])); + uint8_t* combine_input_base = params.combine_input_buffers[params.ep_rank]; // CFT combine stages local contributions compactly into its dedicated receive region. - uint8_t* const region_c_base - = params.use_cft_for_combine ? static_cast(const_cast(params.cft_le_combine_recv)) : nullptr; + uint8_t* const combine_recv_base = params.use_cft_for_combine ? params.cft_combine_recv_payload : nullptr; int const grid = std::max(params.prepare_num_tokens, 1); - // Zero LE-backed counters from HOST before kernel launch. - // NOTE: Combine LE counters are zeroed in prepare_dispatch_launch (before any fabric activity). - // Zeroing them here (after dispatch's fabric puts) corrupts subsequent counter increments - // because cudaDeviceSynchronize does NOT wait for fabric engine completion. + // Preserve params.cft_le_combine_counters and params.cft_combine_counter_baseline + // across rounds. The CFT reduce kernel advances each slot's baseline after its wait completes. SWITCH_BOOL(params.use_low_precision, LOW_PRECISION, { SWITCH_DTYPE(params.dtype, SrcT, { auto kernel_fn = moeA2APrepareCombineKernel; launchWithPdlWhenEnabled("moeA2APrepareCombineKernel", kernel_fn, grid, kBlockSize, 0, params.stream, - recv_buffer_bytes, params.source_payload, params.elements_per_token, params.ep_size, + combine_input_base, params.source_payload, params.elements_per_token, params.ep_size, params.max_tokens_per_rank, params.flag_val, params.recv_counters, params.source_stride_per_token, - params.workspace_stride_per_token, params.prepare_first_token, params.prepare_num_tokens, region_c_base, - params.ep_rank); + params.workspace_stride_per_token, params.prepare_first_token, params.prepare_num_tokens, + combine_recv_base, params.ep_rank); }); }); } @@ -2162,6 +2017,8 @@ void moe_a2a_prepare_combine_launch(MoeA2ACombineParams const& params) void moe_a2a_combine_launch(MoeA2ACombineParams const& params) { + constexpr int kBlockSize = 256; + // Validate parameters TLLM_CHECK(params.top_k > 0 && params.top_k <= kMaxTopK); TLLM_CHECK(params.ep_size > 0 && params.ep_size <= kMaxRanks); @@ -2187,21 +2044,16 @@ void moe_a2a_combine_launch(MoeA2ACombineParams const& params) { cft_grid = 1; } - int const cft_block = tensorrt_llm::common::getEnvMoeA2ACombineBlockSize(); CombineKernelPointers kp = {}; - kp.src_data_ptrs[0] = params.output_data; - for (int rank = 0; rank < params.ep_size; rank++) - { - kp.recv_buffers[rank][0] = params.recv_buffers[rank]; - } + kp.output = params.output_data; for (int i = 0; i < params.ep_size; i++) { kp.completion_flags[i] = params.completion_flags[i]; } kp.flag_val = params.flag_val; kp.topk_target_ranks = params.topk_target_ranks; - kp.topk_send_indices = params.topk_send_indices; + kp.topk_target_indices = params.topk_target_indices; // CFT combine metadata. kp.combine_counters = params.cft_le_combine_counters; @@ -2213,28 +2065,21 @@ void moe_a2a_combine_launch(MoeA2ACombineParams const& params) kp.active_rank_mask[w] = params.active_rank_mask[w]; } - // Offset-trick gather: peers' pushed data lands in THIS rank's region C - // (cft_le_combine_recv). recv_buffers[P] = region_C_base + (P - S) * stride so the - // reduce reads peer P's contribution at the same slot layout as the fence path. - uint8_t const* combine_base = static_cast(params.cft_le_combine_recv); - int const element_size = static_cast(tensorrt_llm::common::getDTypeSize(params.dtype)); - // Every contribution uses the same compact wire layout in the receive region. - int const bytes_per_token = params.elements_per_token * (params.use_low_precision ? 1 : element_size); - int64_t const peer_src_stride_per_rank = static_cast(params.max_tokens_per_rank) * bytes_per_token; + // A source rank's pushed tokens occupy one runtime-packed slice in the local inbox. + int64_t const peer_stride = static_cast(params.max_tokens_per_rank) * params.wire_bytes_per_token; for (int rank = 0; rank < params.ep_size; rank++) { - // Local tokens occupy the zero-offset slice; peer slices are addressed relative to it. - kp.recv_buffers[rank][0] = combine_base + (rank - params.ep_rank) * peer_src_stride_per_rank; + kp.source_buffers[rank] = params.cft_combine_recv_payload + rank * peer_stride; } SWITCH_BOOL(params.enable_rank_mask, ENABLE_RANK_MASK, { SWITCH_DTYPE(params.dtype, T, { SWITCH_BOOL(params.use_low_precision, LOW_PRECISION, { SWITCH_TOP_K(params.top_k, TOP_K, { - auto kernel_fn = moeA2ACombineCountedWriteKernel; - launchWithPdlWhenEnabled("moeA2ACombineCountedWriteKernel", kernel_fn, cft_grid, cft_block, 0, + auto kernel_fn = moeA2ACombineKernel_Cft; + launchWithPdlWhenEnabled("moeA2ACombineKernel_Cft", kernel_fn, cft_grid, kBlockSize, 0, params.stream, kp, params.max_tokens_per_rank, params.elements_per_token, - params.local_num_tokens, params.ep_rank); + params.local_num_tokens, params.ep_rank, params.ep_size); }); }); }); @@ -2243,7 +2088,6 @@ void moe_a2a_combine_launch(MoeA2ACombineParams const& params) } // Configure kernel launch (one block per token). - int const kBlockSize = tensorrt_llm::common::getEnvMoeA2ACombineBlockSize(); int grid = params.local_num_tokens; // If local_num_tokens is 0, we still need to launch a minimal kernel to participate in the synchronization. if (grid == 0) @@ -2255,13 +2099,14 @@ void moe_a2a_combine_launch(MoeA2ACombineParams const& params) CombineKernelPointers kernel_ptrs = {}; // Zero-initialize kernel_ptrs.timeout_cycles = params.timeout_cycles; - // Set output data pointer in src_data_ptrs[0] - kernel_ptrs.src_data_ptrs[0] = params.output_data; + kernel_ptrs.output = params.output_data; - // Fill recv buffer pointers + // Each expert rank stores this origin rank's tokens in its own combine input region. + int64_t const origin_offset + = static_cast(params.ep_rank) * params.max_tokens_per_rank * params.reduce_stride_per_token; for (int rank = 0; rank < params.ep_size; rank++) { - kernel_ptrs.recv_buffers[rank][0] = params.recv_buffers[rank]; + kernel_ptrs.source_buffers[rank] = params.combine_input_buffers[rank] + origin_offset; } // Copy completion flag pointers @@ -2273,7 +2118,7 @@ void moe_a2a_combine_launch(MoeA2ACombineParams const& params) // Copy communication tracking pointers kernel_ptrs.topk_target_ranks = params.topk_target_ranks; - kernel_ptrs.topk_send_indices = params.topk_send_indices; + kernel_ptrs.topk_target_indices = params.topk_target_indices; // Copy active-rank bitmask into the kernel pointers struct for (int w = 0; w < kRankMaskWords; ++w) diff --git a/cpp/tensorrt_llm/kernels/moe/communication/moeAlltoAllKernels.h b/cpp/tensorrt_llm/kernels/moe/communication/moeAlltoAllKernels.h index 0ab369f4a5aa..316e7bc7b9fb 100644 --- a/cpp/tensorrt_llm/kernels/moe/communication/moeAlltoAllKernels.h +++ b/cpp/tensorrt_llm/kernels/moe/communication/moeAlltoAllKernels.h @@ -45,14 +45,16 @@ static constexpr size_t kCftCounterStrideU64 = kCftCounterStride / sizeof(uint64 static constexpr int kCftMbarrierSlotBytes = 64; // Fixed-size peer metadata passed by value to the CFT combine push kernel. -struct CftPeerLeIds +struct CftCombinePeerInfo { uint32_t ids[kMaxRanks]; + uint32_t* completion_flags[kMaxRanks]; uint64_t active_rank_mask[kRankMaskWords]; }; -// Default completion-flag wait budget: 300 s at an assumed 2 GHz clock64 rate. -static constexpr int64_t kDefaultTimeoutCycles = 300ll * 2000ll * 1000ll * 1000ll; +// Nominal clock64 rate used to convert timeout seconds to SM cycles. +static constexpr int64_t kAssumedClockHz = 2000ll * 1000ll * 1000ll; +static constexpr int64_t kDefaultTimeoutCycles = 300ll * kAssumedClockHz; // Default per-block dynamic shared-memory cap on sm_90+; larger requests must opt in via // cudaFuncAttributeMaxDynamicSharedMemorySize. @@ -92,8 +94,8 @@ struct DispatchKernelPointers int* local_token_counter; // Atomic counter for completed tokens // Top-K compact routing info per local token (size: [local_num_tokens, top_k]) - int* topk_target_ranks; // target rank per k, -1 for invalid or duplicate routes - int* topk_send_indices; // dst index per k, -1 for invalid or duplicate routes + int* topk_target_ranks; // target rank per k, -1 for invalid or duplicate routes + int* topk_target_indices; // dst index per k, -1 for invalid or duplicate routes // Optional: Statistics for EPLB int const* eplb_local_stats; // [eplb_stats_num_experts] @@ -120,25 +122,25 @@ struct DispatchKernelPointers // The local rank's own bit must always be set; this is checked at launch time. uint64_t active_rank_mask[kRankMaskWords]; - // Completion-flag wait budget in clock64() cycles; see moeA2AGetTimeoutCycles(). + // Host-selected wait budget in clock64() cycles. int64_t timeout_cycles{kDefaultTimeoutCycles}; }; -// Combine kernel pointers - non-const output in src_data_ptrs[0], const recv buffers +// Gather one contribution slice per expert rank. struct CombineKernelPointers { - // Payload pointers - void* src_data_ptrs[kMaxPayloads]; // src_data_ptrs[0] is output - void const* recv_buffers[kMaxRanks][kMaxPayloads]; // 2D array of receive buffer pointers (const) + void* output; + // Fence: peer input slices. CFT: peer contributions in the local receive inbox. + uint8_t const* source_buffers[kMaxRanks]; - // Completion flags for synchronization (fence-based path) + // Combine readiness flags shared by the fence and CFT paths. uint32_t* completion_flags[kMaxRanks]; // If completion_flags[target_rank][source_rank] == *flag_val, then source // rank has signaled the target rank uint32_t* flag_val; // The value of the flag for this round (stored on the local rank) // Top-K compact routing info per local token (size: [local_num_tokens, top_k]) - int const* topk_target_ranks; // target rank per k, -1 for invalid or duplicate routes - int const* topk_send_indices; // dst index per k, -1 for invalid or duplicate routes + int const* topk_target_ranks; // target rank per k, -1 for invalid or duplicate routes + int const* topk_target_indices; // dst index per k, -1 for invalid or duplicate routes // ---- CFT combine (counted-write) fields. Unused by the fence combine path. ---- // Local LE combine counters: per receive-slot HW-incremented byte counters. @@ -150,7 +152,7 @@ struct CombineKernelPointers // completion flag writes/waits to/from inactive peers. uint64_t active_rank_mask[kRankMaskWords]; - // Completion-flag wait budget in clock64() cycles; see moeA2AGetTimeoutCycles(). + // Host-selected wait budget in clock64() cycles. int64_t timeout_cycles{kDefaultTimeoutCycles}; }; @@ -180,8 +182,8 @@ struct MoeA2ADispatchParams int* send_counters; // [ep_size] atomic counters - tracks tokens sent to each target rank int* topk_target_ranks; // Top-K compact routing info per local token (size: [local_num_tokens, top_k]), target rank // per k, -1 for duplicates - int* topk_send_indices; // Top-K compact routing info per local token (size: [local_num_tokens, top_k]), dst index - // per k, -1 for duplicates + int* topk_target_indices; // Top-K compact routing info per local token (size: [local_num_tokens, top_k]), dst index + // per k, -1 for duplicates // Distributed aux data and recv buffers // Each rank owns recv_counters[parity][source_rank]. The two parity banks @@ -222,22 +224,13 @@ struct MoeA2ADispatchParams // CUDA graph replay until generation-scoped invalidation and recapture are available. uint64_t active_rank_mask[kRankMaskWords] = {~uint64_t{0}, ~uint64_t{0}, ~uint64_t{0}, ~uint64_t{0}}; - // Completion-flag wait budget in clock64() cycles; see moeA2AGetTimeoutCycles(). + // Host-selected wait budget in clock64() cycles. int64_t timeout_cycles{kDefaultTimeoutCycles}; // CUDA stream cudaStream_t stream; }; -// Resolve the completion-flag wait budget, in clock64() cycles. -// -// No collective separates a rank's first-touch JIT/autotune work from its dispatch -// launch, so this device-side budget is in effect a deadline on the slowest peer's -// host-side progress. Warmup therefore uses a larger budget than steady state. -// Overridable via TRTLLM_MOE_A2A_TIMEOUT_SEC / TRTLLM_MOE_A2A_WARMUP_TIMEOUT_SEC. -// See nvbugs/6482566. -int64_t moeA2AGetTimeoutCycles(bool is_warmup); - // Dispatch kernels void moe_a2a_dispatch_launch(MoeA2ADispatchParams const& params); // Prepare for dispatch: zero send_counters, local_token_counter and increment flag_val @@ -279,8 +272,8 @@ struct MoeA2ACombineParams uint32_t* flag_val; // The value of the flag for this round (stored on the local rank) int* topk_target_ranks; // Top-K compact routing info per local token (size: [local_num_tokens, top_k]), target rank // per k, -1 for duplicates - int* topk_send_indices; // Top-K compact routing info per local token (size: [local_num_tokens, top_k]), dst index - // per k, -1 for duplicates + int* topk_target_indices; // Top-K compact routing info per local token (size: [local_num_tokens, top_k]), dst index + // per k, -1 for duplicates // Local recv_counters[parity][source_rank]. The two parity banks alternate // between A2A rounds. int const* recv_counters; @@ -288,17 +281,17 @@ struct MoeA2ACombineParams // Distributed aux data and recv buffers uint32_t* completion_flags[kMaxRanks]; // If completion_flags[target_rank][source_rank] == *flag_val, then source // rank has signaled the target rank - void const* recv_buffers[kMaxRanks]; // Per-rank receive buffers (only for single payload) + uint8_t* combine_input_buffers[kMaxRanks]; // Expert-output/staging region on each rank // ---- CFT combine (counted-write) path. Gated by use_cft_for_combine. ---- // When true, moe_a2a_combine_launch takes the CFT push+reduce path and the base fence // combine below is bypassed. The base fence combine is unaffected when false. bool use_cft_for_combine; uint32_t cft_peer_le_ids[kMaxRanks]; // LE ID per target rank - uint64_t cft_le_combine_payload_base; // LE byte offset for combine payload (region C) + uint64_t cft_combine_recv_offset; // LE byte offset of the combine receive inbox uint64_t cft_le_combine_counter_base; // LE byte offset for combine counters uint64_t* cft_le_combine_counters; // Direct pointer to local LE combine counters - void* cft_le_combine_recv; // Direct pointer to local LE combine payload region (C) + uint8_t* cft_combine_recv_payload; // Local CFT receive inbox, including the self contribution uint64_t* cft_combine_counter_baseline; // [ep_size * max_tokens_per_rank] regular device memory int combine_counter_ep_stride = 0; // STABLE static stride (maxNumTokens) for counter/baseline slot indexing @@ -312,7 +305,7 @@ struct MoeA2ACombineParams // CUDA graph replay until generation-scoped invalidation and recapture are available. uint64_t active_rank_mask[kRankMaskWords] = {~uint64_t{0}, ~uint64_t{0}, ~uint64_t{0}, ~uint64_t{0}}; - // Completion-flag wait budget in clock64() cycles; see moeA2AGetTimeoutCycles(). + // Host-selected wait budget in clock64() cycles. int64_t timeout_cycles{kDefaultTimeoutCycles}; // CUDA stream diff --git a/cpp/tensorrt_llm/thop/moe/communication/moeAlltoAllMeta.h b/cpp/tensorrt_llm/thop/moe/communication/moeAlltoAllMeta.h index 410034a8ae85..0e8f1aa10eab 100644 --- a/cpp/tensorrt_llm/thop/moe/communication/moeAlltoAllMeta.h +++ b/cpp/tensorrt_llm/thop/moe/communication/moeAlltoAllMeta.h @@ -17,6 +17,7 @@ #pragma once #include "tensorrt_llm/common/config.h" +#include "tensorrt_llm/kernels/moe/communication/moeAlltoAllKernels.h" #include #include @@ -30,50 +31,112 @@ namespace torch_ext namespace moe_comm { -// Enum for indexing into moe_a2a_metainfo tensor +// Per-rank layout: round state, dispatch control/payload, combine control/input/receive. +// Fence combine pulls from peer input buffers; CFT pushes into the local receive inbox. +// Region boundaries are fixed at allocation; token slices remain runtime-packed. +// OFFSET fields are byte offsets from the per-rank workspace base; SIZE fields are byte capacities. +// Array extents below use allocation-time max_tokens, top_k and ep_size, not runtime token counts. +// CFT-only offsets and sizes are zero when CFT storage is not allocated. enum MoeA2AMetaInfoIndex : int64_t { + // Shared round state. + // uint32_t scalar advanced by dispatch/combine prepare; supplies sync epochs and round parity. FLAG_VAL_OFFSET_INDEX = 0, + + // Dispatch control, routing metadata and receive payload. + // int32_t scalar counting token CTAs to elect the last CTA that publishes dispatch counts. LOCAL_TOKEN_COUNTER_OFFSET_INDEX = 1, + // int32_t[ep_size]: outgoing token counts / slot allocators, indexed by destination rank. SEND_COUNTERS_OFFSET_INDEX = 2, + // int32_t[2][ep_size]: incoming counts by round parity and sender; dispatch writes, combine reads. RECV_COUNTERS_OFFSET_INDEX = 3, + // uint32_t[ep_size]: per-peer epoch flags for fence dispatch synchronization. DISPATCH_COMPLETION_FLAGS_OFFSET_INDEX = 4, - COMBINE_COMPLETION_FLAGS_OFFSET_INDEX = 5, - // Counted-write counters: uint64 per slot, kCftCounterStride-aligned. - DISPATCH_COUNTED_WRITE_COUNTERS_OFFSET_INDEX = 6, - COMBINE_COUNTED_WRITE_COUNTERS_OFFSET_INDEX = 7, - TOPK_TARGET_RANKS_OFFSET_INDEX = 8, - TOPK_SEND_INDICES_OFFSET_INDEX = 9, - EPLB_GATHERED_STATS_OFFSET_INDEX = 10, - PAYLOAD_DATA_OFFSET_INDEX = 11, - // Static max tokens/rank (a count, not a byte offset). - MAX_NUM_TOKENS_INDEX = 12, - DISPATCH_COUNTER_BASELINE_OFFSET_INDEX = 13, + // CFT-only cumulative received-byte counters per sender; uint64_t at kCftCounterStride byte spacing. + DISPATCH_COUNTED_WRITE_COUNTERS_OFFSET_INDEX = 5, + // CFT-only uint64_t[ep_size]: consumed-byte baselines for dispatch counters; ordinary device memory. + DISPATCH_COUNTER_BASELINE_OFFSET_INDEX = 6, + // int32_t[max_tokens][top_k]: destination ranks for local tokens, reused by combine to gather results. + TOPK_TARGET_RANKS_OFFSET_INDEX = 7, + // int32_t[max_tokens][top_k]: matching receive-slot indices within each destination's sender slice. + TOPK_TARGET_INDICES_OFFSET_INDEX = 8, + // int32_t[ep_size][eplb_stats_num_experts]: gathered per-rank expert statistics; empty without EPLB. + EPLB_GATHERED_STATS_OFFSET_INDEX = 9, + // Receive region for dispatched activations, scales, expert IDs/weights and optional extra payloads. + DISPATCH_PAYLOAD_OFFSET_INDEX = 10, + // Reserved dispatch payload capacity, including payload alignment padding. + DISPATCH_PAYLOAD_SIZE_INDEX = 11, + + // Combine control, expert-output staging and the CFT-only receive inbox. + // uint32_t[ep_size]: peer readiness epochs, used by fence combine and the CFT cross-round guard. + COMBINE_COMPLETION_FLAGS_OFFSET_INDEX = 12, + // CFT-only cumulative received-byte counters per expert-rank/token slot, kCftCounterStride bytes apart. + COMBINE_COUNTED_WRITE_COUNTERS_OFFSET_INDEX = 13, + // CFT-only uint64_t[ep_size][max_tokens]: consumed-byte baselines for combine counters. COMBINE_COUNTER_BASELINE_OFFSET_INDEX = 14, - NUM_METAINFO_FIELDS = 15 + // Expert-output / staging region: fence peers pull from it; MoE may write directly into it. + COMBINE_INPUT_OFFSET_INDEX = 15, + // Reserved combine input capacity, large enough for original-dtype MoE output even with FP8 wire data. + COMBINE_INPUT_SIZE_INDEX = 16, + // CFT-only receive inbox: peer pushes and the self contribution are gathered here for local reduction. + COMBINE_RECV_OFFSET_INDEX = 17, + // CFT receive payload capacity in wire-format bytes, excluding counters and baselines. + COMBINE_RECV_SIZE_INDEX = 18, + + // Allocation-time configuration used to construct and validate the layout, not mutable round state. + // Maximum input token count per rank; bounds routing storage and per-rank token slots. + MAX_NUM_TOKENS_INDEX = 19, + // Configured experts selected per token; determines routing-table capacity. + TOP_K_INDEX = 20, + // Number of ranks in the EP group; determines peer-array sizes. + EP_SIZE_INDEX = 21, + // Number of expert statistics gathered from each rank; zero disables EPLB storage. + EPLB_STATS_NUM_EXPERTS_INDEX = 22, + // Whether CFT counters, baselines and receive storage are allocated; not the current transport choice. + CFT_ENABLED_INDEX = 23, + // Total required bytes per rank, including all control buffers, payloads and alignment padding. + WORKSPACE_SIZE_INDEX = 24, + // Number of int64_t entries in the metadata tensor; not a stored layout field. + NUM_METAINFO_FIELDS = 25 }; -using MoeA2ADataOffsets = std::array; +using MoeA2AWorkspaceLayout = std::array; +static constexpr int64_t kWorkspaceAlignment = 256; inline std::vector> getMoeA2AMetaInfoIndexPairs() { + using namespace tensorrt_llm::kernels::moe_comm; return { {"MOE_A2A_FLAG_VAL_OFFSET_INDEX", FLAG_VAL_OFFSET_INDEX}, {"MOE_A2A_LOCAL_TOKEN_COUNTER_OFFSET_INDEX", LOCAL_TOKEN_COUNTER_OFFSET_INDEX}, {"MOE_A2A_SEND_COUNTERS_OFFSET_INDEX", SEND_COUNTERS_OFFSET_INDEX}, {"MOE_A2A_RECV_COUNTERS_OFFSET_INDEX", RECV_COUNTERS_OFFSET_INDEX}, {"MOE_A2A_DISPATCH_COMPLETION_FLAGS_OFFSET_INDEX", DISPATCH_COMPLETION_FLAGS_OFFSET_INDEX}, - {"MOE_A2A_COMBINE_COMPLETION_FLAGS_OFFSET_INDEX", COMBINE_COMPLETION_FLAGS_OFFSET_INDEX}, {"MOE_A2A_DISPATCH_COUNTED_WRITE_COUNTERS_OFFSET_INDEX", DISPATCH_COUNTED_WRITE_COUNTERS_OFFSET_INDEX}, - {"MOE_A2A_COMBINE_COUNTED_WRITE_COUNTERS_OFFSET_INDEX", COMBINE_COUNTED_WRITE_COUNTERS_OFFSET_INDEX}, + {"MOE_A2A_DISPATCH_COUNTER_BASELINE_OFFSET_INDEX", DISPATCH_COUNTER_BASELINE_OFFSET_INDEX}, {"MOE_A2A_TOPK_TARGET_RANKS_OFFSET_INDEX", TOPK_TARGET_RANKS_OFFSET_INDEX}, - {"MOE_A2A_TOPK_SEND_INDICES_OFFSET_INDEX", TOPK_SEND_INDICES_OFFSET_INDEX}, + {"MOE_A2A_TOPK_TARGET_INDICES_OFFSET_INDEX", TOPK_TARGET_INDICES_OFFSET_INDEX}, {"MOE_A2A_EPLB_GATHERED_STATS_OFFSET_INDEX", EPLB_GATHERED_STATS_OFFSET_INDEX}, - {"MOE_A2A_PAYLOAD_DATA_OFFSET_INDEX", PAYLOAD_DATA_OFFSET_INDEX}, - {"MOE_A2A_MAX_NUM_TOKENS_INDEX", MAX_NUM_TOKENS_INDEX}, - {"MOE_A2A_DISPATCH_COUNTER_BASELINE_OFFSET_INDEX", DISPATCH_COUNTER_BASELINE_OFFSET_INDEX}, + {"MOE_A2A_DISPATCH_PAYLOAD_OFFSET_INDEX", DISPATCH_PAYLOAD_OFFSET_INDEX}, + {"MOE_A2A_DISPATCH_PAYLOAD_SIZE_INDEX", DISPATCH_PAYLOAD_SIZE_INDEX}, + {"MOE_A2A_COMBINE_COMPLETION_FLAGS_OFFSET_INDEX", COMBINE_COMPLETION_FLAGS_OFFSET_INDEX}, + {"MOE_A2A_COMBINE_COUNTED_WRITE_COUNTERS_OFFSET_INDEX", COMBINE_COUNTED_WRITE_COUNTERS_OFFSET_INDEX}, {"MOE_A2A_COMBINE_COUNTER_BASELINE_OFFSET_INDEX", COMBINE_COUNTER_BASELINE_OFFSET_INDEX}, + {"MOE_A2A_COMBINE_INPUT_OFFSET_INDEX", COMBINE_INPUT_OFFSET_INDEX}, + {"MOE_A2A_COMBINE_INPUT_SIZE_INDEX", COMBINE_INPUT_SIZE_INDEX}, + {"MOE_A2A_COMBINE_RECV_OFFSET_INDEX", COMBINE_RECV_OFFSET_INDEX}, + {"MOE_A2A_COMBINE_RECV_SIZE_INDEX", COMBINE_RECV_SIZE_INDEX}, + {"MOE_A2A_MAX_NUM_TOKENS_INDEX", MAX_NUM_TOKENS_INDEX}, + {"MOE_A2A_TOP_K_INDEX", TOP_K_INDEX}, + {"MOE_A2A_EP_SIZE_INDEX", EP_SIZE_INDEX}, + {"MOE_A2A_EPLB_STATS_NUM_EXPERTS_INDEX", EPLB_STATS_NUM_EXPERTS_INDEX}, + {"MOE_A2A_CFT_ENABLED_INDEX", CFT_ENABLED_INDEX}, + {"MOE_A2A_WORKSPACE_SIZE_INDEX", WORKSPACE_SIZE_INDEX}, {"MOE_A2A_NUM_METAINFO_FIELDS", NUM_METAINFO_FIELDS}, + {"MOE_A2A_MAX_RANKS", kMaxRanks}, + {"MOE_A2A_MAX_TOP_K", kMaxTopK}, + {"MOE_A2A_MAX_PAYLOADS", kMaxPayloads}, + {"MOE_A2A_WORKSPACE_ALIGNMENT", kWorkspaceAlignment}, }; } diff --git a/cpp/tensorrt_llm/thop/moe/communication/moeAlltoAllOp.cpp b/cpp/tensorrt_llm/thop/moe/communication/moeAlltoAllOp.cpp index ea3d0552a61f..cdecce0bd757 100644 --- a/cpp/tensorrt_llm/thop/moe/communication/moeAlltoAllOp.cpp +++ b/cpp/tensorrt_llm/thop/moe/communication/moeAlltoAllOp.cpp @@ -26,6 +26,7 @@ #include #include #include +#include #include #include #include @@ -39,13 +40,15 @@ namespace torch_ext namespace moe_comm { -// Whether the engine is in its startup warmup phase, which uses a larger -// completion-flag budget. See moeA2AGetTimeoutCycles(). -static std::atomic gInWarmup{false}; +// Process-wide budget copied into subsequent dispatch/combine launch arguments. +static std::atomic gTimeoutCycles{tensorrt_llm::kernels::moe_comm::kDefaultTimeoutCycles}; -void moeA2ASetWarmupOp(bool in_warmup) +void moeA2ASetTimeoutOp(int64_t timeoutSec) { - gInWarmup.store(in_warmup, std::memory_order_relaxed); + constexpr int64_t kMaxTimeoutSec = 24 * 60 * 60; + TORCH_CHECK(timeoutSec > 0 && timeoutSec <= kMaxTimeoutSec, "MoE all-to-all timeout must be in 1..", kMaxTimeoutSec, + " seconds"); + gTimeoutCycles.store(timeoutSec * tensorrt_llm::kernels::moe_comm::kAssumedClockHz, std::memory_order_relaxed); } static constexpr size_t CACHELINE_ALIGNMENT = 128; @@ -114,6 +117,16 @@ inline size_t alignOffset(size_t offset, size_t alignment) return (offset + alignment - 1) & ~(alignment - 1); } +MoeA2AWorkspaceLayout const& readWorkspaceLayout(torch::Tensor const& metainfo) +{ + CHECK_CPU(metainfo); + CHECK_TYPE(metainfo, torch::kInt64); + CHECK_CONTIGUOUS(metainfo); + TORCH_CHECK(metainfo.dim() == 1 && metainfo.numel() == NUM_METAINFO_FIELDS, + "metainfo must contain the complete MoE A2A workspace layout"); + return *reinterpret_cast(metainfo.data_ptr()); +} + inline bool hasActiveRankMask(torch::optional const& maskTensor) { return maskTensor.has_value() && maskTensor.value().defined(); @@ -145,161 +158,139 @@ inline void resolveActiveRankMask(torch::optional const& maskTens ") as active"); } -// Calculate auxiliary data offsets -MoeA2ADataOffsets calculateOffsets(int epSize, int maxNumTokens, int eplbStatsNumExperts, bool canUseCft) +// All offsets and capacities are per rank. Payload capacities are padded so +// the following phase's control buffers remain aligned. +MoeA2AWorkspaceLayout calculateWorkspaceLayout(int64_t epSize, int64_t maxNumTokens, int64_t topK, + int64_t dispatchBytes, int64_t combineInputBytes, int64_t combineRecvBytes, int64_t eplbStatsNumExperts, + bool canUseCft) { - // TODO: Use lambdas to encapsulate offset and alignment for each entry, which is less error prone and easier to - // read. - constexpr size_t kSizeOfInt32 = sizeof(int32_t); - - MoeA2ADataOffsets offsets{}; - size_t offset = 0; - - // flag_val - offsets[FLAG_VAL_OFFSET_INDEX] = offset; - offset += kSizeOfInt32; - - // local_token_counter - offsets[LOCAL_TOKEN_COUNTER_OFFSET_INDEX] = offset; - offset += kSizeOfInt32; - - // send_counters - offsets[SEND_COUNTERS_OFFSET_INDEX] = offset; - offset += epSize * kSizeOfInt32; - - // recv_counters[parity][source_rank] stores the token count received from source_rank. - // The two parity banks alternate between A2A rounds. - offsets[RECV_COUNTERS_OFFSET_INDEX] = offset; - offset += 2 * epSize * kSizeOfInt32; - - // dispatch completion flags - offset = alignOffset(offset, CACHELINE_ALIGNMENT); - offsets[DISPATCH_COMPLETION_FLAGS_OFFSET_INDEX] = offset; - offset += epSize * kSizeOfInt32; - - // combine completion flags - offset = alignOffset(offset, CACHELINE_ALIGNMENT); - offsets[COMBINE_COMPLETION_FLAGS_OFFSET_INDEX] = offset; - offset += epSize * kSizeOfInt32; - - // topk_target_ranks: [maxNumTokens, kMaxTopK] - offset = alignOffset(offset, CACHELINE_ALIGNMENT); - offsets[TOPK_TARGET_RANKS_OFFSET_INDEX] = offset; - offset += static_cast(maxNumTokens) * static_cast(tensorrt_llm::kernels::moe_comm::kMaxTopK) - * kSizeOfInt32; - - // topk_send_indices: [maxNumTokens, kMaxTopK] - offset = alignOffset(offset, CACHELINE_ALIGNMENT); - offsets[TOPK_SEND_INDICES_OFFSET_INDEX] = offset; - offset += static_cast(maxNumTokens) * static_cast(tensorrt_llm::kernels::moe_comm::kMaxTopK) - * kSizeOfInt32; - - // eplb gathered stats: [epSize, eplbStatsNumExperts] - offset = alignOffset(offset, CACHELINE_ALIGNMENT); - offsets[EPLB_GATHERED_STATS_OFFSET_INDEX] = offset; - offset += static_cast(epSize) * static_cast(eplbStatsNumExperts) * kSizeOfInt32; - - // Counted write counters: each 8B counter must be kCftCounterStride-aligned, so that - // concurrent counter updates do not contend for the same L2 port. + using tensorrt_llm::kernels::moe_comm::kMaxRanks; + using tensorrt_llm::kernels::moe_comm::kMaxTopK; using tensorrt_llm::kernels::moe_comm::kCftCounterStride; - - // CFT-only regions; unused offsets stay 0 to keep the field count fixed. + TORCH_CHECK(epSize > 0 && epSize <= kMaxRanks, "Invalid EP size: ", epSize); + TORCH_CHECK(maxNumTokens > 0 && maxNumTokens <= std::numeric_limits::max(), + "Invalid allocation-time token capacity: ", maxNumTokens); + TORCH_CHECK(topK > 0 && topK <= kMaxTopK, "Invalid top_k: ", topK); + TORCH_CHECK(eplbStatsNumExperts >= 0 && eplbStatsNumExperts <= std::numeric_limits::max(), + "Invalid EPLB expert capacity: ", eplbStatsNumExperts); + TORCH_CHECK(canUseCft || combineRecvBytes == 0, "A combine receive payload requires CFT support"); + + auto alignedBytes = [](int64_t bytes) + { + TORCH_CHECK(bytes >= 0 && bytes <= std::numeric_limits::max() - kWorkspaceAlignment, + "Invalid workspace payload capacity: ", bytes); + return static_cast(alignOffset(bytes, kWorkspaceAlignment)); + }; + dispatchBytes = alignedBytes(dispatchBytes); + combineInputBytes = alignedBytes(combineInputBytes); + combineRecvBytes = alignedBytes(combineRecvBytes); + + MoeA2AWorkspaceLayout layout{}; + int64_t offset = 0; + auto reserve = [&](MoeA2AMetaInfoIndex index, int64_t bytes, int64_t alignment = sizeof(int32_t)) + { + TORCH_CHECK( + offset <= std::numeric_limits::max() - alignment, "Workspace layout alignment overflows int64"); + offset = static_cast(alignOffset(offset, alignment)); + TORCH_CHECK(bytes >= 0 && bytes <= std::numeric_limits::max() - offset, + "Workspace layout size overflows int64"); + layout[index] = offset; + offset += bytes; + }; + int64_t const rankCountsBytes = epSize * sizeof(int32_t); + int64_t const combineSlots = epSize * maxNumTokens; + + reserve(FLAG_VAL_OFFSET_INDEX, sizeof(uint32_t)); + + // Dispatch writes these counters and routes; combine reuses the completed routes. + reserve(LOCAL_TOKEN_COUNTER_OFFSET_INDEX, sizeof(int32_t)); + reserve(SEND_COUNTERS_OFFSET_INDEX, rankCountsBytes); + reserve(RECV_COUNTERS_OFFSET_INDEX, 2 * rankCountsBytes); + reserve(DISPATCH_COMPLETION_FLAGS_OFFSET_INDEX, rankCountsBytes, CACHELINE_ALIGNMENT); if (canUseCft) { - // dispatch counted write counters: [ep_size] uint64_t, kCftCounterStride stride - offset = alignOffset(offset, kCftCounterStride); - offsets[DISPATCH_COUNTED_WRITE_COUNTERS_OFFSET_INDEX] = offset; - offset += epSize * kCftCounterStride; - - // combine counted write counters (CFT combine path): per receive-slot uint64 - // counters, [ep_size * maxNumTokens], kCftCounterStride stride to avoid L2 XBAR camping. - offset = alignOffset(offset, kCftCounterStride); - offsets[COMBINE_COUNTED_WRITE_COUNTERS_OFFSET_INDEX] = offset; - offset += static_cast(epSize) * static_cast(maxNumTokens) * kCftCounterStride; - - // dispatch counter baseline: [ep_size] uint64 - offset = alignOffset(offset, CACHELINE_ALIGNMENT); - offsets[DISPATCH_COUNTER_BASELINE_OFFSET_INDEX] = offset; - offset += static_cast(epSize) * sizeof(uint64_t); - - // combine counter baseline: [ep_size * maxNumTokens] uint64 - offset = alignOffset(offset, CACHELINE_ALIGNMENT); - offsets[COMBINE_COUNTER_BASELINE_OFFSET_INDEX] = offset; - offset += static_cast(epSize) * static_cast(maxNumTokens) * sizeof(uint64_t); + reserve(DISPATCH_COUNTED_WRITE_COUNTERS_OFFSET_INDEX, epSize * kCftCounterStride, kCftCounterStride); + reserve(DISPATCH_COUNTER_BASELINE_OFFSET_INDEX, epSize * sizeof(uint64_t), CACHELINE_ALIGNMENT); + } + reserve(TOPK_TARGET_RANKS_OFFSET_INDEX, maxNumTokens * topK * sizeof(int32_t), CACHELINE_ALIGNMENT); + reserve(TOPK_TARGET_INDICES_OFFSET_INDEX, maxNumTokens * topK * sizeof(int32_t), CACHELINE_ALIGNMENT); + reserve(EPLB_GATHERED_STATS_OFFSET_INDEX, epSize * eplbStatsNumExperts * sizeof(int32_t), CACHELINE_ALIGNMENT); + reserve(DISPATCH_PAYLOAD_OFFSET_INDEX, dispatchBytes, kWorkspaceAlignment); + layout[DISPATCH_PAYLOAD_SIZE_INDEX] = dispatchBytes; + + // Fence pulls from combine input; CFT pushes into the separate receive inbox. + reserve(COMBINE_COMPLETION_FLAGS_OFFSET_INDEX, rankCountsBytes, CACHELINE_ALIGNMENT); + if (canUseCft) + { + reserve(COMBINE_COUNTED_WRITE_COUNTERS_OFFSET_INDEX, combineSlots * kCftCounterStride, kCftCounterStride); + reserve(COMBINE_COUNTER_BASELINE_OFFSET_INDEX, combineSlots * sizeof(uint64_t), CACHELINE_ALIGNMENT); + } + reserve(COMBINE_INPUT_OFFSET_INDEX, combineInputBytes, kWorkspaceAlignment); + layout[COMBINE_INPUT_SIZE_INDEX] = combineInputBytes; + if (canUseCft) + { + reserve(COMBINE_RECV_OFFSET_INDEX, combineRecvBytes, kWorkspaceAlignment); + layout[COMBINE_RECV_SIZE_INDEX] = combineRecvBytes; } - // payload data - offset = alignOffset(offset, CACHELINE_ALIGNMENT); - offsets[PAYLOAD_DATA_OFFSET_INDEX] = offset; - - // Stable combine slot stride (a count, not a byte offset). - offsets[MAX_NUM_TOKENS_INDEX] = maxNumTokens; - - return offsets; + layout[MAX_NUM_TOKENS_INDEX] = maxNumTokens; + layout[TOP_K_INDEX] = topK; + layout[EP_SIZE_INDEX] = epSize; + layout[EPLB_STATS_NUM_EXPERTS_INDEX] = eplbStatsNumExperts; + layout[CFT_ENABLED_INDEX] = canUseCft; + layout[WORKSPACE_SIZE_INDEX] = offset; + return layout; } -// Initialize auxiliary data in workspace -// This function sets up the initial values for flag_val and completion_flags -// -// Inputs: -// - workspace: [ep_size, size_per_rank] unified virtual memory workspace -// - epRank: Current expert parallel rank -// - epSize: Total expert parallel size -// - maxNumTokens: Maximum number of tokens supported -// - eplbStatsNumExperts: (Optional) Number of experts used for EPLB stats -// -// Returns: -// - metainfo: Tensor containing offsets for auxiliary data -torch::Tensor moeA2AInitializeOp(torch::Tensor const& workspace, int64_t epRank, int64_t epSize, int64_t maxNumTokens, - torch::optional eplbStatsNumExperts, bool canUseCftCountedWrites) +torch::Tensor moeA2AGetWorkspaceLayoutOp(int64_t epSize, int64_t maxNumTokens, int64_t topK, int64_t dispatchBytes, + int64_t combineInputBytes, int64_t combineRecvBytes, torch::optional eplbStatsNumExperts, bool canUseCft) { - using tensorrt_llm::kernels::moe_comm::kMaxRanks; + auto const layout = calculateWorkspaceLayout(epSize, maxNumTokens, topK, dispatchBytes, combineInputBytes, + combineRecvBytes, eplbStatsNumExperts.value_or(0), canUseCft); + auto metainfo + = torch::empty({NUM_METAINFO_FIELDS}, torch::TensorOptions().dtype(torch::kInt64).device(torch::kCPU)); + std::copy(layout.begin(), layout.end(), metainfo.data_ptr()); + return metainfo; +} - // Validate inputs +// Initialize control state after allocation; layout construction itself is CPU-only. +void moeA2AInitializeOp(torch::Tensor const& workspace, torch::Tensor const& metainfo, int64_t epRank, int64_t epSize) +{ CHECK_TH_CUDA(workspace); CHECK_TYPE(workspace, torch::kUInt8); - TORCH_CHECK(workspace.dim() == 2, "workspace must be a 2D tensor of shape [epSize, sizePerRank]"); - TORCH_CHECK(workspace.size(0) == epSize, "workspace first dimension must equal epSize"); - TORCH_CHECK(epSize > 0 && epSize <= kMaxRanks, "epSize must be in the range (0, ", kMaxRanks, "]"); - TORCH_CHECK(epRank >= 0 && epRank < epSize, "epRank must be in the range [0, epSize)"); - - int64_t eplbStatsNumExpertsValue = eplbStatsNumExperts.value_or(0); - TORCH_CHECK(eplbStatsNumExpertsValue >= 0, "eplbStatsNumExperts must be positive if not None."); - - // Calculate auxiliary data offsets - MoeA2ADataOffsets offsets - = calculateOffsets(epSize, maxNumTokens, static_cast(eplbStatsNumExpertsValue), canUseCftCountedWrites); - - // Initialize workspace to zero, then mark both recv-counter parities empty. - workspace[epRank].zero_(); - uint8_t* rankWorkSpacePtr = workspace.data_ptr() + epRank * workspace.stride(0); - cudaMemsetAsync(rankWorkSpacePtr + offsets[RECV_COUNTERS_OFFSET_INDEX], 0xFF, + TORCH_CHECK(workspace.dim() == 2 && workspace.size(0) == epSize && workspace.stride(1) == 1, + "workspace must have shape [ep_size, bytes_per_rank] with contiguous rank slices"); + TORCH_CHECK(epRank >= 0 && epRank < epSize, "Invalid EP rank"); + auto const& layout = readWorkspaceLayout(metainfo); + TORCH_CHECK(layout[EP_SIZE_INDEX] == epSize, "Workspace layout EP size mismatch"); + auto const expected = calculateWorkspaceLayout(epSize, layout[MAX_NUM_TOKENS_INDEX], layout[TOP_K_INDEX], + layout[DISPATCH_PAYLOAD_SIZE_INDEX], layout[COMBINE_INPUT_SIZE_INDEX], layout[COMBINE_RECV_SIZE_INDEX], + layout[EPLB_STATS_NUM_EXPERTS_INDEX], layout[CFT_ENABLED_INDEX]); + TORCH_CHECK(layout == expected, "Workspace layout metadata is inconsistent"); + TORCH_CHECK(layout[DISPATCH_PAYLOAD_SIZE_INDEX] > 0 && layout[COMBINE_INPUT_SIZE_INDEX] > 0, + "Dispatch and combine input payload capacities must be positive"); + TORCH_CHECK(!layout[CFT_ENABLED_INDEX] || layout[COMBINE_RECV_SIZE_INDEX] > 0, + "CFT requires a combine receive payload capacity"); + TORCH_CHECK(workspace.size(1) >= layout[WORKSPACE_SIZE_INDEX], "Workspace needs ", layout[WORKSPACE_SIZE_INDEX], + " bytes per rank, got ", workspace.size(1)); + + workspace[epRank].narrow(0, 0, layout[WORKSPACE_SIZE_INDEX]).zero_(); + uint8_t* rankWorkspace = workspace.data_ptr() + epRank * workspace.stride(0); + cudaMemsetAsync(rankWorkspace + layout[RECV_COUNTERS_OFFSET_INDEX], 0xFF, 2 * static_cast(epSize) * sizeof(int32_t), at::cuda::getCurrentCUDAStream()); - - // Return metainfo as a tensor containing offsets - torch::Tensor metainfo = torch::empty( - {static_cast(NUM_METAINFO_FIELDS)}, torch::TensorOptions().dtype(torch::kInt64).device(torch::kCPU)); - - for (int i = 0; i < static_cast(NUM_METAINFO_FIELDS); i++) - { - metainfo[i] = static_cast(offsets[i]); - } - // Synchronize among ranks. Under a non-MPI orchestrator (Ray) MpiComm throws // "MPI is disabled, DON'T USE MPI" from mpiUtils.h, which made the whole // NVLinkOneSided strategy unusable there; fall back to the Torch process // group the Ray workers already initialise, as pg_utils does elsewhere. cudaDeviceSynchronize(); moeA2ABarrier(); - - return metainfo; } // ============================================================================ // CFT Handle-Based Counted Writes Initialization // ============================================================================ -// Static CftLeManager — lives for the process lifetime (like workspace). +// One CFT binding per process, released before its backing workspace is freed. static std::unique_ptr g_cft_manager; // Initialize CFT Logical Endpoints by binding the LE to the MNNVL workspace. @@ -396,6 +387,25 @@ void moeA2ACftInitializeOp(torch::Tensor const& workspace, int64_t workspaceMemH } } +// All ranks must finish using the workspace before releasing their local binding. +void moeA2ACftDestroyOp(torch::Tensor const& workspace, int64_t epRank) +{ + CHECK_TH_CUDA(workspace); + CHECK_TYPE(workspace, torch::kUInt8); + TORCH_CHECK(workspace.dim() == 2, "workspace must be a 2D tensor"); + TORCH_CHECK(epRank >= 0 && epRank < workspace.size(0), "epRank is outside the workspace"); + if (!g_cft_manager || !g_cft_manager->isInitialized()) + { + return; + } + auto const workspaceRankPtr + = reinterpret_cast(workspace.data_ptr() + epRank * workspace.stride(0)); + TORCH_CHECK(g_cft_manager->getLocalBackingPtr() == workspaceRankPtr, + "Cannot destroy CFT endpoints bound to a different workspace"); + TORCH_CHECK(cudaDeviceSynchronize() == cudaSuccess, "CUDA synchronization failed before CFT endpoint release"); + g_cft_manager.reset(); +} + // MoE All-to-All Dispatch Operation // This operation dispatches tokens and their associated payloads to different expert ranks. // @@ -444,15 +454,12 @@ std::tuple, int64_t, torch::Tensor> moeA2ADispatchOp( TORCH_CHECK(tokenSelectedExperts.dim() == 2, "tokenSelectedExperts must be a 2D tensor"); TORCH_CHECK(tokenSelectedExperts.size(1) == topK, "tokenSelectedExperts must have topK columns"); - CHECK_CPU(metainfo); - CHECK_TYPE(metainfo, torch::kInt64); - TORCH_CHECK(metainfo.dim() == 1, "metainfo must be a 1D tensor"); - TORCH_CHECK(metainfo.size(0) == static_cast(NUM_METAINFO_FIELDS), - "metainfo must have NUM_METAINFO_FIELDS elements"); - MoeA2ADataOffsets const& offsets = *reinterpret_cast(metainfo.data_ptr()); + auto const& offsets = readWorkspaceLayout(metainfo); int64_t localNumTokens = tokenSelectedExperts.size(0); TORCH_CHECK(runtimeMaxTokensPerRank > 0, "runtimeMaxTokensPerRank must be positive"); + TORCH_CHECK(runtimeMaxTokensPerRank <= offsets[MAX_NUM_TOKENS_INDEX], + "runtimeMaxTokensPerRank exceeds the allocation-time token capacity"); TORCH_CHECK(epSize > 0 && epSize <= kMaxRanks, "epSize must be in the range (0, ", kMaxRanks, "]"); TORCH_CHECK(epRank >= 0 && epRank < epSize, "epRank must be in the range [0, epSize)"); TORCH_CHECK(topK > 0 && topK <= kMaxTopK, "topK must be in the range (0, kMaxTopK]"); @@ -491,7 +498,7 @@ std::tuple, int64_t, torch::Tensor> moeA2ADispatchOp( // Record the cacheline aligned start offset for each payload's recv buffer. // 1. We assume the base workspace ptr of each rank is aligned (checked in this OP) - // 2. offsets[PAYLOAD_DATA_OFFSET_INDEX] is aligned (ensured in calculateOffsets) + // 2. offsets[DISPATCH_PAYLOAD_OFFSET_INDEX] is aligned (fixed by the workspace layout) // 3. We align the currentOffset during update. // In this way, it is guaranteed that the recv buffer is (over-)aligned, sufficient for 128bit vectorized ld/st. @@ -500,7 +507,7 @@ std::tuple, int64_t, torch::Tensor> moeA2ADispatchOp( std::vector payloadRecvBufferOffsets; // Start offset for the first payload - size_t currentOffset = static_cast(offsets[PAYLOAD_DATA_OFFSET_INDEX]); + size_t currentOffset = static_cast(offsets[DISPATCH_PAYLOAD_OFFSET_INDEX]); for (auto const& payload : inputPayloads) { CHECK_CONTIGUOUS(payload); @@ -547,14 +554,18 @@ std::tuple, int64_t, torch::Tensor> moeA2ADispatchOp( TORCH_CHECK(workspace.dim() == 2, "workspace must be a 2D tensor of shape [epSize, sizePerRank]"); TORCH_CHECK(workspace.size(0) == epSize, "workspace first dimension must equal epSize"); - // Validate workspace size - must include space for auxiliary data + payloads - int64_t sizePerRank = workspace.size(1); - int64_t requiredSize = static_cast(currentOffset); - TORCH_CHECK(sizePerRank >= requiredSize, - "Workspace size per rank insufficient for dispatch. " - "Need at least ", - requiredSize, " bytes (", offsets[PAYLOAD_DATA_OFFSET_INDEX], " for auxiliary data + payloads), but got ", - sizePerRank); + TORCH_CHECK(epSize == offsets[EP_SIZE_INDEX] && topK == offsets[TOP_K_INDEX], + "Dispatch EP size/top_k differs from its workspace layout"); + TORCH_CHECK(localNumTokens <= runtimeMaxTokensPerRank, "Local token count exceeds the runtime capacity"); + TORCH_CHECK(eplbStatsNumExperts <= offsets[EPLB_STATS_NUM_EXPERTS_INDEX], + "EPLB statistics exceed their workspace capacity"); + TORCH_CHECK(!useCftCountedWrites || offsets[CFT_ENABLED_INDEX], "Workspace was allocated without CFT support"); + TORCH_CHECK(workspace.size(1) >= offsets[WORKSPACE_SIZE_INDEX], "Workspace is smaller than its layout"); + int64_t const combinePayloadOffset = offsets[COMBINE_INPUT_OFFSET_INDEX]; + int64_t const payloadCapacity = offsets[DISPATCH_PAYLOAD_SIZE_INDEX]; + TORCH_CHECK(currentOffset <= static_cast(offsets[DISPATCH_PAYLOAD_OFFSET_INDEX] + payloadCapacity), + "Dispatch payload exceeds its workspace capacity: need ", + currentOffset - offsets[DISPATCH_PAYLOAD_OFFSET_INDEX], " bytes, capacity ", payloadCapacity); // Get base workspace pointer uint8_t* workspacePtr = workspace.data_ptr(); @@ -592,7 +603,7 @@ std::tuple, int64_t, torch::Tensor> moeA2ADispatchOp( params.local_token_counter = reinterpret_cast(rankWorkSpacePtr + offsets[LOCAL_TOKEN_COUNTER_OFFSET_INDEX]); params.send_counters = reinterpret_cast(rankWorkSpacePtr + offsets[SEND_COUNTERS_OFFSET_INDEX]); params.topk_target_ranks = reinterpret_cast(rankWorkSpacePtr + offsets[TOPK_TARGET_RANKS_OFFSET_INDEX]); - params.topk_send_indices = reinterpret_cast(rankWorkSpacePtr + offsets[TOPK_SEND_INDICES_OFFSET_INDEX]); + params.topk_target_indices = reinterpret_cast(rankWorkSpacePtr + offsets[TOPK_TARGET_INDICES_OFFSET_INDEX]); for (int target_rank = 0; target_rank < epSize; target_rank++) { @@ -688,8 +699,7 @@ std::tuple, int64_t, torch::Tensor> moeA2ADispatchOp( } params.stream = at::cuda::getCurrentCUDAStream(); - params.timeout_cycles - = tensorrt_llm::kernels::moe_comm::moeA2AGetTimeoutCycles(gInWarmup.load(std::memory_order_relaxed)); + params.timeout_cycles = gTimeoutCycles.load(std::memory_order_relaxed); // Prepare for dispatch (zero counters/indices and increment flag_val) moe_a2a_prepare_dispatch_launch(params); @@ -705,23 +715,12 @@ std::tuple, int64_t, torch::Tensor> moeA2ADispatchOp( for (int payload_idx = 0; payload_idx < num_payloads; payload_idx++) { auto const& payload = inputPayloads[payload_idx]; - void* recvDataPtr; - if (useCftCountedWrites) - { - // LE IS workspace — recv data is at the same workspace offset regardless of CFT. - recvDataPtr = rankWorkSpacePtr + payloadRecvBufferOffsets[payload_idx]; - } - else - { - recvDataPtr = rankWorkSpacePtr + payloadRecvBufferOffsets[payload_idx]; - } + void* recvDataPtr = rankWorkSpacePtr + payloadRecvBufferOffsets[payload_idx]; auto recvTensor = torch::from_blob( recvDataPtr, {epSize, runtimeMaxTokensPerRank, payloadElementsPerToken[payload_idx]}, payload.options()); recvTensors.push_back(recvTensor); } - // Compute aligned offset after dispatch payloads for combine payload region - int64_t combinePayloadOffset = static_cast(alignOffset(currentOffset, CACHELINE_ALIGNMENT)); torch::Tensor eplbGatheredStats; if (enableEplb) { @@ -741,9 +740,9 @@ std::tuple, int64_t, torch::Tensor> moeA2ADispatchOp( // MoE All-to-All Combine Operation // Combine the per-rank expert outputs into the originating tokens' buffers on the local rank. // -// The payload may be external or a view of the normal combine workspace region. Callers that place -// the MoE output directly in the workspace pass payloadInWorkspace=true to skip staging; callers -// that cannot choose the MoE output tensor leave it false and prepareCombine stages the payload. +// The payload may be external or a view of the combine input region. Recognize the input +// region by its address; payloadInWorkspace=true additionally requires that zero-copy path. +// Other sources, including dispatch payload views, are staged when needed. // Fence combine reads from 'combinePayloadOffset'. CFT combine stages the local slice and receives // peer slices in a dedicated counted-write region before reduction. torch::Tensor moeA2ACombineOp(torch::Tensor const& payload, int64_t localNumTokens, torch::Tensor const& workspace, @@ -796,12 +795,7 @@ torch::Tensor moeA2ACombineOp(torch::Tensor const& payload, int64_t localNumToke } // use_low_precision is passed through to the kernel via params.use_low_precision; dtype is not mutated. - CHECK_CPU(metainfo); - CHECK_TYPE(metainfo, torch::kInt64); - TORCH_CHECK(metainfo.dim() == 1, "metainfo must be a 1D tensor"); - TORCH_CHECK(metainfo.size(0) == static_cast(NUM_METAINFO_FIELDS), - "metainfo must have NUM_METAINFO_FIELDS elements"); - MoeA2ADataOffsets const& offsets = *reinterpret_cast(metainfo.data_ptr()); + auto const& offsets = readWorkspaceLayout(metainfo); // Validate workspace and set synchronization pointers CHECK_TH_CUDA(workspace); @@ -810,22 +804,24 @@ torch::Tensor moeA2ACombineOp(torch::Tensor const& payload, int64_t localNumToke uint8_t* workspacePtr = workspace.data_ptr(); int64_t sizePerRank = workspace.size(1); uint8_t* rankWorkSpacePtr = workspacePtr + epRank * workspace.stride(0); - TORCH_CHECK(combinePayloadOffset >= 0, "combinePayloadOffset must be non-negative"); + TORCH_CHECK(epSize == offsets[EP_SIZE_INDEX] && topK == offsets[TOP_K_INDEX], + "Combine EP size/top_k differs from its workspace layout"); + TORCH_CHECK(sizePerRank >= offsets[WORKSPACE_SIZE_INDEX], "Workspace is smaller than its layout"); + TORCH_CHECK(!useCftCountedWrites || offsets[CFT_ENABLED_INDEX], "Workspace was allocated without CFT support"); + int64_t const regionSize = offsets[COMBINE_INPUT_SIZE_INDEX]; + TORCH_CHECK(combinePayloadOffset == offsets[COMBINE_INPUT_OFFSET_INDEX], + "combinePayloadOffset must address the fixed combine source region"); + TORCH_CHECK(runtimeMaxTokensPerRank <= offsets[MAX_NUM_TOKENS_INDEX], + "runtimeMaxTokensPerRank exceeds the allocation-time token capacity"); uint8_t* combinePayloadPtr = rankWorkSpacePtr + combinePayloadOffset; // If the caller claims the payload is in the workspace, ensure it really is: a mismatch would // otherwise silently fall back to staging and lose the zero-copy path the caller asked for. - if (payloadInWorkspace) - { - TORCH_CHECK(payload.data_ptr() == combinePayloadPtr, - "payload_in_workspace is true but 'payload' dataptr does not match combinePayloadOffset"); - } - + bool const inputIsWorkspace = payload.data_ptr() == combinePayloadPtr; + TORCH_CHECK(!payloadInWorkspace || inputIsWorkspace, + "payload_in_workspace is true but payload does not address the combine input region"); int64_t payloadSize = payload.numel() * payload.element_size(); - TORCH_CHECK(combinePayloadOffset + payloadSize <= sizePerRank, - "Workspace size per rank insufficient for combine. " - "Need at least ", - combinePayloadOffset + payloadSize, " bytes (", combinePayloadOffset, " for offset + ", payloadSize, - " for payload), but got ", sizePerRank); + TORCH_CHECK(payloadSize <= regionSize, "Combine payload exceeds its fixed workspace region: need ", payloadSize, + " bytes, capacity ", regionSize); // Create output tensor (local on current rank), no need for initialization // Typically, newly allocated GPU torch tensors are at least 16-byte aligned. @@ -849,11 +845,11 @@ torch::Tensor moeA2ACombineOp(torch::Tensor const& payload, int64_t localNumToke params.wire_bytes_per_token = static_cast(elementsPerToken) * (useLowPrecision ? 1 : static_cast(payload.element_size())); params.workspace_stride_per_token - = useLowPrecision && !payloadInWorkspace ? params.wire_bytes_per_token : params.source_stride_per_token; + = useLowPrecision && !inputIsWorkspace ? params.wire_bytes_per_token : params.source_stride_per_token; params.flag_val = reinterpret_cast(rankWorkSpacePtr + offsets[FLAG_VAL_OFFSET_INDEX]); params.topk_target_ranks = reinterpret_cast(rankWorkSpacePtr + offsets[TOPK_TARGET_RANKS_OFFSET_INDEX]); - params.topk_send_indices = reinterpret_cast(rankWorkSpacePtr + offsets[TOPK_SEND_INDICES_OFFSET_INDEX]); + params.topk_target_indices = reinterpret_cast(rankWorkSpacePtr + offsets[TOPK_TARGET_INDICES_OFFSET_INDEX]); params.recv_counters = reinterpret_cast(rankWorkSpacePtr + offsets[RECV_COUNTERS_OFFSET_INDEX]); for (int target_rank = 0; target_rank < epSize; target_rank++) @@ -861,7 +857,7 @@ torch::Tensor moeA2ACombineOp(torch::Tensor const& payload, int64_t localNumToke uint8_t* target_workspace_ptr = workspacePtr + target_rank * workspace.stride(0); params.completion_flags[target_rank] = reinterpret_cast(target_workspace_ptr + offsets[COMBINE_COMPLETION_FLAGS_OFFSET_INDEX]); - params.recv_buffers[target_rank] = target_workspace_ptr + combinePayloadOffset; + params.combine_input_buffers[target_rank] = target_workspace_ptr + combinePayloadOffset; } // CFT requires the payload to be 16B-aligned (fabric.try_put.counted operates on 16B chunks). @@ -871,8 +867,8 @@ torch::Tensor moeA2ACombineOp(torch::Tensor const& payload, int64_t localNumToke " bytes per token; CFT counted writes require 16-byte alignment"); } - // ---- CFT combine wiring (counted writes). Sets up dedicated receive region C, - // per-slot combine counters, and single-buffer baselines. Fence combine ignores these. ---- + // CFT receives peer pushes and the local contribution into a dedicated local inbox. + // Fence combine instead reads the peer combine input buffers directly. params.use_cft_for_combine = useCftCountedWrites; if (useCftCountedWrites) { @@ -885,15 +881,16 @@ torch::Tensor moeA2ACombineOp(torch::Tensor const& payload, int64_t localNumToke } // Dedicated combine receive region: prepare writes the local slice and fabric pushes write peer slices. - int64_t combineRecvRegionOffset = alignOffset(combinePayloadOffset + payloadSize, CACHELINE_ALIGNMENT); - TORCH_CHECK(combineRecvRegionOffset + payloadSize <= sizePerRank, - "CFT combine: workspace too small for combine receive region C: need ", - combineRecvRegionOffset + payloadSize, " bytes, got ", sizePerRank); - params.cft_le_combine_payload_base = static_cast(combineRecvRegionOffset); + int64_t const combineRecvRegionOffset = offsets[COMBINE_RECV_OFFSET_INDEX]; + int64_t const receiveBytes = epSize * runtimeMaxTokensPerRank * params.wire_bytes_per_token; + TORCH_CHECK(receiveBytes <= offsets[COMBINE_RECV_SIZE_INDEX], + "CFT combine receive payload exceeds its workspace capacity: need ", receiveBytes, " bytes, capacity ", + offsets[COMBINE_RECV_SIZE_INDEX]); + params.cft_combine_recv_offset = static_cast(combineRecvRegionOffset); params.cft_le_combine_counter_base = offsets[COMBINE_COUNTED_WRITE_COUNTERS_OFFSET_INDEX]; params.cft_le_combine_counters = reinterpret_cast(rankWorkSpacePtr + offsets[COMBINE_COUNTED_WRITE_COUNTERS_OFFSET_INDEX]); - params.cft_le_combine_recv = reinterpret_cast(rankWorkSpacePtr + combineRecvRegionOffset); + params.cft_combine_recv_payload = rankWorkSpacePtr + combineRecvRegionOffset; { int const staticMaxTokens = static_cast(offsets[MAX_NUM_TOKENS_INDEX]); @@ -932,7 +929,7 @@ torch::Tensor moeA2ACombineOp(torch::Tensor const& payload, int64_t localNumToke } else { - params.prepare_num_tokens = payloadInWorkspace ? 0 : params.ep_size * params.max_tokens_per_rank; + params.prepare_num_tokens = inputIsWorkspace ? 0 : params.ep_size * params.max_tokens_per_rank; } params.cft_push_payload = params.use_low_precision ? combinePayloadPtr : params.source_payload; @@ -942,8 +939,7 @@ torch::Tensor moeA2ACombineOp(torch::Tensor const& payload, int64_t localNumToke = params.use_cft_for_combine ? params.wire_bytes_per_token : params.workspace_stride_per_token; params.stream = at::cuda::getCurrentCUDAStream(); - params.timeout_cycles - = tensorrt_llm::kernels::moe_comm::moeA2AGetTimeoutCycles(gInWarmup.load(std::memory_order_relaxed)); + params.timeout_cycles = gTimeoutCycles.load(std::memory_order_relaxed); moe_a2a_prepare_combine_launch(params); @@ -972,12 +968,7 @@ void moeA2ASanitizeExpertIdsOp(torch::Tensor& expert_ids, torch::Tensor& workspa int runtime_max_tokens_per_rank = static_cast(expert_ids.size(1)); int top_k = static_cast(expert_ids.size(2)); - CHECK_CPU(metainfo); - CHECK_TYPE(metainfo, torch::kInt64); - TORCH_CHECK(metainfo.dim() == 1, "metainfo must be a 1D tensor"); - TORCH_CHECK(metainfo.size(0) == static_cast(NUM_METAINFO_FIELDS), - "metainfo must have NUM_METAINFO_FIELDS elements"); - MoeA2ADataOffsets const& offsets = *reinterpret_cast(metainfo.data_ptr()); + auto const& offsets = readWorkspaceLayout(metainfo); uint8_t* rankWorkSpacePtr = workspace.data_ptr() + epRank * workspace.stride(0); int* recv_counters = reinterpret_cast(rankWorkSpacePtr + offsets[RECV_COUNTERS_OFFSET_INDEX]); @@ -989,8 +980,8 @@ void moeA2ASanitizeExpertIdsOp(torch::Tensor& expert_ids, torch::Tensor& workspa } // Return a workspace-backed tensor for combine payload region using from_blob -torch::Tensor moeA2AGetCombinePayloadTensorOp(torch::Tensor const& workspace, int64_t epRank, int64_t epSize, - int64_t runtimeMaxTokensPerRank, int64_t combinePayloadOffset, c10::ScalarType outDtype, int64_t hiddenSize) +torch::Tensor moeA2AGetCombinePayloadTensorOp(torch::Tensor const& workspace, torch::Tensor const& metainfo, + int64_t epRank, int64_t epSize, int64_t runtimeMaxTokensPerRank, c10::ScalarType outDtype, int64_t hiddenSize) { CHECK_TH_CUDA(workspace); CHECK_TYPE(workspace, torch::kUInt8); @@ -1003,10 +994,14 @@ torch::Tensor moeA2AGetCombinePayloadTensorOp(torch::Tensor const& workspace, in int64_t sizePerRank = workspace.size(1); // bytes int64_t elementSize = static_cast(c10::elementSize(outDtype)); int64_t bytesNeeded = epSize * runtimeMaxTokensPerRank * hiddenSize * elementSize; - TORCH_CHECK(combinePayloadOffset >= 0, "combine_payload_offset must be non-negative"); - TORCH_CHECK(combinePayloadOffset + bytesNeeded <= sizePerRank, - "workspace does not have enough space for combine payload tensor. combine payload offset=", - combinePayloadOffset, ", payload size needed=", bytesNeeded, ", workspace size per rank=", sizePerRank); + auto const& layout = readWorkspaceLayout(metainfo); + TORCH_CHECK(epSize == layout[EP_SIZE_INDEX], "Combine view EP size differs from its workspace layout"); + TORCH_CHECK(runtimeMaxTokensPerRank <= layout[MAX_NUM_TOKENS_INDEX], + "Combine view exceeds the allocation-time token capacity"); + TORCH_CHECK(sizePerRank >= layout[WORKSPACE_SIZE_INDEX], "Workspace is smaller than its layout"); + int64_t const combinePayloadOffset = layout[COMBINE_INPUT_OFFSET_INDEX]; + TORCH_CHECK(bytesNeeded <= layout[COMBINE_INPUT_SIZE_INDEX], "Combine view exceeds its input region: need ", + bytesNeeded, " bytes, capacity ", layout[COMBINE_INPUT_SIZE_INDEX]); uint8_t* base = workspace.data_ptr(); uint8_t* rankBase = base + epRank * workspace.stride(0); @@ -1017,17 +1012,6 @@ torch::Tensor moeA2AGetCombinePayloadTensorOp(torch::Tensor const& workspace, in return t; } -// Return the size of auxiliary data in workspace -int64_t moeA2AGetAuxDataSizeOp( - int64_t epSize, int64_t maxNumTokens, torch::optional eplbStatsNumExperts, bool canUseCftCountedWrites) -{ - int64_t eplbStatsNumExpertsValue = eplbStatsNumExperts.value_or(0); - TORCH_CHECK(eplbStatsNumExpertsValue >= 0, "eplbStatsNumExperts must be positive if not None."); - MoeA2ADataOffsets offsets = calculateOffsets(static_cast(epSize), static_cast(maxNumTokens), - static_cast(eplbStatsNumExpertsValue), canUseCftCountedWrites); - return static_cast(offsets[PAYLOAD_DATA_OFFSET_INDEX]); -} - } // namespace moe_comm } // namespace torch_ext @@ -1062,21 +1046,20 @@ TORCH_LIBRARY_FRAGMENT(trtllm, module) module.def( "moe_a2a_cft_initialize(Tensor(a!) workspace, int workspace_mem_handle, " "int workspace_size_per_rank, int ep_rank, int ep_size) -> ()"); - module.def( - "moe_a2a_initialize(Tensor(a!) workspace, int ep_rank, int ep_size, int max_num_tokens_per_rank, " - "int? eplb_stats_num_experts=None, bool can_use_cft_counted_writes=False) -> Tensor"); + module.def("moe_a2a_cft_destroy(Tensor(a!) workspace, int ep_rank) -> ()"); + module.def("moe_a2a_initialize(Tensor(a!) workspace, Tensor metainfo, int ep_rank, int ep_size) -> ()"); module.def( "moe_a2a_sanitize_expert_ids(Tensor(a!) expert_ids, Tensor(a!) workspace, Tensor metainfo, int ep_rank, int " "invalid_expert_id) -> ()"); module.def( - "moe_a2a_get_combine_payload_tensor(Tensor(a) workspace, int ep_rank, int ep_size, int " - "runtime_max_tokens_per_rank, " - "int combine_payload_offset, ScalarType out_dtype, int hidden_size) -> Tensor(a)"); - module.def("moe_a2a_set_warmup(bool in_warmup) -> ()", &tensorrt_llm::torch_ext::moe_comm::moeA2ASetWarmupOp); + "moe_a2a_get_combine_payload_tensor(Tensor(a) workspace, Tensor metainfo, int ep_rank, int ep_size, " + "int runtime_max_tokens_per_rank, ScalarType out_dtype, int hidden_size) -> Tensor(a)"); + module.def("moe_a2a_set_timeout(int timeout_sec) -> ()", &tensorrt_llm::torch_ext::moe_comm::moeA2ASetTimeoutOp); module.def( - "moe_a2a_get_aux_data_size(int ep_size, int max_num_tokens, int? eplb_stats_num_experts=None, " - "bool can_use_cft_counted_writes=False) -> int", - &tensorrt_llm::torch_ext::moe_comm::moeA2AGetAuxDataSizeOp); + "moe_a2a_get_workspace_layout(int ep_size, int max_num_tokens_per_rank, int top_k, " + "int dispatch_payload_bytes, int combine_input_bytes, int combine_recv_bytes, " + "int? eplb_stats_num_experts=None, bool can_use_cft_counted_writes=False) -> Tensor", + &tensorrt_llm::torch_ext::moe_comm::moeA2AGetWorkspaceLayoutOp); } TORCH_LIBRARY_IMPL(trtllm, CUDA, module) @@ -1088,4 +1071,5 @@ TORCH_LIBRARY_IMPL(trtllm, CUDA, module) module.impl( "moe_a2a_get_combine_payload_tensor", &tensorrt_llm::torch_ext::moe_comm::moeA2AGetCombinePayloadTensorOp); module.impl("moe_a2a_cft_initialize", &tensorrt_llm::torch_ext::moe_comm::moeA2ACftInitializeOp); + module.impl("moe_a2a_cft_destroy", &tensorrt_llm::torch_ext::moe_comm::moeA2ACftDestroyOp); } diff --git a/legacy-files.txt b/legacy-files.txt index 712874114b32..87c84be1c835 100644 --- a/legacy-files.txt +++ b/legacy-files.txt @@ -160,7 +160,6 @@ tensorrt_llm/_torch/device_mesh.py tensorrt_llm/_torch/disaggregation/kv_cache_transceiver.py tensorrt_llm/_torch/distributed/__init__.py tensorrt_llm/_torch/distributed/communicator.py -tensorrt_llm/_torch/moe/fused_moe/communication/moe_alltoall.py tensorrt_llm/_torch/distributed/ops.py tensorrt_llm/_torch/distributed/pg_utils.py tensorrt_llm/_torch/moe/expert_statistic.py @@ -586,7 +585,6 @@ tests/unittest/_torch/multi_gpu/test_linear.py tests/unittest/_torch/multi_gpu/test_lowprecision_allreduce.py tests/unittest/_torch/multi_gpu/test_mnnvl_allreduce.py tests/unittest/_torch/multi_gpu/test_mnnvl_memory.py -tests/unittest/_torch/moe/multi_gpu/test_moe_a2a.py tests/unittest/_torch/multi_gpu/test_user_buffers.py tests/unittest/_torch/multi_gpu_modeling/test_deepseek.py tests/unittest/_torch/multimodal/test_external_embedding.py diff --git a/pyproject.toml b/pyproject.toml index 9f98bd665fd8..2e838b5874ef 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -326,7 +326,6 @@ exclude = [ "tensorrt_llm/_torch/modules/triton_linear.py", "tensorrt_llm/_torch/moe/expert_statistic.py", "tensorrt_llm/_torch/moe/fused_moe/__init__.py", - "tensorrt_llm/_torch/moe/fused_moe/communication/moe_alltoall.py", "tensorrt_llm/_torch/moe/fused_moe/create_moe.py", "tensorrt_llm/_torch/moe/fused_moe/deep_ep_utils.py", "tensorrt_llm/_torch/moe/fused_moe/fused_moe_cute_dsl.py", @@ -635,7 +634,6 @@ exclude = [ "tests/unittest/_torch/modeling/test_modeling_vila.py", "tests/unittest/_torch/modules/test_group_rmn_norm.py", "tests/unittest/_torch/modules/test_triton_linear.py", - "tests/unittest/_torch/moe/multi_gpu/test_moe_a2a.py", "tests/unittest/_torch/moe/test_fused_moe.py", "tests/unittest/_torch/moe/test_moe_host_sharer.py", "tests/unittest/_torch/moe/test_moe_load_balancer.py", diff --git a/requirements-dev.txt b/requirements-dev.txt index bf2771e32109..e537349e1fc0 100644 --- a/requirements-dev.txt +++ b/requirements-dev.txt @@ -66,8 +66,8 @@ rapidfuzz==3.14.5 aiperf==0.8.0 nanobind>=2.9.0 nixl-cu13==1.4.0 -cupti-python>=13.0,<13.2 -nvidia-cuda-cupti>=13.0,<13.2 +cupti-python>=13.4,<13.5 +nvidia-cuda-cupti>=13.4,<13.5 cxxfilt hf-transfer==0.1.9 line_profiler diff --git a/ruff-legacy.toml b/ruff-legacy.toml index f2a888f10906..c013a3c897d5 100644 --- a/ruff-legacy.toml +++ b/ruff-legacy.toml @@ -282,7 +282,6 @@ include = [ "tensorrt_llm/_torch/modules/triton_linear.py", "tensorrt_llm/_torch/moe/expert_statistic.py", "tensorrt_llm/_torch/moe/fused_moe/__init__.py", - "tensorrt_llm/_torch/moe/fused_moe/communication/moe_alltoall.py", "tensorrt_llm/_torch/moe/fused_moe/create_moe.py", "tensorrt_llm/_torch/moe/fused_moe/deep_ep_utils.py", "tensorrt_llm/_torch/moe/fused_moe/fused_moe_cute_dsl.py", @@ -591,7 +590,6 @@ include = [ "tests/unittest/_torch/modeling/test_modeling_vila.py", "tests/unittest/_torch/modules/test_group_rmn_norm.py", "tests/unittest/_torch/modules/test_triton_linear.py", - "tests/unittest/_torch/moe/multi_gpu/test_moe_a2a.py", "tests/unittest/_torch/moe/test_fused_moe.py", "tests/unittest/_torch/moe/test_moe_host_sharer.py", "tests/unittest/_torch/moe/test_moe_load_balancer.py", diff --git a/tensorrt_llm/__init__.py b/tensorrt_llm/__init__.py index 516e31ad2c52..11359ee36080 100644 --- a/tensorrt_llm/__init__.py +++ b/tensorrt_llm/__init__.py @@ -52,7 +52,6 @@ import tensorrt_llm.runtime as runtime import tensorrt_llm.tools as tools - from ._mnnvl_utils import MnnvlMemory, MnnvlMoe, MoEAlltoallInfo from ._utils import (default_gpus_per_node, local_mpi_rank, local_mpi_size, mpi_barrier, mpi_comm, mpi_rank, mpi_world_size, set_mpi_comm, str_dtype_to_torch) @@ -74,9 +73,6 @@ 'quantization': ('tensorrt_llm.quantization', None), 'runtime': ('tensorrt_llm.runtime', None), 'tools': ('tensorrt_llm.tools', None), - 'MnnvlMemory': ('tensorrt_llm._mnnvl_utils', 'MnnvlMemory'), - 'MnnvlMoe': ('tensorrt_llm._mnnvl_utils', 'MnnvlMoe'), - 'MoEAlltoallInfo': ('tensorrt_llm._mnnvl_utils', 'MoEAlltoallInfo'), 'default_gpus_per_node': ('tensorrt_llm._utils', 'default_gpus_per_node'), 'local_mpi_rank': ('tensorrt_llm._utils', 'local_mpi_rank'), 'local_mpi_size': ('tensorrt_llm._utils', 'local_mpi_size'), @@ -150,9 +146,6 @@ def __dir__(): 'mpi_world_size', 'torch_models', 'Mapping', - 'MnnvlMemory', - 'MnnvlMoe', - 'MoEAlltoallInfo', 'runtime', 'models', 'quantization', diff --git a/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py b/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py index 80288cd04992..c8aa2bdea308 100644 --- a/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py +++ b/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py @@ -747,13 +747,11 @@ def _( @torch.library.register_fake("trtllm::moe_a2a_initialize") def _( workspace: torch.Tensor, + metainfo: torch.Tensor, ep_rank: int, ep_size: int, - max_num_tokens_per_rank: int, - eplb_stats_num_experts: Optional[int] = None, - can_use_cft_counted_writes: bool = False, - ) -> torch.Tensor: - return torch.empty((10, ), dtype=torch.int64, device="cpu") + ) -> None: + return None @torch.library.register_fake("trtllm::moe_a2a_sanitize_expert_ids") def _( @@ -768,10 +766,10 @@ def _( @torch.library.register_fake("trtllm::moe_a2a_get_combine_payload_tensor") def _( workspace: torch.Tensor, + metainfo: torch.Tensor, ep_rank: int, ep_size: int, runtime_max_tokens_per_rank: int, - combine_payload_offset: int, out_dtype: torch.dtype, hidden_size: int, ) -> torch.Tensor: diff --git a/tensorrt_llm/_torch/distributed/__init__.py b/tensorrt_llm/_torch/distributed/__init__.py index 833f78ba1fd1..41e08e21890e 100644 --- a/tensorrt_llm/_torch/distributed/__init__.py +++ b/tensorrt_llm/_torch/distributed/__init__.py @@ -1,3 +1,6 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + from tensorrt_llm.functional import AllReduceFusionOp from .communicator import Distributed, MPIDist, TorchDist diff --git a/tensorrt_llm/_torch/distributed/communicator.py b/tensorrt_llm/_torch/distributed/communicator.py index 03ab0b82a03f..f6e6ec65f45d 100644 --- a/tensorrt_llm/_torch/distributed/communicator.py +++ b/tensorrt_llm/_torch/distributed/communicator.py @@ -1,3 +1,18 @@ +# SPDX-FileCopyrightText: Copyright (c) 2022-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + import math import pickle # nosec B403 from abc import ABC, abstractmethod @@ -17,7 +32,7 @@ except Exception: MPI = None # deferred; functions will error if used when ENABLE_MULTI_DEVICE is True -from tensorrt_llm._mnnvl_utils import init_helix_cp_comm +from tensorrt_llm._torch.distributed.mnnvl_memory import init_helix_cp_comm from tensorrt_llm._utils import (local_mpi_size, mpi_allgather, mpi_barrier, mpi_comm, mpi_disabled, mpi_isend, mpi_isend_object, mpi_recv, mpi_recv_object, diff --git a/tensorrt_llm/_mnnvl_utils.py b/tensorrt_llm/_torch/distributed/mnnvl_memory.py similarity index 79% rename from tensorrt_llm/_mnnvl_utils.py rename to tensorrt_llm/_torch/distributed/mnnvl_memory.py index 2d1ce779a3dd..29df0a7933e9 100644 --- a/tensorrt_llm/_mnnvl_utils.py +++ b/tensorrt_llm/_torch/distributed/mnnvl_memory.py @@ -21,7 +21,7 @@ import time from dataclasses import dataclass from enum import Enum -from typing import Any, List, Optional, Protocol, Union +from typing import Any, List, Optional, Protocol import pynvml import torch @@ -33,10 +33,10 @@ from torch.utils._python_dispatch import _disable_current_modes -from ._dlpack_utils import pack_strided_memory -from ._utils import get_sm_version, mpi_comm, mpi_disabled -from .logger import logger -from .mapping import Mapping +from tensorrt_llm._dlpack_utils import pack_strided_memory +from tensorrt_llm._utils import get_sm_version, mpi_comm, mpi_disabled +from tensorrt_llm.logger import logger +from tensorrt_llm.mapping import Mapping class ProcessGroupComm: @@ -1149,351 +1149,3 @@ def init_helix_cp_comm(mapping: Mapping) -> None: """ if mapping.has_cp_helix() and not mapping.cp_config.get("use_nccl_for_alltoall", True): HelixCpMnnvlMemory.get_comm(mapping) - - -@dataclass -class MoEAlltoallInfo: - local_gather_indices: torch.Tensor - send_rank_count_cumsum: torch.Tensor - send_rank_local_indices: torch.Tensor - recv_rank_count_cumsum: torch.Tensor - recv_rank_local_indices: torch.Tensor - backward_recv_rank_local_indices: torch.Tensor - local_token_allocation_count: int - - -class MnnvlMoe: - moe_workspace: MnnvlMemory = None - moe_prepare_workspace: MnnvlMemory = None - moe_workspace_tensor: torch.Tensor = None - moe_prepare_workspace_tensor: torch.Tensor = None - moe_mapping: Mapping = None - - @staticmethod - def get_moe_workspaces(mapping: Mapping): - if MnnvlMoe.moe_workspace is not None: - assert mapping == MnnvlMoe.moe_mapping, "only one moe mapping supported now" - return MnnvlMoe.moe_workspace_tensor - - MnnvlMoe.moe_mapping = mapping - workspace_size_per_rank = torch.ops.trtllm.get_moe_commworkspace_size_per_rank( - mapping.moe_ep_size - ) - MnnvlMoe.moe_workspace = MnnvlMemory(mapping, workspace_size_per_rank) - MnnvlMoe.moe_workspace_tensor = MnnvlMoe.moe_workspace.as_torch_strided_tensor(torch.uint64) - torch.ops.trtllm.moe_initialize_workspace( - MnnvlMoe.moe_workspace_tensor, mapping.moe_ep_rank, mapping.moe_ep_size - ) - torch.cuda.synchronize() - MnnvlMoe.moe_workspace.comm.barrier() - return MnnvlMoe.moe_workspace_tensor - - @staticmethod - def get_moe_prepare_workspace(mapping: Mapping): - if MnnvlMoe.moe_prepare_workspace_tensor is not None: - assert mapping == MnnvlMoe.moe_mapping, "only one moe mapping supported now" - return MnnvlMoe.moe_prepare_workspace_tensor - workspace_size_per_rank = torch.ops.trtllm.get_moe_prepare_workspace_size_per_rank( - mapping.moe_ep_size - ) - MnnvlMoe.moe_prepare_workspace = MnnvlMemory(mapping, workspace_size_per_rank) - MnnvlMoe.moe_prepare_workspace_tensor = ( - MnnvlMoe.moe_prepare_workspace.as_torch_strided_tensor(torch.uint64) - ) - return MnnvlMoe.moe_prepare_workspace_tensor - - @staticmethod - def checkpoint_prepare() -> None: - """Detach TRT-native two-sided MoE workspaces for checkpointing.""" - for workspace in (MnnvlMoe.moe_workspace, MnnvlMoe.moe_prepare_workspace): - if workspace is not None: - workspace.checkpoint_prepare() - - @staticmethod - def checkpoint_restore(comm: MnnvlCheckpointCommunicator) -> None: - """Restore TRT-native two-sided MoE workspaces at their original virtual addresses.""" - workspaces = (MnnvlMoe.moe_workspace, MnnvlMoe.moe_prepare_workspace) - restored_workspaces = [] - try: - for workspace in workspaces: - if workspace is not None and workspace.checkpoint_restore(comm): - restored_workspaces.append(workspace) - if not restored_workspaces: - return - restored_main_workspace = any( - workspace is MnnvlMoe.moe_workspace for workspace in restored_workspaces - ) - local_error = None - try: - if restored_main_workspace and MnnvlMoe.moe_workspace_tensor is not None: - assert MnnvlMoe.moe_mapping is not None - torch.ops.trtllm.moe_initialize_workspace( - MnnvlMoe.moe_workspace_tensor, - MnnvlMoe.moe_mapping.moe_ep_rank, - MnnvlMoe.moe_mapping.moe_ep_size, - ) - torch.cuda.synchronize() - except Exception as error: - local_error = f"{type(error).__name__}: {error}" - readiness_errors = _checkpoint_allgather( - comm, - local_error, - operation="two-sided frontend readiness", - ) - failed_ranks = [ - f"rank {rank}: {error}" - for rank, error in enumerate(readiness_errors) - if error is not None - ] - if failed_ranks: - raise RuntimeError( - "Native two-sided MoE restore failed on one or more ranks:\n" - + "\n".join(failed_ranks) - ) - except Exception: - for workspace in restored_workspaces: - workspace._checkpoint_restore_failed() - raise - for workspace in restored_workspaces: - workspace._checkpoint_restore_complete() - - @staticmethod - def require_mapped() -> None: - """Reject kernel access while either native MoE workspace is detached.""" - for workspace in (MnnvlMoe.moe_workspace, MnnvlMoe.moe_prepare_workspace): - if workspace is not None and not workspace.mapped: - raise RuntimeError("Native MoE All-to-All workspace handles are unmapped") - - @staticmethod - def compute_target_rank_id( - token_selected_experts: torch.Tensor, expert_count: int, ep_size: int - ): - assert expert_count % ep_size == 0, "expert_count should be divisible by ep_size" - expert_per_rank = expert_count // ep_size - token_target_rank_ids = token_selected_experts // expert_per_rank - return token_target_rank_ids - - @staticmethod - def mnnvl_moe_alltoallv_prepare_without_allgather( - expert_ids: torch.Tensor, - expert_statics: Optional[torch.Tensor], - workspace: torch.Tensor, - max_token_count_per_rank: int, - ep_rank: int, - ep_size: int, - expert_count: int, - slot_count: int, - top_k: int, - ): - ( - local_send_rank_count_cumsum, - local_send_rank_indices, - local_recv_rank_count_cumsum, - local_recv_rank_indices, - backward_local_recv_rank_indices, - gathered_expert_statics, - ) = torch.ops.trtllm.mnnvl_moe_alltoallv_prepare_without_allgather( - expert_ids, - expert_statics, - workspace, - max_token_count_per_rank, - ep_rank, - ep_size, - expert_count, - slot_count, - top_k, - ) - - local_token_allocation_count = max_token_count_per_rank * ep_size - # Looks like we don't need this. - local_gather_indices = None - - alltoall_info = MoEAlltoallInfo( - local_gather_indices, - local_send_rank_count_cumsum, - local_send_rank_indices, - local_recv_rank_count_cumsum, - local_recv_rank_indices, - backward_local_recv_rank_indices, - local_token_allocation_count, - ) - - return alltoall_info, gathered_expert_statics - - @staticmethod - def mnnvl_moe_expert_static_allgather( - expert_ids: torch.Tensor, - workspace: torch.Tensor, - ep_rank: int, - ep_size: int, - expert_count: int, - ): - gathered_expert_ids = torch.ops.trtllm.mnnvl_moe_expert_static_allgather( - expert_ids, workspace, ep_rank, ep_size, expert_count - ) - return gathered_expert_ids - - @staticmethod - def mnnvl_moe_alltoallv_prepare( - gathered_target_rank_ids: torch.Tensor, - real_rank_token_count_cumsum: Optional[torch.Tensor], - gathered_expert_ids: torch.Tensor, - gathered_scales: Optional[torch.Tensor], - max_token_count_per_rank: int, - expert_count: int, - top_k: int, - ep_rank: int, - ep_size: int, - ): - ( - local_gather_indices, - send_rank_count_cumsum, - send_rank_local_indices, - recv_rank_count_cumsum, - recv_rank_local_indices, - backward_recv_rank_local_indices, - ) = torch.ops.trtllm.moe_comm_prepare_indices( - gathered_target_rank_ids, - real_rank_token_count_cumsum, - max_token_count_per_rank, - expert_count, - top_k, - ep_rank, - ep_size, - ) - - local_token_allocation_count = max_token_count_per_rank * ep_size - - local_expert_ids = torch.empty( - local_token_allocation_count, top_k, dtype=torch.int32, device=torch.device("cuda") - ) - if gathered_scales is None: - local_scales = None - else: - local_scales = torch.empty( - local_token_allocation_count, - top_k, - dtype=torch.float32, - device=torch.device("cuda"), - ) - - torch.ops.trtllm.moe_local_gather( - recv_rank_count_cumsum, - local_gather_indices, - gathered_expert_ids, - gathered_scales, - local_expert_ids, - local_scales, - max_token_count_per_rank, - expert_count, - top_k, - ep_rank, - ep_size, - ) - - alltoall_info = MoEAlltoallInfo( - local_gather_indices, - send_rank_count_cumsum, - send_rank_local_indices, - recv_rank_count_cumsum, - recv_rank_local_indices, - backward_recv_rank_local_indices, - local_token_allocation_count, - ) - return alltoall_info, local_expert_ids, local_scales - - @staticmethod - def mnnvl_moe_alltoallv( - x: Union[torch.Tensor, List[Optional[torch.Tensor]]], - alltoall_info: MoEAlltoallInfo, - workspace: torch.Tensor, - ep_rank: int, - ep_size: int, - ) -> Union[torch.Tensor, List[Optional[torch.Tensor]]]: - # Convert single tensor to list for unified handling - is_single_tensor = not isinstance(x, list) - if is_single_tensor: - assert x.dim() == 2, "only 2D tensor supported, please reshape." - x = [x] - - assert len(x) > 0, "Empty tensor list not supported" - - # Filter out None values - valid_list = [tensor is not None for tensor in x] - valid_tensors = [tensor for tensor in x if tensor is not None] - - if len(valid_tensors) == 0: - # All tensors are None, return list of None - result = [None] * len(x) - else: - first_dim = None - for tensor in valid_tensors: - # Validate dimensions of valid tensors - assert tensor.dim() == 2, "only 2D tensor supported, please reshape." - if first_dim is None: - first_dim = tensor.shape[0] - else: - assert tensor.shape[0] == first_dim, ( - f"All tensors must have the same first dimension, got {tensor.shape[0]} vs {first_dim}" - ) - - # Process only valid tensors - output_tensors = torch.ops.trtllm.moe_comm( - valid_tensors, - alltoall_info.send_rank_count_cumsum, - alltoall_info.send_rank_local_indices, - alltoall_info.recv_rank_count_cumsum, - alltoall_info.recv_rank_local_indices, - workspace, - alltoall_info.local_token_allocation_count, - ep_rank, - ep_size, - ) - - # Restore None positions in output - idx = 0 - result = [] - for is_valid in valid_list: - if is_valid: - result.append(output_tensors[idx]) - idx += 1 - else: - result.append(None) - - # If input was a single tensor, return a single tensor - if is_single_tensor: - result = result[0] - - return result - - @staticmethod - def mnnvl_moe_alltoallv_combine( - x: torch.Tensor, - alltoall_info: MoEAlltoallInfo, - workspace: torch.Tensor, - ep_rank: int, - ep_size: int, - top_k: int, - token_count: int, - use_low_precision_combine: bool = False, - do_reduce: bool = True, - ): - assert x.dim() == 2, "2D tensor supported, please reshape." - output_tensors = torch.ops.trtllm.moe_comm( - [x], - alltoall_info.recv_rank_count_cumsum, - alltoall_info.recv_rank_local_indices, - alltoall_info.send_rank_count_cumsum, - alltoall_info.backward_recv_rank_local_indices, - workspace, - token_count * top_k, - ep_rank, - ep_size, - [True], - use_low_precision_combine, - ) - output_tensor = output_tensors[0].reshape(token_count, top_k, x.shape[1]) - if do_reduce: - return torch.sum(output_tensor, dim=1, keepdim=False) - else: - return output_tensor diff --git a/tensorrt_llm/_torch/distributed/ops.py b/tensorrt_llm/_torch/distributed/ops.py index 6d20a9aa6394..e9d6d5ff7b46 100644 --- a/tensorrt_llm/_torch/distributed/ops.py +++ b/tensorrt_llm/_torch/distributed/ops.py @@ -22,9 +22,10 @@ import torch from torch import nn -from tensorrt_llm._mnnvl_utils import HelixCpMnnvlMemory, MnnvlMemory from tensorrt_llm._torch.distributed.allreduce_helper import \ CustomAllReduceHelper +from tensorrt_llm._torch.distributed.mnnvl_memory import (HelixCpMnnvlMemory, + MnnvlMemory) from tensorrt_llm._torch.distributed.symm_mem_allreduce import \ SymmetricMemoryAllReduce from tensorrt_llm._torch.utils import get_model_extra_attrs @@ -813,7 +814,7 @@ def is_mnnvl(mapping: Mapping, where MNNVL is the clear win; an explicit request is honoured on a single node too, as long as the hardware supports it. """ - from tensorrt_llm._mnnvl_utils import MnnvlMemory + from tensorrt_llm._torch.distributed.mnnvl_memory import MnnvlMemory arch = platform.machine().lower() is_on_aarch64 = "aarch64" in arch diff --git a/tensorrt_llm/_torch/mnnvl_alltoall_workspace.py b/tensorrt_llm/_torch/mnnvl_alltoall_workspace.py index 07dde82ee935..7238597c5b08 100644 --- a/tensorrt_llm/_torch/mnnvl_alltoall_workspace.py +++ b/tensorrt_llm/_torch/mnnvl_alltoall_workspace.py @@ -19,17 +19,17 @@ import torch -from tensorrt_llm._mnnvl_utils import ( - MnnvlCheckpointCommunicator, - MnnvlMemory, - _checkpoint_allgather, -) from tensorrt_llm._torch.alltoall_watchdog import ( AlltoAllWatchdog, AlltoAllWatchdogCoordinator, AlltoAllWatchdogTimeout, EPGroupHealthLike, ) +from tensorrt_llm._torch.distributed.mnnvl_memory import ( + MnnvlCheckpointCommunicator, + MnnvlMemory, + _checkpoint_allgather, +) _WORKSPACE_LIFECYCLE_KEY = "mnnvl_alltoall_workspace_lifecycle" diff --git a/tensorrt_llm/_torch/modules/dwdp/transport.py b/tensorrt_llm/_torch/modules/dwdp/transport.py index 5a4abbb6ac53..c0a797562720 100644 --- a/tensorrt_llm/_torch/modules/dwdp/transport.py +++ b/tensorrt_llm/_torch/modules/dwdp/transport.py @@ -76,7 +76,7 @@ # to dup an FD from a sibling DWDP MPI worker into the local fd table so # ``cuMemImportFromShareableHandle(fd, POSIX_FILE_DESCRIPTOR)`` accepts it. # Mirrors ``MnnvlMemory.open_mnnvl_memory`` in -# ``tensorrt_llm/_mnnvl_utils.py``. +# ``tensorrt_llm/_torch/distributed/mnnvl_memory.py``. _SYS_pidfd_open = 434 _SYS_pidfd_getfd = 438 diff --git a/tensorrt_llm/_torch/modules/dwdp/vmm.py b/tensorrt_llm/_torch/modules/dwdp/vmm.py index f7cf16a10a1d..b472152d0bb5 100644 --- a/tensorrt_llm/_torch/modules/dwdp/vmm.py +++ b/tensorrt_llm/_torch/modules/dwdp/vmm.py @@ -201,7 +201,7 @@ def peer_handle_type() -> cuda.CUmemAllocationHandleType: ``CUDA_ERROR_NOT_PERMITTED`` (800), so we use ``CU_MEM_HANDLE_TYPE_POSIX_FILE_DESCRIPTOR`` and exchange the FDs between sibling MPI workers via ``pidfd_open`` / ``pidfd_getfd`` - (mirrors ``MnnvlMemory.get_allocation_prop`` in ``_mnnvl_utils.py``). + (mirrors ``MnnvlMemory.get_allocation_prop`` in ``mnnvl_memory.py``). """ arch = platform.machine().lower() if "aarch64" in arch: diff --git a/tensorrt_llm/_torch/moe/fused_moe/communication/deep_ep.py b/tensorrt_llm/_torch/moe/fused_moe/communication/deep_ep.py index ae71d3f60653..9964bfac129e 100644 --- a/tensorrt_llm/_torch/moe/fused_moe/communication/deep_ep.py +++ b/tensorrt_llm/_torch/moe/fused_moe/communication/deep_ep.py @@ -25,7 +25,7 @@ import torch -from tensorrt_llm._mnnvl_utils import MnnvlMemory +from tensorrt_llm._torch.distributed.mnnvl_memory import MnnvlMemory from tensorrt_llm._torch.moe.fused_moe.deep_ep_utils import buffer_pool, deep_ep_installed from tensorrt_llm._utils import local_mpi_size from tensorrt_llm.mapping import Mapping diff --git a/tensorrt_llm/_torch/moe/fused_moe/communication/deep_ep_low_latency.py b/tensorrt_llm/_torch/moe/fused_moe/communication/deep_ep_low_latency.py index dc033ea9ba67..eec67d706bf5 100644 --- a/tensorrt_llm/_torch/moe/fused_moe/communication/deep_ep_low_latency.py +++ b/tensorrt_llm/_torch/moe/fused_moe/communication/deep_ep_low_latency.py @@ -25,7 +25,7 @@ import torch -from tensorrt_llm._mnnvl_utils import MnnvlMemory +from tensorrt_llm._torch.distributed.mnnvl_memory import MnnvlMemory from tensorrt_llm._torch.moe.fused_moe.deep_ep_utils import buffer_pool, deep_ep_installed from tensorrt_llm._utils import get_sm_version from tensorrt_llm.mapping import Mapping diff --git a/tensorrt_llm/_torch/moe/fused_moe/communication/moe_alltoall.py b/tensorrt_llm/_torch/moe/fused_moe/communication/moe_alltoall.py deleted file mode 100644 index e9bb8035e9ec..000000000000 --- a/tensorrt_llm/_torch/moe/fused_moe/communication/moe_alltoall.py +++ /dev/null @@ -1,725 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. -""" -MoE All-to-All Operations - -This module provides a high-level interface for MoE all-to-all dispatch and combine operations -with proper workspace management and synchronization. -""" - -# ruff: noqa: E501 - -import os -import sys -from dataclasses import dataclass -from typing import Callable, Dict, Optional - -import torch - -from tensorrt_llm._mnnvl_utils import (CftMnnvlMemory, - MnnvlCheckpointCommunicator, MnnvlMemory) -from tensorrt_llm._torch.alltoall_watchdog import ( - DEFAULT_ALLTOALL_WATCHDOG_POLL_INTERVAL_S, - DEFAULT_ALLTOALL_WATCHDOG_TIMEOUT_S, ActiveRankMaskSnapshot, - AlltoAllWatchdog, AlltoAllWatchdogCoordinator, AlltoAllWatchdogTimeout, - EPGroupHealthLike, reject_rank_mask_cuda_graph_capture) -from tensorrt_llm._torch.mnnvl_alltoall_workspace import \ - _MnnvlAlltoAllWorkspaceLifecycle -from tensorrt_llm.bindings import internal as _tllm_internal -from tensorrt_llm.logger import logger as tllm_logger -from tensorrt_llm.mapping import Mapping -from tensorrt_llm.math_utils import pad_up - -_CFT_DEFAULT_MAX_BATCH_FOR_DISPATCH = 128 -_CFT_MAX_BATCH_FOR_DISPATCH_ENV = "TRTLLM_MOE_A2A_CFT_MAX_BATCH_FOR_DISPATCH" -# CFT combine wins at small/medium batch and ties/regresses at large batch, so -# it is gated by the same per-call token-count threshold as dispatch. -_CFT_DEFAULT_MAX_BATCH_FOR_COMBINE = 128 -_CFT_MAX_BATCH_FOR_COMBINE_ENV = "TRTLLM_MOE_A2A_CFT_MAX_BATCH_FOR_COMBINE" -FORCE_CFT_ENV = "TRTLLM_MOE_A2A_FORCE_CFT" -_CFT_ALIGNMENT_BYTES = 16 - - -def get_force_cft() -> bool | None: - value = os.environ.get(FORCE_CFT_ENV) - if value == "0": - return False - if value == "1": - return True - return None - - -def resolve_can_use_cft(can_use_cft_counted_writes: bool) -> bool: - """Apply the TRTLLM_MOE_A2A_FORCE_CFT override to a caller's request. - - Workspace sizing and workspace layout both depend on this, so they must - resolve it identically: a caller that sizes without the override and then - constructs with it would lay out the CFT region in an undersized buffer. - """ - force_cft = get_force_cft() - if force_cft is None: - return can_use_cft_counted_writes - return force_cft - - -def should_use_cft( - can_use_cft: bool, - force_cft: bool | None, - max_batch: int | None, - runtime_max_tokens_per_rank: int, -) -> bool: - if not can_use_cft: - return False - if force_cft is not None: - return force_cft - if max_batch is None: - return True - return runtime_max_tokens_per_rank <= max_batch - - -def _use_cft_for_dispatch_payloads(use_cft: bool, - payloads: list[torch.Tensor]) -> bool: - if not use_cft: - return False - for payload_index, payload in enumerate(payloads): - bytes_per_token = payload.shape[1] * payload.element_size() - if bytes_per_token % _CFT_ALIGNMENT_BYTES != 0: - tllm_logger.warning_once( - "CFT counted writes disabled: dispatch payload " - f"{payload_index} has {bytes_per_token} bytes per token, which is not " - f"{_CFT_ALIGNMENT_BYTES}-byte aligned. Falling back to fence-based dispatch.", - key= - f"moe_a2a_cft_dispatch_alignment_{payload_index}_{bytes_per_token}", - ) - return False - return True - - -def _use_cft_for_combine_payload(use_cft: bool, payload: torch.Tensor, - use_low_precision: bool) -> bool: - if not use_cft: - return False - wire_element_size = 1 if use_low_precision else payload.element_size() - bytes_per_token = payload.shape[-1] * wire_element_size - if bytes_per_token % _CFT_ALIGNMENT_BYTES != 0: - tllm_logger.warning_once( - "CFT counted writes disabled: combine payload has " - f"{bytes_per_token} bytes per token, which is not " - f"{_CFT_ALIGNMENT_BYTES}-byte aligned. Falling back to fence-based combine.", - key=f"moe_a2a_cft_combine_alignment_{bytes_per_token}", - ) - return False - return True - - -def _get_cft_max_batch(env_name: str, default: int) -> int: - env_value = os.environ.get(env_name) - if env_value is None: - return default - try: - threshold = int(env_value) - except ValueError as e: - raise ValueError(f"{env_name} must be an integer") from e - if threshold < 0: - raise ValueError(f"{env_name} must be non-negative") - return threshold - - -def _get_cft_max_batch_for_dispatch() -> int | None: - return _get_cft_max_batch(_CFT_MAX_BATCH_FOR_DISPATCH_ENV, - _CFT_DEFAULT_MAX_BATCH_FOR_DISPATCH) - - -def _get_cft_max_batch_for_combine() -> int | None: - return _get_cft_max_batch(_CFT_MAX_BATCH_FOR_COMBINE_ENV, - _CFT_DEFAULT_MAX_BATCH_FOR_COMBINE) - - -@dataclass -class _A2AState: - phase: str = "idle" # idle | dispatched - local_num_tokens: int | None = None - combine_payload_offset: int | None = None - eplb_gathered_stats: torch.Tensor | None = None - active_rank_mask_snapshot: ActiveRankMaskSnapshot | None = None - - -class MoeAlltoAll: - """ - Manages MoE All-to-All operations with proper workspace allocation and synchronization. - - This class encapsulates the dispatch and combine operations, managing workspace memory - and auxiliary data structures needed for cross-GPU communication. - """ - - # Shared workspace/memory across the process, separated by handle type. - _WORKSPACES: Dict[bool, dict] = {} - - _METAINFO_INDEX: Dict[str, int] | None = None - - @staticmethod - def get_aux_data_size( - ep_size: int, - max_num_tokens: int, - eplb_stats_num_experts: Optional[int] = None, - can_use_cft_counted_writes: bool = False, - ) -> int: - return torch.ops.trtllm.moe_a2a_get_aux_data_size( - ep_size, max_num_tokens, eplb_stats_num_experts, - can_use_cft_counted_writes) - - @staticmethod - def calculate_required_workspace_size( - ep_size: int, - top_k: int, - max_num_tokens: int, - hidden_size: int, - dtype: torch.dtype, - eplb_stats_num_experts: Optional[int] = None, - extra_payload_bytes_per_token: int = 0, - can_use_cft_counted_writes: bool = False) -> int: - can_use_cft_counted_writes = resolve_can_use_cft( - can_use_cft_counted_writes) - element_size = dtype.itemsize - - # Auxiliary data size - workspace_size = MoeAlltoAll.get_aux_data_size( - ep_size, max_num_tokens, eplb_stats_num_experts, - can_use_cft_counted_writes) - - # Dispatch needs workspace for [ep_size, max_tokens] tokens, - # but due to the variety of quantization recipes, we cannot know the exact size, so we conservatively estimate assuming no quantization. - # Meanwhile, we consider the alignment requirement as in moeA2ADispatchOp and moeA2ACombineOp. - # (Unquantized) token hidden states - workspace_size += ep_size * max_num_tokens * hidden_size * element_size - workspace_size = pad_up(workspace_size, 128) - # token_selected_experts - workspace_size += ep_size * max_num_tokens * top_k * 4 - workspace_size = pad_up(workspace_size, 128) - # token_final_scales - workspace_size += ep_size * max_num_tokens * top_k * 4 - workspace_size = pad_up(workspace_size, 128) - # extra payload bytes per token - workspace_size += ep_size * max_num_tokens * extra_payload_bytes_per_token - workspace_size = pad_up(workspace_size, 128) - - # Required workspace for combine [ep_size, max_tokens] tokens - workspace_size += ep_size * max_num_tokens * hidden_size * element_size - workspace_size = pad_up(workspace_size, 128) - - # CFT combine: dedicated combine RECEIVE region C (peer pushes land here; - # prepareCombine never touches it -> no proxy aliasing). Same size as the combine region. - if can_use_cft_counted_writes: - workspace_size += ep_size * max_num_tokens * hidden_size * element_size - workspace_size = pad_up(workspace_size, 128) - - return workspace_size - - @classmethod - def _init_constants(cls): - """Initialize constants from C++ if not already done.""" - # TODO: Can we avoid such code duplication? - if cls._METAINFO_INDEX is None: - thop = _tllm_internal.thop - cls._METAINFO_INDEX = { - "FLAG_VAL_OFFSET_INDEX": - int(thop.MOE_A2A_FLAG_VAL_OFFSET_INDEX), - "LOCAL_TOKEN_COUNTER_OFFSET_INDEX": - int(thop.MOE_A2A_LOCAL_TOKEN_COUNTER_OFFSET_INDEX), - "SEND_COUNTERS_OFFSET_INDEX": - int(thop.MOE_A2A_SEND_COUNTERS_OFFSET_INDEX), - "RECV_COUNTERS_OFFSET_INDEX": - int(thop.MOE_A2A_RECV_COUNTERS_OFFSET_INDEX), - "DISPATCH_COMPLETION_FLAGS_OFFSET_INDEX": - int(thop.MOE_A2A_DISPATCH_COMPLETION_FLAGS_OFFSET_INDEX), - "COMBINE_COMPLETION_FLAGS_OFFSET_INDEX": - int(thop.MOE_A2A_COMBINE_COMPLETION_FLAGS_OFFSET_INDEX), - "DISPATCH_COUNTED_WRITE_COUNTERS_OFFSET_INDEX": - int(thop.MOE_A2A_DISPATCH_COUNTED_WRITE_COUNTERS_OFFSET_INDEX), - "TOPK_TARGET_RANKS_OFFSET_INDEX": - int(thop.MOE_A2A_TOPK_TARGET_RANKS_OFFSET_INDEX), - "TOPK_SEND_INDICES_OFFSET_INDEX": - int(thop.MOE_A2A_TOPK_SEND_INDICES_OFFSET_INDEX), - "EPLB_GATHERED_STATS_OFFSET_INDEX": - int(thop.MOE_A2A_EPLB_GATHERED_STATS_OFFSET_INDEX), - "PAYLOAD_DATA_OFFSET_INDEX": - int(thop.MOE_A2A_PAYLOAD_DATA_OFFSET_INDEX), - "NUM_METAINFO_FIELDS": - int(thop.MOE_A2A_NUM_METAINFO_FIELDS), - } - - def __init__( - self, - mapping: Mapping, - max_num_tokens: int, - top_k: int, - num_slots: int, - workspace_size_per_rank: int, - num_experts: Optional[int] = None, - can_use_cft_counted_writes: bool = False, - ep_group_health: Optional[EPGroupHealthLike] = None, - alltoall_watchdog_timeout_s: Optional[float] = None, - alltoall_watchdog_poll_interval_s: - float = DEFAULT_ALLTOALL_WATCHDOG_POLL_INTERVAL_S, - alltoall_watchdog_on_timeout: Optional[Callable[ - [AlltoAllWatchdogTimeout], None]] = None, - ) -> None: - """ - Initialize MoeAlltoAll with workspace allocation. - - Args: - mapping: TensorRT-LLM Mapping object containing rank information - max_num_tokens: Maximum number of tokens supported. Should be ModelConfig.max_num_tokens. - workspace_size_per_rank: Size of workspace per rank in bytes - num_slots: Number of routing slots (token_selected_experts values are in [0, num_slots)). - Note: The terminology is mapped to `num_experts` in this class and the kernels. - num_experts: (Optional) Number of experts for EPLB stats (must be <= num_slots). DO NOT provide this parameter if EPLB is not enabled. - Note: The terminology is mapped to `eplb_stats_num_experts` in this class and the kernels. - can_use_cft_counted_writes: If True, allow CFT handle-based counted - writes (fabric.try_put.counted via Logical Endpoints) for dispatch. - Requires sm_100+ (Blackwell or later), a build against CUDA 13.4+, an - NVLink fabric, and a driver exporting the CUDA logical endpoint API. - ep_group_health: Optional read-only committed EP membership. When present, rank-mask handling is - enabled in the CUDA kernels, and its mask defines the peers expected by the watchdog. Timeout - detection never mutates it. CUDA graphs are rejected until membership-scoped recapture lands. - alltoall_watchdog_timeout_s: Optional timeout for the host-side AlltoAll watchdog. If None, the - watchdog is disabled. - alltoall_watchdog_poll_interval_s: Poll interval for the watchdog thread. - alltoall_watchdog_on_timeout: Optional callback invoked when the watchdog reports suspects. - """ - # Check for environment variable override - workspace_mb_env = os.environ.get("TRTLLM_MOE_A2A_WORKSPACE_MB") - if workspace_mb_env: - workspace_size_env = int(workspace_mb_env) * 1024 * 1024 - tllm_logger.warning( - f"Overriding automatically calculated workspace_size_per_rank ({workspace_size_per_rank} bytes) with " - f"TRTLLM_MOE_A2A_WORKSPACE_MB={workspace_mb_env} ({workspace_size_env} bytes)." - f"Automatically calculated workspace_size_per_rank is conservatively large, please only consider overriding it if you have a specific reason." - ) - workspace_size_per_rank = workspace_size_env - - # Initialize constants from C++ - self._init_constants() - - # Initialize or reuse workspace - MnnvlMemory.initialize() - - self.workspace_size_per_rank = workspace_size_per_rank - self.max_num_tokens = max_num_tokens - self.ep_size = mapping.moe_ep_size - self.ep_rank = mapping.moe_ep_rank - - self.top_k = top_k - self.num_experts = num_slots - - if not isinstance(self.top_k, int) or self.top_k <= 0: - raise ValueError("top_k must be a positive int") - if not isinstance(self.num_experts, int) or self.num_experts <= 0: - raise ValueError("num_slots must be a positive int") - - if num_experts is not None: - assert num_experts > 0 and num_experts <= num_slots, "num_experts must be in (0, num_slots]" - tllm_logger.info( - "NVLinkOneSided AlltoAll: EPLB is enabled, with num_slots=" - f"{num_slots} and num_experts={num_experts}") - self.enable_eplb = num_experts is not None - self.eplb_stats_num_experts = num_experts - self._force_cft = get_force_cft() - # Opt-in only: no caller passes can_use_cft_counted_writes=True, so - # without the override the CFT path cannot be reached at all. Leaving - # the variable unset keeps CFT disabled, as before. - can_use_cft_counted_writes = resolve_can_use_cft( - can_use_cft_counted_writes) - self.can_use_cft_counted_writes = can_use_cft_counted_writes - if self._force_cft is None: - self.cft_max_batch_for_dispatch = _get_cft_max_batch_for_dispatch() - self.cft_max_batch_for_combine = _get_cft_max_batch_for_combine() - else: - self.cft_max_batch_for_dispatch = None - self.cft_max_batch_for_combine = None - - workspace_key = self.can_use_cft_counted_writes - workspace_entry = self._WORKSPACES.get(workspace_key) - memory_cls = CftMnnvlMemory if self.can_use_cft_counted_writes else MnnvlMemory - - if workspace_entry is None: - tllm_logger.info( - f"NVLinkOneSided AlltoAll: Allocating workspace with size {workspace_size_per_rank} bytes. ep_rank: {self.ep_rank}, ep_size: {self.ep_size}, max_num_tokens: {self.max_num_tokens}" - ) - mnnvl_mem = memory_cls(mapping, workspace_size_per_rank) - workspace = mnnvl_mem.as_torch_strided_tensor(torch.uint8) - metainfo = torch.ops.trtllm.moe_a2a_initialize( - workspace, self.ep_rank, self.ep_size, self.max_num_tokens, - self.eplb_stats_num_experts, self.can_use_cft_counted_writes) - workspace_entry = { - "workspace_size_per_rank": workspace_size_per_rank, - "max_num_tokens": self.max_num_tokens, - "ep_rank": self.ep_rank, - "ep_size": self.ep_size, - "eplb_stats_num_experts": self.eplb_stats_num_experts, - "can_use_cft_counted_writes": self.can_use_cft_counted_writes, - "mnnvl_mem": mnnvl_mem, - "workspace": workspace, - "metainfo": metainfo, - "cft_initialized": False, - } - MoeAlltoAll._WORKSPACES[workspace_key] = workspace_entry - else: - assert workspace_entry[ - "workspace_size_per_rank"] == workspace_size_per_rank, "mistakenly reusing workspace with different workspace_size_per_rank" - assert workspace_entry[ - "max_num_tokens"] == self.max_num_tokens, "mistakenly reusing workspace with different max_num_tokens" - assert workspace_entry[ - "ep_rank"] == self.ep_rank, "mistakenly reusing workspace with different ep_rank" - assert workspace_entry[ - "ep_size"] == self.ep_size, "mistakenly reusing workspace with different ep_size" - assert workspace_entry[ - "eplb_stats_num_experts"] == self.eplb_stats_num_experts, ( - "reuse workspace with different eplb_stats_num_experts") - assert workspace_entry[ - "can_use_cft_counted_writes"] == self.can_use_cft_counted_writes, "reuse workspace with different CFT mode" - - workspace_state = workspace_entry - self.mnnvl_mem = workspace_entry["mnnvl_mem"] - self.workspace = workspace_entry["workspace"] - # Internal state - self._state: _A2AState = _A2AState() - self.ep_group_health = ep_group_health - # Keep the kernel specialization stable for this communicator's lifetime. - self._rank_mask_enabled = ep_group_health is not None - self._workspace_state = workspace_state - if (alltoall_watchdog_timeout_s is None - and self.ep_group_health is not None): - alltoall_watchdog_timeout_s = DEFAULT_ALLTOALL_WATCHDOG_TIMEOUT_S - metainfo_index = self._METAINFO_INDEX - assert metainfo_index is not None - self._workspace_lifecycle = ( - _MnnvlAlltoAllWorkspaceLifecycle.get_or_create( - workspace_state=workspace_state, - memory=self.mnnvl_mem, - workspace=self.workspace, - metainfo=workspace_state["metainfo"], - metainfo_index=metainfo_index, - ep_rank=self.ep_rank, - ep_size=self.ep_size, - health=self.ep_group_health, - )) - self._destroyed = False - self._workspace_registered = False - self._workspace_lifecycle.register( - self, - watchdog_timeout_s=alltoall_watchdog_timeout_s, - watchdog_poll_interval_s=alltoall_watchdog_poll_interval_s, - watchdog_on_timeout=alltoall_watchdog_on_timeout, - ) - self._workspace_registered = True - - @property - def metainfo(self) -> torch.Tensor: - return self._workspace_lifecycle.metainfo - - @property - def _watchdog_coordinator(self) -> AlltoAllWatchdogCoordinator: - return self._workspace_lifecycle.coordinator - - @property - def _alltoall_watchdog(self) -> AlltoAllWatchdog | None: - return self._workspace_lifecycle.watchdog_for(self) - - def checkpoint_resource_key(self) -> int: - """Identify wrappers sharing the same MNNVL workspace lifecycle.""" - return id(self._workspace_lifecycle) - - def destroy(self) -> None: - """Stop background watchdog resources owned by this wrapper.""" - if getattr(self, "_destroyed", False): - return - self._destroyed = True - lifecycle = getattr(self, "_workspace_lifecycle", None) - if lifecycle is not None and getattr(self, "_workspace_registered", - False): - lifecycle.unregister(self) - self._workspace_registered = False - - def __del__(self) -> None: - if not sys.is_finalizing(): - self.destroy() - - def use_cft_for_dispatch(self, runtime_max_tokens_per_rank: int) -> bool: - return should_use_cft(self.can_use_cft_counted_writes, self._force_cft, - self.cft_max_batch_for_dispatch, - runtime_max_tokens_per_rank) - - def use_cft_for_combine(self, runtime_max_tokens_per_rank: int) -> bool: - return should_use_cft(self.can_use_cft_counted_writes, self._force_cft, - self.cft_max_batch_for_combine, - runtime_max_tokens_per_rank) - - def cft_initialize(self) -> None: - """ - Initialize CFT Logical Endpoints by binding the LE to the MNNVL workspace. - Must be called once before the first dispatch when can_use_cft_counted_writes=True. - """ - if not self.can_use_cft_counted_writes: - raise ValueError( - "cft_initialize called but can_use_cft_counted_writes is False") - torch.ops.trtllm.moe_a2a_cft_initialize( - self.workspace, - self.mnnvl_mem.local_mem_handle, - int(self.workspace.size(1)), - self.ep_rank, - self.ep_size, - ) - tllm_logger.info( - f"CFT LE initialized (workspace-bound): ep_rank={self.ep_rank}, ep_size={self.ep_size}" - ) - - def _require_mapped(self) -> None: - if not self.mnnvl_mem.mapped: - raise RuntimeError( - "Native MoE All-to-All workspace handles are unmapped") - - def checkpoint_prepare(self) -> None: - """Collectively detach handles after every shared owner is idle.""" - if self.can_use_cft_counted_writes: - raise RuntimeError( - "Checkpointing a CFT-backed MoE All-to-All workspace is not supported" - ) - self._workspace_lifecycle.checkpoint_prepare() - - def checkpoint_restore( - self, - comm: MnnvlCheckpointCommunicator | None = None, - ) -> None: - """Collectively restore handles and all shared frontend state. - - Args: - comm: An mpi4py-like communicator exposing ``Get_rank()``, - ``Get_size()``, ``allgather()``, and ``barrier()``. Its local - rank and size must match the communicator used for the - original allocation. Every rank must call this method - symmetrically. - """ - if comm is None: - comm = self.mnnvl_mem.comm - if comm is None: - raise RuntimeError( - "MNNVL workspace communicator is not initialized") - self._workspace_lifecycle.checkpoint_restore( - comm, - lambda: torch.ops.trtllm.moe_a2a_initialize( - self.workspace, - self.ep_rank, - self.ep_size, - self.max_num_tokens, - self.eplb_stats_num_experts, - self.can_use_cft_counted_writes, - ), - ) - - def _mnnvl_checkpoint_is_idle(self) -> bool: - return self._state.phase == "idle" - - def _mnnvl_checkpoint_reset(self) -> None: - self.reset_state() - - def dispatch(self, - token_selected_experts: torch.Tensor, - input_payloads: list[torch.Tensor], - runtime_max_tokens_per_rank: int, - invalid_token_expert_id: Optional[int] = None, - expert_id_payload_index: Optional[int] = None, - eplb_local_stats: Optional[torch.Tensor] = None, - active_rank_mask: Optional[torch.Tensor] = None): - """ - Perform MoE all-to-all dispatch operation. - - Args: - token_selected_experts: [local_num_tokens, top_k] tensor of expert indices - input_payloads: List of tensors to dispatch, each has shape [local_num_tokens, payload_num_elements_per_token] - runtime_max_tokens_per_rank: Maximum of the number of tokens of each DP rank's local batch. - invalid_token_expert_id: If not None, set the token_selected_experts of the invalid tokens to this expert id. This is used to notify the MoE to skip these tokens for GroupGEMM. - expert_id_payload_index: The index of token_selected_experts in the input_payloads. Must be provided if invalid_token_expert_id is not None. - eplb_local_stats: (Optional) [num_experts] tensor containing local statistics for EPLB - active_rank_mask: Optional uint64 CPU tensor overriding committed membership in rank-mask mode. When - omitted, the committed mask and generation are captured together. Combine reuses that mask and - fails closed if the committed generation changes first. The masked kernel rejects inactive routes - before remote access; that sentinel is an internal abort artifact, not valid model output. - - Returns: - recv_tensors: List of tensors received, each has shape [ep_size, max_tokens_per_rank, payload_num_elements_per_token] - """ - self._require_mapped() - assert self._state.phase == "idle", "dispatch called twice without an intervening combine" - reject_rank_mask_cuda_graph_capture(self._rank_mask_enabled) - assert runtime_max_tokens_per_rank <= self.max_num_tokens, "runtime_max_tokens_per_rank must not exceed max_num_tokens" - can_use_cft_for_dispatch = self.use_cft_for_dispatch( - runtime_max_tokens_per_rank) - can_use_cft_for_dispatch = _use_cft_for_dispatch_payloads( - can_use_cft_for_dispatch, input_payloads) - # Auto-initialize CFT LEs on first dispatch only - if self.can_use_cft_counted_writes and not self._workspace_state.get( - 'cft_initialized', False): - self.cft_initialize() - self._workspace_state['cft_initialized'] = True - if eplb_local_stats is not None: - assert self.enable_eplb, "eplb_local_stats provided but enable_eplb is False" - assert eplb_local_stats.dim( - ) == 1, "eplb_local_stats must be a 1D tensor" - assert eplb_local_stats.size( - 0 - ) == self.eplb_stats_num_experts, "eplb_local_stats size must match eplb_stats_num_experts" - can_fuse_sanitize = (can_use_cft_for_dispatch - and invalid_token_expert_id is not None - and expert_id_payload_index is not None) - - requested_active_rank_mask = active_rank_mask - if (not self._rank_mask_enabled - and requested_active_rank_mask is not None): - raise ValueError( - "active_rank_mask requires committed EP group health") - active_rank_mask_snapshot = self._watchdog_coordinator.capture_active_rank_mask( - requested_active_rank_mask) - active_rank_mask = active_rank_mask_snapshot.active_rank_mask - recv_tensors, combine_payload_offset, eplb_gathered_stats = torch.ops.trtllm.moe_a2a_dispatch( - token_selected_experts, - input_payloads, - self.workspace, - self.metainfo, - runtime_max_tokens_per_rank, - self.ep_rank, - self.ep_size, - self.top_k, - self.num_experts, - eplb_local_stats, - can_use_cft_for_dispatch, - expert_id_payload_index if can_fuse_sanitize else None, - invalid_token_expert_id if can_fuse_sanitize else None, - self._rank_mask_enabled, - active_rank_mask, - ) - self._watchdog_coordinator.watch_collective(self._alltoall_watchdog, - "dispatch", - active_rank_mask) - if eplb_gathered_stats.numel() == 0: - eplb_gathered_stats = None - - # Update state together after successful dispatch - self._state.local_num_tokens = token_selected_experts.size(0) - self._state.combine_payload_offset = combine_payload_offset - self._state.eplb_gathered_stats = eplb_gathered_stats - self._state.active_rank_mask_snapshot = active_rank_mask_snapshot - self._state.phase = "dispatched" - - if invalid_token_expert_id is not None and not can_fuse_sanitize: - assert expert_id_payload_index is not None, "expert_id_payload_index must be provided if invalid_token_expert_id is not None" - # Sanitize expert IDs for invalid tokens directly on the recv tensor payload - recv_token_selected_experts = recv_tensors[expert_id_payload_index] - torch.ops.trtllm.moe_a2a_sanitize_expert_ids( - recv_token_selected_experts, - self.workspace, - self.metainfo, - self.ep_rank, - invalid_token_expert_id, - ) - - return recv_tensors - - def combine( - self, - payload, - runtime_max_tokens_per_rank: int, - payload_in_workspace: bool = False, - use_low_precision_combine: bool = False, - active_rank_mask: Optional[torch.Tensor] = None, - ): - """ - Perform MoE all-to-all combine operation. - - Args: - payload: [ep_size, max_tokens_per_rank, num_elements_per_token] tensor to combine. The dtype must be float32, bfloat16 or float16. - runtime_max_tokens_per_rank: Maximum of the number of tokens of each DP rank's local batch. - payload_in_workspace: If True, 'payload' is a view into 'workspace' at 'combine_payload_offset' and no staging copy is needed. If False, the op stages 'payload' into the workspace region before combining. Callers that cannot direct the MoE kernel's output into the workspace must leave this False. - use_low_precision_combine: If True, quantize the combine payload to FP8 for NVLink transfer (halves NVLink bandwidth usage, output precision is preserved). - active_rank_mask: Optional uint64 CPU tensor. In rank-mask mode, it must match the mask captured by - dispatch for this collective when supplied. A committed-generation change since dispatch aborts - the collective epoch. - - Returns: - combined_output: [local_num_tokens, num_elements_per_token] tensor of combined results - """ - self._require_mapped() - assert self._state.phase == "dispatched", "combine called before a successful dispatch" - reject_rank_mask_cuda_graph_capture(self._rank_mask_enabled) - assert runtime_max_tokens_per_rank <= self.max_num_tokens, "runtime_max_tokens_per_rank must not exceed max_num_tokens" - - active_rank_mask_snapshot = self._state.active_rank_mask_snapshot - assert active_rank_mask_snapshot is not None - requested_active_rank_mask = active_rank_mask - if (not self._rank_mask_enabled - and requested_active_rank_mask is not None): - raise ValueError( - "active_rank_mask requires committed EP group health") - active_rank_mask = self._watchdog_coordinator.active_rank_mask_for_combine( - active_rank_mask_snapshot, requested_active_rank_mask) - use_cft_for_combine = _use_cft_for_combine_payload( - self.use_cft_for_combine(runtime_max_tokens_per_rank), payload, - use_low_precision_combine) - output = torch.ops.trtllm.moe_a2a_combine( - payload, self._state.local_num_tokens, self.workspace, - self.metainfo, runtime_max_tokens_per_rank, self.ep_rank, - self.ep_size, self.top_k, self._state.combine_payload_offset, - payload_in_workspace, use_low_precision_combine, - use_cft_for_combine, self._rank_mask_enabled, active_rank_mask) - self._watchdog_coordinator.watch_collective(self._alltoall_watchdog, - "combine", active_rank_mask) - - # Reset state for next round - self.reset_state() - - return output - - def reset_state(self) -> None: - """Reset the dispatch/combine state machine to ``idle``. - - Safe to call between forward passes (or from an error handler) to - recover from a forward that called ``dispatch`` but did not reach - ``combine`` — e.g. because an OOM aborted the forward. Without this, - the next ``dispatch`` would fire the assert at line 239. - """ - self._state = _A2AState() - - def get_combine_payload_tensor_in_workspace( - self, runtime_max_tokens_per_rank: int, hidden_size: int, - dtype: torch.dtype) -> torch.Tensor: - """ - Return the combine payload tensor in the workspace, which could be used as the output of MoE kernel to avoid extra copy. - Passing the returned tensor to combine lets the C++ op detect workspace ownership. - """ - self._require_mapped() - if self._state.phase != "dispatched": - raise RuntimeError( - "get_combine_payload_tensor_in_workspace called before a successful dispatch" - ) - - return torch.ops.trtllm.moe_a2a_get_combine_payload_tensor( - self.workspace, - self.ep_rank, - self.ep_size, - runtime_max_tokens_per_rank, - self._state.combine_payload_offset, - dtype, - hidden_size, - ) diff --git a/tensorrt_llm/_torch/moe/fused_moe/communication/nvlink_one_sided.py b/tensorrt_llm/_torch/moe/fused_moe/communication/nvlink_one_sided.py index 90ddc3405a65..3798e68098b6 100644 --- a/tensorrt_llm/_torch/moe/fused_moe/communication/nvlink_one_sided.py +++ b/tensorrt_llm/_torch/moe/fused_moe/communication/nvlink_one_sided.py @@ -25,12 +25,13 @@ """ import os +import re import sys from typing import Callable, Dict, List, Optional, Tuple +import pynvml import torch -from tensorrt_llm._mnnvl_utils import CftMnnvlMemory, MnnvlCheckpointCommunicator, MnnvlMemory from tensorrt_llm._torch.alltoall_watchdog import ( DEFAULT_ALLTOALL_WATCHDOG_POLL_INTERVAL_S, DEFAULT_ALLTOALL_WATCHDOG_TIMEOUT_S, @@ -41,6 +42,12 @@ EPGroupHealthLike, reject_rank_mask_cuda_graph_capture, ) +from tensorrt_llm._torch.distributed.mnnvl_memory import ( + CftMnnvlMemory, + MnnvlCheckpointCommunicator, + MnnvlMemory, + cuda, +) from tensorrt_llm._torch.mnnvl_alltoall_workspace import _MnnvlAlltoAllWorkspaceLifecycle from tensorrt_llm.bindings import internal as _tllm_internal from tensorrt_llm.logger import logger as tllm_logger @@ -50,13 +57,44 @@ from .base import Communication _CFT_DEFAULT_MAX_BATCH_FOR_DISPATCH = 128 -_CFT_MAX_BATCH_FOR_DISPATCH_ENV = "TRTLLM_MOE_A2A_CFT_MAX_BATCH_FOR_DISPATCH" -# CFT combine wins at small/medium batch and ties/regresses at large batch, so it is gated by the -# same per-call token-count threshold as dispatch. +_CFT_MAX_BATCH_FOR_DISPATCH_ENV = "TRTLLM_NVLINK_ONE_SIDED_A2A_CFT_MAX_BATCH_FOR_DISPATCH" _CFT_DEFAULT_MAX_BATCH_FOR_COMBINE = 128 -_CFT_MAX_BATCH_FOR_COMBINE_ENV = "TRTLLM_MOE_A2A_CFT_MAX_BATCH_FOR_COMBINE" -FORCE_CFT_ENV = "TRTLLM_MOE_A2A_FORCE_CFT" +_CFT_MAX_BATCH_FOR_COMBINE_ENV = "TRTLLM_NVLINK_ONE_SIDED_A2A_CFT_MAX_BATCH_FOR_COMBINE" +FORCE_CFT_ENV = "TRTLLM_NVLINK_ONE_SIDED_A2A_FORCE_CFT" _CFT_ALIGNMENT_BYTES = 16 +_CFT_MIN_DRIVER_BRANCH = 615 +_TIMEOUT_ENV = "TRTLLM_NVLINK_ONE_SIDED_A2A_TIMEOUT_SEC" +_WARMUP_TIMEOUT_ENV = "TRTLLM_NVLINK_ONE_SIDED_A2A_WARMUP_TIMEOUT_SEC" +_DEFAULT_TIMEOUT_SEC = 300 +_DEFAULT_WARMUP_TIMEOUT_SEC = 1800 +_MAX_TIMEOUT_SEC = 24 * 60 * 60 + + +def get_timeout_seconds(in_warmup: bool = False) -> int: + """Resolve the nominal collective timeout for the engine's execution phase. + + First-touch JIT compilation, autotuning and module loading can delay peer + ranks by minutes during warmup, so that phase gets a larger default budget. + """ + name = _WARMUP_TIMEOUT_ENV if in_warmup else _TIMEOUT_ENV + default = _DEFAULT_WARMUP_TIMEOUT_SEC if in_warmup else _DEFAULT_TIMEOUT_SEC + value = os.environ.get(name) + if not value: + return default + try: + if re.fullmatch(r"\s*[+-]?[0-9]+", value) is not None: + seconds = int(value) + if 0 < seconds <= _MAX_TIMEOUT_SEC: + return seconds + except ValueError: + # Very long integer strings may exceed Python's conversion limit. + pass + tllm_logger.warning_once( + f'Ignoring invalid {name}="{value}" (expected 1..{_MAX_TIMEOUT_SEC} seconds); ' + f"using {default} s", + key=f"{name}_invalid_{value}", + ) + return default def get_force_cft() -> bool | None: @@ -68,17 +106,87 @@ def get_force_cft() -> bool | None: return None -def resolve_can_use_cft(can_use_cft_counted_writes: bool) -> bool: - """Apply the TRTLLM_MOE_A2A_FORCE_CFT override to a caller's request. +def _get_nvidia_driver_version() -> str | None: + try: + try: + pynvml.nvmlDeviceGetCount() + except pynvml.NVMLError_Uninitialized: + pynvml.nvmlInit() + value = pynvml.nvmlSystemGetDriverVersion() + except pynvml.NVMLError as error: + tllm_logger.warning_once( + "CFT counted writes disabled: failed to query the NVIDIA driver " + f"version via NVML ({error}). Falling back to fence-based dispatch.", + key="moe_a2a_cft_driver_query_failed", + ) + return None + if isinstance(value, bytes): + return value.decode(errors="replace") + return str(value) - Workspace sizing and workspace layout both depend on this, so they must - resolve it identically: a caller that sizes without the override and then - constructs with it would lay out the CFT region in an undersized buffer. - """ - force_cft = get_force_cft() - if force_cft is None: - return can_use_cft_counted_writes - return force_cft + +def cft_driver_is_supported(driver_version: str | bytes | None) -> bool: + if isinstance(driver_version, bytes): + driver_version = driver_version.decode(errors="replace") + if not driver_version: + return False + match = re.match(r"^(\d+)(?:\.|$)", driver_version.strip()) + return bool(match and int(match.group(1)) >= _CFT_MIN_DRIVER_BRANCH) + + +def resolve_cft_counted_writes( + force_cft: bool | None, + driver_version: str | bytes | None, +) -> bool: + """Allow automatic or forced CFT only on a supported driver.""" + return force_cft is not False and cft_driver_is_supported(driver_version) + + +def _cft_device_support_reason() -> str | None: + """Return why the current device cannot use counted-write fabric endpoints.""" + major, minor = torch.cuda.get_device_capability() + if major < 10: + return f"SM{major}{minor} requires SM100 or newer" + try: + attributes = ( + cuda.CUdevice_attribute.CU_DEVICE_ATTRIBUTE_HANDLE_TYPE_FABRIC_SUPPORTED, + cuda.CUdevice_attribute.CU_DEVICE_ATTRIBUTE_LOGICAL_ENDPOINT_UNICAST_SUPPORTED, + cuda.CUdevice_attribute.CU_DEVICE_ATTRIBUTE_LOGICAL_ENDPOINT_COUNTED_OPS_SUPPORTED, + ) + except AttributeError: + return "CUDA Python bindings do not expose Logical Endpoint capabilities" + device = torch.cuda.current_device() + for attribute in attributes: + status, supported = cuda.cuDeviceGetAttribute(attribute, device) + if status != cuda.CUresult.CUDA_SUCCESS: + return f"querying {attribute.name} failed with {status.name}" + if not supported: + return f"device does not support {attribute.name}" + return None + + +def select_cft_counted_writes(force_cft: bool | None) -> bool: + """Resolve CFT availability identically for workspace sizing and construction.""" + driver_version = None + if force_cft is not False: + driver_version = _get_nvidia_driver_version() + if not resolve_cft_counted_writes(force_cft, driver_version): + if driver_version is not None: + tllm_logger.warning_once( + "CFT counted writes disabled: NVIDIA driver " + f"{driver_version} is below required {_CFT_MIN_DRIVER_BRANCH}.00. " + "Falling back to fence-based dispatch.", + key=f"moe_a2a_cft_driver_unsupported_{driver_version}", + ) + return False + unsupported_reason = _cft_device_support_reason() + if unsupported_reason is not None: + tllm_logger.warning_once( + f"CFT counted writes disabled: {unsupported_reason}. Falling back to fence.", + key=f"moe_a2a_cft_device_unsupported_{unsupported_reason}", + ) + return False + return True def should_use_cft( @@ -143,11 +251,11 @@ def _get_cft_max_batch(env_name: str, default: int) -> int: return threshold -def _get_cft_max_batch_for_dispatch() -> int | None: +def _get_cft_max_batch_for_dispatch() -> int: return _get_cft_max_batch(_CFT_MAX_BATCH_FOR_DISPATCH_ENV, _CFT_DEFAULT_MAX_BATCH_FOR_DISPATCH) -def _get_cft_max_batch_for_combine() -> int | None: +def _get_cft_max_batch_for_combine() -> int: return _get_cft_max_batch(_CFT_MAX_BATCH_FOR_COMBINE_ENV, _CFT_DEFAULT_MAX_BATCH_FOR_COMBINE) @@ -164,14 +272,32 @@ class NVLinkOneSided(Communication): """ # Constants from C++ (must match moeAlltoAllKernels.h) - MAX_RANKS = 256 - MAX_TOP_K = 8 - MAX_PAYLOADS = 8 + MAX_RANKS = int(_tllm_internal.thop.MOE_A2A_MAX_RANKS) + MAX_TOP_K = int(_tllm_internal.thop.MOE_A2A_MAX_TOP_K) + MAX_PAYLOADS = int(_tllm_internal.thop.MOE_A2A_MAX_PAYLOADS) # Shared workspaces/memory across the process, keyed by payload layout and CFT mode. _WORKSPACES: Dict[Tuple[object, ...], dict] = {} _WORKSPACE_REFCOUNTS: Dict[Tuple[object, ...], int] = {} _WORKSPACE: dict | None = None + _timeout_initialized = False + + @staticmethod + def set_timeout(timeout_sec: int) -> None: + """Set the process-wide timeout for subsequent one-sided A2A launches. + + Args: + timeout_sec: Integer in [1, 86400], nominal seconds at an assumed + 2 GHz SM clock. Applies to both dispatch/combine and CFT/fence. + + Existing CUDA graphs retain the budget recorded at capture time. + """ + if isinstance(timeout_sec, bool) or not isinstance(timeout_sec, int): + raise TypeError("timeout_sec must be an integer number of seconds") + if not 0 < timeout_sec <= _MAX_TIMEOUT_SEC: + raise ValueError(f"timeout_sec must be in 1..{_MAX_TIMEOUT_SEC} seconds") + torch.ops.trtllm.moe_a2a_set_timeout(timeout_sec) + NVLinkOneSided._timeout_initialized = True # MetaInfo indices - initialized from C++ constants FLAG_VAL_OFFSET_INDEX = None @@ -181,7 +307,61 @@ class NVLinkOneSided(Communication): DISPATCH_COMPLETION_FLAGS_OFFSET_INDEX = None COMBINE_COMPLETION_FLAGS_OFFSET_INDEX = None EPLB_GATHERED_STATS_OFFSET_INDEX = None - PAYLOAD_DATA_OFFSET_INDEX = None + DISPATCH_PAYLOAD_OFFSET_INDEX = None + DISPATCH_PAYLOAD_SIZE_INDEX = None + COMBINE_INPUT_OFFSET_INDEX = None + COMBINE_INPUT_SIZE_INDEX = None + COMBINE_RECV_OFFSET_INDEX = None + COMBINE_RECV_SIZE_INDEX = None + WORKSPACE_SIZE_INDEX = None + + @staticmethod + def _make_workspace_layout( + ep_size: int, + top_k: int, + max_num_tokens: int, + hidden_size: int, + dtype: torch.dtype, + eplb_stats_num_experts: Optional[int], + extra_payload_bytes_per_token: int, + can_use_cft_counted_writes: bool, + use_low_precision_combine: bool, + ) -> torch.Tensor: + """Plan per-rank control buffers and independently sized payload regions.""" + if hidden_size <= 0 or extra_payload_bytes_per_token < 0: + raise ValueError("hidden_size must be positive and extra payload size non-negative") + tokens = ep_size * max_num_tokens + # The wrapper accepts raw activations or quantized activations plus scales. + # A FP32 scale per 16 elements bounds the supported block-scale formats; + # accounting for its separate alignment also covers very small payloads. + activations = pad_up(tokens * hidden_size * dtype.itemsize, 128) + quantized = pad_up(tokens * hidden_size, 128) + pad_up( + tokens * ((hidden_size + 15) // 16) * 4, 128 + ) + dispatch_bytes = ( + max(activations, quantized) + + 2 * pad_up(tokens * top_k * 4, 128) + + pad_up(tokens * extra_payload_bytes_per_token, 128) + ) + # MoE may write its original-dtype output directly into the input region, + # even when the communication wire format is FP8. + combine_element_size = max(dtype.itemsize, 2) + combine_input_bytes = tokens * hidden_size * combine_element_size + combine_recv_bytes = ( + tokens * hidden_size * (1 if use_low_precision_combine else combine_element_size) + if can_use_cft_counted_writes + else 0 + ) + return torch.ops.trtllm.moe_a2a_get_workspace_layout( + ep_size, + max_num_tokens, + top_k, + dispatch_bytes, + combine_input_bytes, + combine_recv_bytes, + eplb_stats_num_experts, + can_use_cft_counted_writes, + ) @staticmethod def get_aux_data_size( @@ -189,10 +369,20 @@ def get_aux_data_size( max_num_tokens: int, eplb_stats_num_experts: Optional[int] = None, can_use_cft_counted_writes: bool = False, + top_k: Optional[int] = None, ) -> int: - return torch.ops.trtllm.moe_a2a_get_aux_data_size( - ep_size, max_num_tokens, eplb_stats_num_experts, can_use_cft_counted_writes + """Control-buffer bytes; omitted top_k reserves the native routing limit.""" + layout = torch.ops.trtllm.moe_a2a_get_workspace_layout( + ep_size, + max_num_tokens, + NVLinkOneSided.MAX_TOP_K if top_k is None else top_k, + 0, + 0, + 0, + eplb_stats_num_experts, + can_use_cft_counted_writes, ) + return int(layout[_tllm_internal.thop.MOE_A2A_WORKSPACE_SIZE_INDEX]) @staticmethod def calculate_required_workspace_size( @@ -203,41 +393,24 @@ def calculate_required_workspace_size( dtype: torch.dtype, eplb_stats_num_experts: Optional[int] = None, extra_payload_bytes_per_token: int = 0, - can_use_cft_counted_writes: bool = False, + can_use_cft_counted_writes: Optional[bool] = None, + use_low_precision_combine: bool = False, ) -> int: - can_use_cft_counted_writes = resolve_can_use_cft(can_use_cft_counted_writes) - element_size = dtype.itemsize - - # Auxiliary data size - workspace_size = NVLinkOneSided.get_aux_data_size( - ep_size, max_num_tokens, eplb_stats_num_experts, can_use_cft_counted_writes + # None sizes for what the constructor would select on this platform. + if can_use_cft_counted_writes is None: + can_use_cft_counted_writes = select_cft_counted_writes(get_force_cft()) + layout = NVLinkOneSided._make_workspace_layout( + ep_size, + top_k, + max_num_tokens, + hidden_size, + dtype, + eplb_stats_num_experts, + extra_payload_bytes_per_token, + can_use_cft_counted_writes, + use_low_precision_combine, ) - - # Dispatch needs workspace for [ep_size, max_tokens] tokens, - # but due to the variety of quantization recipes, we cannot know the exact size, so we conservatively estimate assuming no quantization. - # Meanwhile, we consider the alignment requirement as in moeA2ADispatchOp and moeA2ACombineOp. - # (Unquantized) token hidden states - workspace_size += ep_size * max_num_tokens * hidden_size * element_size - workspace_size = pad_up(workspace_size, 128) - # token_selected_experts - workspace_size += ep_size * max_num_tokens * top_k * 4 - workspace_size = pad_up(workspace_size, 128) - # token_final_scales - workspace_size += ep_size * max_num_tokens * top_k * 4 - workspace_size = pad_up(workspace_size, 128) - # Required workspace for combine [ep_size, max_tokens] tokens - workspace_size += ep_size * max_num_tokens * hidden_size * element_size - workspace_size = pad_up(workspace_size, 128) - # CFT combine: dedicated combine RECEIVE region C (peer pushes land here; - # prepareCombine never touches it -> no proxy aliasing). Same size as the combine region. - if can_use_cft_counted_writes: - workspace_size += ep_size * max_num_tokens * hidden_size * element_size - workspace_size = pad_up(workspace_size, 128) - # extra payload bytes per token - workspace_size += ep_size * max_num_tokens * extra_payload_bytes_per_token - workspace_size = pad_up(workspace_size, 128) - - return workspace_size + return int(layout[_tllm_internal.thop.MOE_A2A_WORKSPACE_SIZE_INDEX]) @classmethod def _init_constants(cls): @@ -259,7 +432,13 @@ def _init_constants(cls): cls.EPLB_GATHERED_STATS_OFFSET_INDEX = int( thop.MOE_A2A_EPLB_GATHERED_STATS_OFFSET_INDEX ) - cls.PAYLOAD_DATA_OFFSET_INDEX = int(thop.MOE_A2A_PAYLOAD_DATA_OFFSET_INDEX) + cls.DISPATCH_PAYLOAD_OFFSET_INDEX = int(thop.MOE_A2A_DISPATCH_PAYLOAD_OFFSET_INDEX) + cls.DISPATCH_PAYLOAD_SIZE_INDEX = int(thop.MOE_A2A_DISPATCH_PAYLOAD_SIZE_INDEX) + cls.COMBINE_INPUT_OFFSET_INDEX = int(thop.MOE_A2A_COMBINE_INPUT_OFFSET_INDEX) + cls.COMBINE_INPUT_SIZE_INDEX = int(thop.MOE_A2A_COMBINE_INPUT_SIZE_INDEX) + cls.COMBINE_RECV_OFFSET_INDEX = int(thop.MOE_A2A_COMBINE_RECV_OFFSET_INDEX) + cls.COMBINE_RECV_SIZE_INDEX = int(thop.MOE_A2A_COMBINE_RECV_SIZE_INDEX) + cls.WORKSPACE_SIZE_INDEX = int(thop.MOE_A2A_WORKSPACE_SIZE_INDEX) def __init__( self, @@ -272,7 +451,6 @@ def __init__( dtype: Optional[torch.dtype] = None, num_experts: Optional[int] = None, use_low_precision_combine: bool = False, - can_use_cft_counted_writes: bool = False, ep_group_health: EPGroupHealthLike | None = None, alltoall_watchdog_timeout_s: Optional[float] = None, alltoall_watchdog_poll_interval_s: float = DEFAULT_ALLTOALL_WATCHDOG_POLL_INTERVAL_S, @@ -281,6 +459,14 @@ def __init__( """ Initialize NVLinkOneSided with workspace allocation. + CFT is selected automatically on supported platforms using separate + dispatch/combine token-count thresholds (128 by default). + TRTLLM_NVLINK_ONE_SIDED_A2A_FORCE_CFT=0 selects fence; 1 bypasses the thresholds, + but not capability or payload-alignment requirements. CFT requires + sm_100+, a build against CUDA 13.4+, an NVLink fabric, and a driver + exporting the Logical Endpoint API (615.00+). Unsupported devices or + payloads fall back to fence with a one-time warning. + Args: mapping: TensorRT-LLM Mapping object containing rank information num_slots: Number of routing slots (token_selected_experts values are in [0, num_slots)). @@ -295,14 +481,6 @@ def __init__( use_low_precision_combine: If True, quantize the combine payload to FP8 for NVLink transfer (halves NVLink bandwidth usage, output precision is preserved). Corresponds to model_config.use_low_precision_moe_combine. - can_use_cft_counted_writes: If True, allow CFT handle-based counted - writes (fabric.try_put.counted via Logical Endpoints) for dispatch. - Requires sm_100+ (Blackwell or later), a build against CUDA 13.4+, an - NVLink fabric, and a driver exporting the CUDA logical endpoint API. - Defaults to False: CFT is opt-in, so the fence-based path remains the - default on every architecture. Callers that have verified the CFT - prerequisites may pass True, or set TRTLLM_MOE_A2A_FORCE_CFT=1 to force - CFT for supported workloads (0 forces the fence path). ep_group_health: Optional read-only committed EP membership. When present, rank-mask handling is enabled in the CUDA kernels, and its mask defines the peers expected by the watchdog. Timeout detection never mutates it. CUDA graphs are rejected until membership-scoped recapture lands. @@ -326,6 +504,11 @@ def __init__( f"NVLinkOneSided supports at most {self.MAX_RANKS} EP ranks, got ep_size={self.ep_size}." ) + # Standalone callers have no engine phase transitions. Initialize their + # steady-state budget without overwriting an explicit or warmup timeout. + if not NVLinkOneSided._timeout_initialized: + NVLinkOneSided.set_timeout(get_timeout_seconds()) + # Store needed parameters self.num_experts = num_slots self.top_k = top_k @@ -343,10 +526,7 @@ def __init__( self.enable_eplb = num_experts is not None self.eplb_stats_num_experts = num_experts self._force_cft = get_force_cft() - # Opt-in only: no caller passes can_use_cft_counted_writes=True, so - # without the override the CFT path cannot be reached at all. Leaving - # the variable unset keeps CFT disabled, as before. - can_use_cft_counted_writes = resolve_can_use_cft(can_use_cft_counted_writes) + can_use_cft_counted_writes = select_cft_counted_writes(self._force_cft) self.can_use_cft_counted_writes = can_use_cft_counted_writes if self._force_cft is None: self.cft_max_batch_for_dispatch = _get_cft_max_batch_for_dispatch() @@ -355,15 +535,13 @@ def __init__( self.cft_max_batch_for_dispatch = None self.cft_max_batch_for_combine = None if can_use_cft_counted_writes: - tllm_logger.info( - "NVLinkOneSided AlltoAll: CFT handle-based counted writes enabled for dispatch" - ) if self._force_cft is True: tllm_logger.info("NVLinkOneSided AlltoAll: CFT forced for supported workloads") - elif self.cft_max_batch_for_dispatch is not None: + else: tllm_logger.info( - "NVLinkOneSided AlltoAll: CFT dispatch disabled above " - f"runtime_max_tokens_per_rank={self.cft_max_batch_for_dispatch}" + "NVLinkOneSided AlltoAll: automatic CFT enabled with token-count limits " + f"dispatch={self.cft_max_batch_for_dispatch}, " + f"combine={self.cft_max_batch_for_combine}" ) else: tllm_logger.info( @@ -375,23 +553,31 @@ def __init__( # Get workspace size auto_workspace_size = None + metainfo = None if hidden_size is not None and dtype is not None: - auto_workspace_size = self.calculate_required_workspace_size( + metainfo = self._make_workspace_layout( self.ep_size, self.top_k, max_num_tokens_per_rank, hidden_size, dtype, - eplb_stats_num_experts=self.eplb_stats_num_experts, - can_use_cft_counted_writes=self.can_use_cft_counted_writes, + self.eplb_stats_num_experts, + 0, + self.can_use_cft_counted_writes, + self.use_low_precision_combine, ) - workspace_mb_env = os.environ.get("TRTLLM_MOE_A2A_WORKSPACE_MB") + auto_workspace_size = int(metainfo[self.WORKSPACE_SIZE_INDEX]) + workspace_mb_env = os.environ.get("TRTLLM_NVLINK_ONE_SIDED_A2A_WORKSPACE_MB") if workspace_mb_env: self.workspace_size_per_rank = int(workspace_mb_env) * 1024 * 1024 - msg = f"NVLinkOneSided: Forcing workspace size to {self.workspace_size_per_rank} bytes (TRTLLM_MOE_A2A_WORKSPACE_MB={workspace_mb_env})." + msg = f"NVLinkOneSided: Forcing workspace size to {self.workspace_size_per_rank} bytes (TRTLLM_NVLINK_ONE_SIDED_A2A_WORKSPACE_MB={workspace_mb_env})." if auto_workspace_size is not None: msg += f"Automatically calculated workspace size is {auto_workspace_size} bytes." - msg += "Auto calculation is conservative, so only consider overriding it if you have a specific reason." + if self.workspace_size_per_rank < auto_workspace_size: + raise ValueError( + f"Workspace override is too small: {self.workspace_size_per_rank} bytes, " + f"layout requires {auto_workspace_size} bytes per rank" + ) tllm_logger.warning(msg) elif auto_workspace_size is not None: self.workspace_size_per_rank = auto_workspace_size @@ -402,11 +588,35 @@ def __init__( ) self.workspace_size_per_rank = 2048 * 1024 * 1024 - # Initialize or reuse workspace. The C++ op computes payload offsets - # from the current tensors at dispatch time, while the Python singleton - # owns the symmetric memory backing those offsets. Keep separate - # workspaces for different payload layouts so one test/layer cannot - # reuse stale one-sided state from another shape. + if metainfo is None: + # Without a model shape, distribute the explicit/default byte budget + # conservatively. The resulting boundaries are still fixed at initialization. + alignment = int(_tllm_internal.thop.MOE_A2A_WORKSPACE_ALIGNMENT) + control_bytes = self.get_aux_data_size( + self.ep_size, + max_num_tokens_per_rank, + self.eplb_stats_num_experts, + self.can_use_cft_counted_writes, + self.top_k, + ) + regions = 3 if self.can_use_cft_counted_writes else 2 + capacity = ( + (self.workspace_size_per_rank - control_bytes) // regions // alignment * alignment + ) + if capacity <= 0: + raise ValueError("Workspace byte budget leaves no room for payloads") + metainfo = torch.ops.trtllm.moe_a2a_get_workspace_layout( + self.ep_size, + max_num_tokens_per_rank, + self.top_k, + capacity, + capacity, + capacity if self.can_use_cft_counted_writes else 0, + self.eplb_stats_num_experts, + self.can_use_cft_counted_writes, + ) + # Fixed region capacities are shared by allocation, native bounds checks, + # and workspace-backed views. Runtime shapes only pack tokens inside them. MnnvlMemory.initialize() self._workspace_key = ( self.workspace_size_per_rank, @@ -420,6 +630,7 @@ def __init__( dtype, self.use_low_precision_combine, self.can_use_cft_counted_writes, + tuple(metainfo.tolist()), ) workspace_state = NVLinkOneSided._WORKSPACES.get(self._workspace_key) @@ -432,14 +643,7 @@ def __init__( ) mnnvl_mem = memory_cls(mapping, self.workspace_size_per_rank) workspace = mnnvl_mem.as_torch_strided_tensor(torch.uint8) - metainfo = torch.ops.trtllm.moe_a2a_initialize( - workspace, - self.ep_rank, - self.ep_size, - self.max_num_tokens_per_rank, - self.eplb_stats_num_experts, - self.can_use_cft_counted_writes, - ) + torch.ops.trtllm.moe_a2a_initialize(workspace, metainfo, self.ep_rank, self.ep_size) workspace_state = { "workspace_size_per_rank": self.workspace_size_per_rank, "max_num_tokens_per_rank": self.max_num_tokens_per_rank, @@ -510,11 +714,6 @@ def __init__( ep_size=self.ep_size, health=self.ep_group_health, ) - # Keep combine storage at a fixed offset across changing dispatch - # payloads, so a later dispatch cannot overwrite a peer's live combine. - if hidden_size is not None and dtype is not None: - self._reserve_combine_region(hidden_size, dtype) - self._workspace_lifecycle.register( self, watchdog_timeout_s=alltoall_watchdog_timeout_s, @@ -579,7 +778,7 @@ def supports_post_quant_dispatch(self) -> bool: return True def destroy(self): - """Release shared state during explicit, rank-coordinated teardown.""" + """Release this instance's reference after all ranks finish using the workspace.""" if getattr(self, "_destroyed", False): return @@ -600,6 +799,8 @@ def destroy(self): if refcount > 0: NVLinkOneSided._WORKSPACE_REFCOUNTS[workspace_key] = refcount else: + if self._workspace_state.get("cft_initialized", False): + torch.ops.trtllm.moe_a2a_cft_destroy(self.workspace, self.ep_rank) NVLinkOneSided._WORKSPACE_REFCOUNTS.pop(workspace_key, None) workspace_state = NVLinkOneSided._WORKSPACES.pop(workspace_key, None) if NVLinkOneSided._WORKSPACE is workspace_state: @@ -689,17 +890,16 @@ def checkpoint_restore( comm = self.mnnvl_mem.comm if comm is None: raise RuntimeError("MNNVL workspace communicator is not initialized") - self._require_workspace_lifecycle().checkpoint_restore( - comm, - lambda: torch.ops.trtllm.moe_a2a_initialize( - self.workspace, - self.ep_rank, - self.ep_size, - self.max_num_tokens_per_rank, - self.eplb_stats_num_experts, - self.can_use_cft_counted_writes, - ), - ) + + # The native op re-validates the planned layout and resets control state; + # the layout itself is unchanged by a restore. + def reinitialize_frontend() -> torch.Tensor: + torch.ops.trtllm.moe_a2a_initialize( + self.workspace, self.moe_a2a_metainfo, self.ep_rank, self.ep_size + ) + return self.moe_a2a_metainfo + + self._require_workspace_lifecycle().checkpoint_restore(comm, reinitialize_frontend) def _mnnvl_checkpoint_is_idle(self) -> bool: return self._dispatch_state.get("phase") == "idle" @@ -707,49 +907,6 @@ def _mnnvl_checkpoint_is_idle(self) -> bool: def _mnnvl_checkpoint_reset(self) -> None: self._dispatch_state = {"phase": "idle"} - def _reserve_combine_region(self, hidden_size: int, dtype: torch.dtype) -> int: - """Keep dispatch and combine storage disjoint across workspace reuse.""" - layout = (hidden_size, dtype.itemsize) - old_layout = self._workspace_state.get("combine_storage_layout") - if old_layout is not None: - if old_layout != layout: - raise ValueError( - "shared A2A workspace requires a stable combine shape and dtype size" - ) - return self._workspace_state["combine_storage_offset"] - - # Native combine stores a dense [EP, runtime_max_tokens, hidden] tensor. - # Match the existing allocation: reserve a second full-size region only - # for CFT-capable workspaces. Low-precision combine fits in this bound. - region_bytes = pad_up( - self.ep_size * self.max_num_tokens_per_rank * hidden_size * dtype.itemsize, 128 - ) - num_regions = 2 if self.can_use_cft_counted_writes else 1 - offset = (self.workspace_size_per_rank - num_regions * region_bytes) // 128 * 128 - aux_bytes = int(self.moe_a2a_metainfo[self.PAYLOAD_DATA_OFFSET_INDEX]) - if offset < aux_bytes: - raise ValueError("A2A workspace is too small for stable combine regions") - dispatch_end = getattr(self, "_dispatch_state", {}).get("dispatch_payload_end", aux_bytes) - if dispatch_end > offset: - raise ValueError("A2A dispatch payload overlaps the reserved combine region") - self._workspace_state["combine_storage_layout"] = layout - self._workspace_state["combine_storage_offset"] = offset - return offset - - def _check_dispatch_region(self, payloads: List[torch.Tensor], max_tokens: int) -> int: - if not 0 < max_tokens <= self.max_num_tokens_per_rank: - raise ValueError("runtime token count exceeds the configured A2A capacity") - end = int(self.moe_a2a_metainfo[self.PAYLOAD_DATA_OFFSET_INDEX]) - for payload in payloads: - end = pad_up( - end + self.ep_size * max_tokens * payload.shape[1] * payload.element_size(), 128 - ) - limit = self._workspace_state.get("combine_storage_offset", self.workspace_size_per_rank) - # Check before the native dispatch can write into any peer's memory. - if end > limit: - raise ValueError("A2A dispatch payload overlaps the reserved combine region") - return end - def dispatch( self, hidden_states: torch.Tensor, @@ -800,7 +957,6 @@ def dispatch( payloads.append(token_selected_slots) if token_final_scales is not None: payloads.append(token_final_scales) - dispatch_payload_end = self._check_dispatch_region(payloads, runtime_max_tokens_per_rank) can_use_cft_for_dispatch = _use_cft_for_dispatch_payloads( can_use_cft_for_dispatch, payloads ) @@ -847,12 +1003,7 @@ def dispatch( if eplb_gathered_stats.numel() == 0: eplb_gathered_stats = None self._dispatch_state["eplb_gathered_stats"] = eplb_gathered_stats - if int(combine_payload_offset) != dispatch_payload_end: - raise RuntimeError("native A2A dispatch layout disagrees with the reserved layout") - self._dispatch_state["dispatch_payload_end"] = dispatch_payload_end - self._dispatch_state["combine_payload_offset"] = self._workspace_state.get( - "combine_storage_offset", dispatch_payload_end - ) + self._dispatch_state["combine_payload_offset"] = int(combine_payload_offset) self._dispatch_state["local_num_tokens"] = token_selected_slots.size(0) self._dispatch_state["runtime_max_tokens_per_rank"] = runtime_max_tokens_per_rank self._dispatch_state["active_rank_mask_snapshot"] = active_rank_mask_snapshot @@ -973,9 +1124,6 @@ def combine( final_hidden_states, self.use_low_precision_combine, ) - combine_payload_offset = self._reserve_combine_region( - final_hidden_states.shape[-1], final_hidden_states.dtype - ) output = torch.ops.trtllm.moe_a2a_combine( final_hidden_states, int(local_num_tokens), @@ -1025,7 +1173,7 @@ def get_combine_payload_tensor_in_workspace( dtype: Data type Returns: - Tensor view into workspace [ep_size, max_tokens_per_rank, hidden_size] + Tensor view into combine input [ep_size * runtime_max_tokens_per_rank, hidden_size] """ self._require_mapped() if self._dispatch_state.get("phase") != "dispatched": @@ -1037,14 +1185,12 @@ def get_combine_payload_tensor_in_workspace( if combine_payload_offset is None: raise RuntimeError("combine_payload_offset not found in dispatch state") - combine_payload_offset = self._reserve_combine_region(hidden_size, dtype) - self._dispatch_state["combine_payload_offset"] = combine_payload_offset result = torch.ops.trtllm.moe_a2a_get_combine_payload_tensor( self.workspace, + self.moe_a2a_metainfo, int(self.ep_rank), int(self.ep_size), int(runtime_max_tokens_per_rank), - int(combine_payload_offset), dtype, int(hidden_size), ) diff --git a/tensorrt_llm/_torch/moe/fused_moe/communication/nvlink_two_sided.py b/tensorrt_llm/_torch/moe/fused_moe/communication/nvlink_two_sided.py index 2582a965490f..fcda1f3042db 100644 --- a/tensorrt_llm/_torch/moe/fused_moe/communication/nvlink_two_sided.py +++ b/tensorrt_llm/_torch/moe/fused_moe/communication/nvlink_two_sided.py @@ -22,18 +22,305 @@ """ import os -from typing import List, Optional, Tuple +from dataclasses import dataclass +from typing import List, Optional, Tuple, Union from weakref import WeakSet import torch -from tensorrt_llm._mnnvl_utils import MnnvlCheckpointCommunicator, MnnvlMemory, MnnvlMoe +from tensorrt_llm._torch.distributed.mnnvl_memory import MnnvlCheckpointCommunicator, MnnvlMemory from tensorrt_llm._torch.mnnvl_alltoall_workspace import _collect_active_ranks from tensorrt_llm.mapping import Mapping from .base import Communication +@dataclass +class MoEAlltoallInfo: + local_gather_indices: torch.Tensor + send_rank_count_cumsum: torch.Tensor + send_rank_local_indices: torch.Tensor + recv_rank_count_cumsum: torch.Tensor + recv_rank_local_indices: torch.Tensor + backward_recv_rank_local_indices: torch.Tensor + local_token_allocation_count: int + + +class MnnvlMoe: + moe_workspace: MnnvlMemory = None + moe_prepare_workspace: MnnvlMemory = None + moe_workspace_tensor: torch.Tensor = None + moe_prepare_workspace_tensor: torch.Tensor = None + moe_mapping: Mapping = None + + @staticmethod + def get_moe_workspaces(mapping: Mapping): + if MnnvlMoe.moe_workspace is not None: + assert mapping == MnnvlMoe.moe_mapping, "only one moe mapping supported now" + return MnnvlMoe.moe_workspace_tensor + + MnnvlMoe.moe_mapping = mapping + workspace_size_per_rank = torch.ops.trtllm.get_moe_commworkspace_size_per_rank( + mapping.moe_ep_size + ) + MnnvlMoe.moe_workspace = MnnvlMemory(mapping, workspace_size_per_rank) + MnnvlMoe.moe_workspace_tensor = MnnvlMoe.moe_workspace.as_torch_strided_tensor(torch.uint64) + torch.ops.trtllm.moe_initialize_workspace( + MnnvlMoe.moe_workspace_tensor, mapping.moe_ep_rank, mapping.moe_ep_size + ) + torch.cuda.synchronize() + MnnvlMoe.moe_workspace.comm.barrier() + return MnnvlMoe.moe_workspace_tensor + + @staticmethod + def get_moe_prepare_workspace(mapping: Mapping): + if MnnvlMoe.moe_prepare_workspace_tensor is not None: + assert mapping == MnnvlMoe.moe_mapping, "only one moe mapping supported now" + return MnnvlMoe.moe_prepare_workspace_tensor + workspace_size_per_rank = torch.ops.trtllm.get_moe_prepare_workspace_size_per_rank( + mapping.moe_ep_size + ) + MnnvlMoe.moe_prepare_workspace = MnnvlMemory(mapping, workspace_size_per_rank) + MnnvlMoe.moe_prepare_workspace_tensor = ( + MnnvlMoe.moe_prepare_workspace.as_torch_strided_tensor(torch.uint64) + ) + return MnnvlMoe.moe_prepare_workspace_tensor + + @staticmethod + def compute_target_rank_id( + token_selected_experts: torch.Tensor, expert_count: int, ep_size: int + ): + assert expert_count % ep_size == 0, "expert_count should be divisible by ep_size" + expert_per_rank = expert_count // ep_size + token_target_rank_ids = token_selected_experts // expert_per_rank + return token_target_rank_ids + + @staticmethod + def mnnvl_moe_alltoallv_prepare_without_allgather( + expert_ids: torch.Tensor, + expert_statics: Optional[torch.Tensor], + workspace: torch.Tensor, + max_token_count_per_rank: int, + ep_rank: int, + ep_size: int, + expert_count: int, + slot_count: int, + top_k: int, + ): + ( + local_send_rank_count_cumsum, + local_send_rank_indices, + local_recv_rank_count_cumsum, + local_recv_rank_indices, + backward_local_recv_rank_indices, + gathered_expert_statics, + ) = torch.ops.trtllm.mnnvl_moe_alltoallv_prepare_without_allgather( + expert_ids, + expert_statics, + workspace, + max_token_count_per_rank, + ep_rank, + ep_size, + expert_count, + slot_count, + top_k, + ) + + local_token_allocation_count = max_token_count_per_rank * ep_size + # Looks like we don't need this. + local_gather_indices = None + + alltoall_info = MoEAlltoallInfo( + local_gather_indices, + local_send_rank_count_cumsum, + local_send_rank_indices, + local_recv_rank_count_cumsum, + local_recv_rank_indices, + backward_local_recv_rank_indices, + local_token_allocation_count, + ) + + return alltoall_info, gathered_expert_statics + + @staticmethod + def mnnvl_moe_expert_static_allgather( + expert_ids: torch.Tensor, + workspace: torch.Tensor, + ep_rank: int, + ep_size: int, + expert_count: int, + ): + gathered_expert_ids = torch.ops.trtllm.mnnvl_moe_expert_static_allgather( + expert_ids, workspace, ep_rank, ep_size, expert_count + ) + return gathered_expert_ids + + @staticmethod + def mnnvl_moe_alltoallv_prepare( + gathered_target_rank_ids: torch.Tensor, + real_rank_token_count_cumsum: Optional[torch.Tensor], + gathered_expert_ids: torch.Tensor, + gathered_scales: Optional[torch.Tensor], + max_token_count_per_rank: int, + expert_count: int, + top_k: int, + ep_rank: int, + ep_size: int, + ): + ( + local_gather_indices, + send_rank_count_cumsum, + send_rank_local_indices, + recv_rank_count_cumsum, + recv_rank_local_indices, + backward_recv_rank_local_indices, + ) = torch.ops.trtllm.moe_comm_prepare_indices( + gathered_target_rank_ids, + real_rank_token_count_cumsum, + max_token_count_per_rank, + expert_count, + top_k, + ep_rank, + ep_size, + ) + + local_token_allocation_count = max_token_count_per_rank * ep_size + + local_expert_ids = torch.empty( + local_token_allocation_count, top_k, dtype=torch.int32, device=torch.device("cuda") + ) + if gathered_scales is None: + local_scales = None + else: + local_scales = torch.empty( + local_token_allocation_count, + top_k, + dtype=torch.float32, + device=torch.device("cuda"), + ) + + torch.ops.trtllm.moe_local_gather( + recv_rank_count_cumsum, + local_gather_indices, + gathered_expert_ids, + gathered_scales, + local_expert_ids, + local_scales, + max_token_count_per_rank, + expert_count, + top_k, + ep_rank, + ep_size, + ) + + alltoall_info = MoEAlltoallInfo( + local_gather_indices, + send_rank_count_cumsum, + send_rank_local_indices, + recv_rank_count_cumsum, + recv_rank_local_indices, + backward_recv_rank_local_indices, + local_token_allocation_count, + ) + return alltoall_info, local_expert_ids, local_scales + + @staticmethod + def mnnvl_moe_alltoallv( + x: Union[torch.Tensor, List[Optional[torch.Tensor]]], + alltoall_info: MoEAlltoallInfo, + workspace: torch.Tensor, + ep_rank: int, + ep_size: int, + ) -> Union[torch.Tensor, List[Optional[torch.Tensor]]]: + # Convert single tensor to list for unified handling + is_single_tensor = not isinstance(x, list) + if is_single_tensor: + assert x.dim() == 2, "only 2D tensor supported, please reshape." + x = [x] + + assert len(x) > 0, "Empty tensor list not supported" + + # Filter out None values + valid_list = [tensor is not None for tensor in x] + valid_tensors = [tensor for tensor in x if tensor is not None] + + if len(valid_tensors) == 0: + # All tensors are None, return list of None + result = [None] * len(x) + else: + first_dim = None + for tensor in valid_tensors: + # Validate dimensions of valid tensors + assert tensor.dim() == 2, "only 2D tensor supported, please reshape." + if first_dim is None: + first_dim = tensor.shape[0] + else: + assert tensor.shape[0] == first_dim, ( + f"All tensors must have the same first dimension, got {tensor.shape[0]} vs {first_dim}" + ) + + # Process only valid tensors + output_tensors = torch.ops.trtllm.moe_comm( + valid_tensors, + alltoall_info.send_rank_count_cumsum, + alltoall_info.send_rank_local_indices, + alltoall_info.recv_rank_count_cumsum, + alltoall_info.recv_rank_local_indices, + workspace, + alltoall_info.local_token_allocation_count, + ep_rank, + ep_size, + ) + + # Restore None positions in output + idx = 0 + result = [] + for is_valid in valid_list: + if is_valid: + result.append(output_tensors[idx]) + idx += 1 + else: + result.append(None) + + # If input was a single tensor, return a single tensor + if is_single_tensor: + result = result[0] + + return result + + @staticmethod + def mnnvl_moe_alltoallv_combine( + x: torch.Tensor, + alltoall_info: MoEAlltoallInfo, + workspace: torch.Tensor, + ep_rank: int, + ep_size: int, + top_k: int, + token_count: int, + use_low_precision_combine: bool = False, + do_reduce: bool = True, + ): + assert x.dim() == 2, "2D tensor supported, please reshape." + output_tensors = torch.ops.trtllm.moe_comm( + [x], + alltoall_info.recv_rank_count_cumsum, + alltoall_info.recv_rank_local_indices, + alltoall_info.send_rank_count_cumsum, + alltoall_info.backward_recv_rank_local_indices, + workspace, + token_count * top_k, + ep_rank, + ep_size, + [True], + use_low_precision_combine, + ) + output_tensor = output_tensors[0].reshape(token_count, top_k, x.shape[1]) + if do_reduce: + return torch.sum(output_tensor, dim=1, keepdim=False) + else: + return output_tensor + + class NVLinkTwoSided(Communication): """ NVLINK two-sided comm AllToAll strategy. diff --git a/tensorrt_llm/_torch/moe/fused_moe/moe_op_backend.py b/tensorrt_llm/_torch/moe/fused_moe/moe_op_backend.py index 19ecceae36a4..61a033f172e3 100644 --- a/tensorrt_llm/_torch/moe/fused_moe/moe_op_backend.py +++ b/tensorrt_llm/_torch/moe/fused_moe/moe_op_backend.py @@ -210,7 +210,9 @@ class TRTLLMOpBackend(MoEOpBackend): """TRTLLM native op backend implementation.""" def __init__(self): - from tensorrt_llm._mnnvl_utils import MnnvlMemory, MnnvlMoe + from tensorrt_llm._torch.distributed.mnnvl_memory import MnnvlMemory + + from .communication.nvlink_two_sided import MnnvlMoe self._MnnvlMemory = MnnvlMemory self._MnnvlMoe = MnnvlMoe diff --git a/tensorrt_llm/_torch/pyexecutor/model_engine.py b/tensorrt_llm/_torch/pyexecutor/model_engine.py index 4fc771e1ce64..98ba2058376c 100644 --- a/tensorrt_llm/_torch/pyexecutor/model_engine.py +++ b/tensorrt_llm/_torch/pyexecutor/model_engine.py @@ -356,12 +356,19 @@ def _set_moe_a2a_warmup(in_warmup: bool) -> None: No-op when the op is unavailable (older bindings). """ + from ..moe.fused_moe.communication.nvlink_one_sided import ( + NVLinkOneSided, get_timeout_seconds) + + timeout_sec = get_timeout_seconds(in_warmup) try: - torch.ops.trtllm.moe_a2a_set_warmup(in_warmup) - logger.info(f"moe_a2a completion-flag budget: in_warmup={in_warmup}") + NVLinkOneSided.set_timeout(timeout_sec) + logger.info( + f"moe_a2a completion-flag budget: in_warmup={in_warmup}, " + f"timeout={timeout_sec} s (nominal, at an assumed 2 GHz clock64 rate)" + ) except (AttributeError, RuntimeError) as e: logger.warning( - f"moe_a2a_set_warmup unavailable, the all-to-all timeout " + f"moe_a2a_set_timeout unavailable, the all-to-all timeout " f"budget was not switched: {type(e).__name__}: {e}") @@ -1778,7 +1785,7 @@ def _prewarm_cute_dsl_indexer_q(self) -> None: completion-flag deadline. It is a partial mitigation only: other first-touch compiles remain inside collective-bearing forwards, and some sit on the all-to-all path itself and cannot be pre-compiled this way. - The runtime budget (``moeA2AGetTimeoutCycles``) covers the general case. + The phase-specific all-to-all timeout covers the general case. Only the fallback tactics are compiled -- what an eager, cache-miss forward selects. The runner's kernel cache key excludes m/n/k, so one @@ -2100,7 +2107,7 @@ def _reset_moe_alltoall_state(self) -> None: """Reset all MoE all-to-all state machines reachable from ``self.model``. Each MoE backend keeps a small dispatch/combine phase state per layer - (``MoeAlltoAll`` or ``NVLinkOneSided``). A forward that calls + (``NVLinkOneSided``). A forward that calls ``dispatch`` but raises before reaching ``combine`` (e.g., a warmup OOM mid-MoE) leaves that state in ``dispatched``, which fails the invariant on the next ``dispatch`` call. This helper walks the model diff --git a/tests/integration/test_lists/test-db/l0_dgx_b200.yml b/tests/integration/test_lists/test-db/l0_dgx_b200.yml index e0f01fa1ce3d..aa4c40bb6a11 100644 --- a/tests/integration/test_lists/test-db/l0_dgx_b200.yml +++ b/tests/integration/test_lists/test-db/l0_dgx_b200.yml @@ -207,6 +207,8 @@ l0_dgx_b200: # - accuracy/test_llm_api_pytorch.py::TestDeepSeekV4FlashBase::test_fp8_4gpus_static_eplb[moe_backend=DEEPGEMM] TIMEOUT (120) - accuracy/test_disaggregated_serving.py::TestNemotron3Super120B::test_ctx_dp2_gen_tp4 TIMEOUT (60) - accuracy/test_disaggregated_serving.py::TestQwen3NextInstruct::test_auto_dtype[use_py_transceiver=True] TIMEOUT (60) + # ------------- MoE communication unit tests (multi-GPU) --------------- + - unittest/_torch/moe/multi_gpu/test_nvlink_one_sided.py TIMEOUT (30) # ------------- VisualGen multi-GPU tests --------------- - unittest/_torch/visual_gen/multi_gpu/test_attn2d_attention.py - unittest/_torch/visual_gen/multi_gpu/test_cosmos3_transformer_parallel.py diff --git a/tests/microbenchmarks/bench_moe/search.py b/tests/microbenchmarks/bench_moe/search.py index 8cfa47369f9a..770cea4c2617 100644 --- a/tests/microbenchmarks/bench_moe/search.py +++ b/tests/microbenchmarks/bench_moe/search.py @@ -29,7 +29,7 @@ except ImportError: from cuda import cuda -from tensorrt_llm._mnnvl_utils import MnnvlMemory +from tensorrt_llm._torch.distributed.mnnvl_memory import MnnvlMemory from tensorrt_llm._torch.moe.fused_moe.impl_contract import ( MoEDeployment, MoEProblem, diff --git a/tests/microbenchmarks/bench_moe_comm.py b/tests/microbenchmarks/bench_moe_comm.py index b2ba7fcbcd30..f19db4b28601 100644 --- a/tests/microbenchmarks/bench_moe_comm.py +++ b/tests/microbenchmarks/bench_moe_comm.py @@ -21,6 +21,11 @@ - Communication.dispatch() - Communication.combine() +dispatch_us and combine_us report CUPTI kernel spans, using CUDA graph replay by +default or eager execution with --no_cuda_graph. If CUPTI is unavailable, timing +falls back to CUDA events and benchmark_metadata.warning records the reason; +otherwise it is null. --kernel_breakdown additionally prints per-kernel statistics. + Launch (examples): ```bash @@ -60,7 +65,6 @@ import torch from mpi4py import MPI from mpi4py.futures import MPIPoolExecutor -from torch.autograd import DeviceType import tensorrt_llm as tllm from tensorrt_llm._torch.model_config import ModelConfig @@ -90,20 +94,51 @@ class Profile: PROFILES: Dict[str, Profile] = { + "gpt_oss": Profile( + name="gpt_oss", + hidden_size=2880, + top_k=4, + num_experts=128, + quant_algo=QuantAlgo.W4A8_MXFP4_MXFP8, + ), "deepseek_v3": Profile( name="deepseek_v3", hidden_size=7168, top_k=8, num_experts=256, + # Notice: Cutlass quantize_input() is a no-op for FP8_BLOCK_SCALES: dispatch + # carries BF16 activations, not post-quantized FP8 payloads. quant_algo=QuantAlgo.FP8_BLOCK_SCALES, ), - "gpt_oss": Profile( - name="gpt_oss", - hidden_size=2880, - top_k=4, - num_experts=128, + "deepseek_v4_flash": Profile( + name="deepseek_v4_flash", + hidden_size=4096, + top_k=6, + num_experts=256, + quant_algo=QuantAlgo.W4A8_MXFP4_MXFP8, + ), + "deepseek_v4_pro": Profile( + name="deepseek_v4_pro", + hidden_size=7168, + top_k=6, + num_experts=384, quant_algo=QuantAlgo.W4A8_MXFP4_MXFP8, ), + "kimi_k3": Profile( + name="kimi_k3", + # All-to-all exchanges latent MoE activations, not the model's 7168-wide states. + hidden_size=3584, + top_k=16, + num_experts=896, + quant_algo=QuantAlgo.W4A8_MXFP4_MXFP8, + ), + "qwen3p8_2p4t_a95b": Profile( + name="qwen3p8_2p4t_a95b", + hidden_size=8192, + top_k=10, + num_experts=512, + quant_algo=QuantAlgo.NO_QUANT, + ), } @@ -211,7 +246,7 @@ def _create_model_config( act_dtype: torch.dtype, max_num_tokens_per_rank: int, quant_config: Optional[QuantConfig], - use_low_precision_moe_combine: bool = False, + use_low_precision_combine: bool = False, ) -> ModelConfig: # Keep it minimal: just enough fields for CommunicationFactory. return ModelConfig( @@ -221,242 +256,8 @@ def _create_model_config( max_num_tokens=int(max_num_tokens_per_rank), moe_max_num_tokens=int(max_num_tokens_per_rank), use_cuda_graph=False, - use_low_precision_moe_combine=use_low_precision_moe_combine, - ) - - -def _time_dispatch_and_combine( - backend: Communication, - *, - hidden_states: torch.Tensor, - hidden_states_sf: Optional[torch.Tensor], - token_selected_slots: torch.Tensor, - token_final_scales: Optional[torch.Tensor], - all_rank_num_tokens: List[int], - hidden_size: int, - warmup: int, - iters: int, - flush_l2: bool = True, -) -> Tuple[List[float], List[float], Dict[str, Any]]: - """Time dispatch and combine using Kineto (torch.profiler with CUPTI). - - Returns: - dispatch_times_us: Per-iteration dispatch GPU times in microseconds - combine_times_us: Per-iteration combine GPU times in microseconds - detailed_stats: Dict containing per-kernel timing breakdown - """ - device = hidden_states.device - - # L2 cache flushing buffer - l2_buffer = None - if flush_l2: - l2_size = torch.cuda.get_device_properties(device).L2_cache_size - # Use 2x L2 size to ensure complete flush - l2_flush_size = (l2_size * 2) // 4 # Size in int32 elements - l2_buffer = torch.empty(l2_flush_size, dtype=torch.int32, device=device) - - # Profile with Kineto - with torch.profiler.profile( - # Include CPU so `record_function("dispatch"/"combine")` ranges appear in - # key_averages() / events(). Without CPU activity those ranges are missing, - # causing dispatch/combine attribution to fail. - activities=[torch.profiler.ProfilerActivity.CUDA, torch.profiler.ProfilerActivity.CPU], - record_shapes=False, - with_stack=False, - ) as prof: - _sync() - - # Warmup iterations (not profiled) - for _ in range(warmup): - if l2_buffer is not None: - l2_buffer.zero_() - backend.prepare_dispatch( - token_selected_slots, all_rank_num_tokens - ) # For most ranks this is no-op except for NVLINK_TWO_SIDED - recv_hidden_states, _, _, _ = backend.dispatch( - hidden_states, - hidden_states_sf, - token_selected_slots, - token_final_scales, - all_rank_num_tokens, - ) - shape = list(recv_hidden_states.shape) - shape[-1] = hidden_size - recv_hidden_states_moe = torch.empty( - tuple(shape), dtype=torch.bfloat16, device=recv_hidden_states.device - ) - _ = backend.combine( - recv_hidden_states_moe, all_rank_max_num_tokens=max(all_rank_num_tokens) - ) - - # Timed iterations - for _ in range(iters): - # L2 cache flushing - if l2_buffer is not None: - l2_buffer.zero_() - - # Mark dispatch operation for aggregated timing - with torch.profiler.record_function("dispatch"): - backend.prepare_dispatch( - token_selected_slots, all_rank_num_tokens - ) # For most ranks this is no-op except for NVLINK_TWO_SIDED - recv_hidden_states, _, _, _ = backend.dispatch( - hidden_states, - hidden_states_sf, - token_selected_slots, - token_final_scales, - all_rank_num_tokens, - ) - - # Simulate MoE computation output - shape = list(recv_hidden_states.shape) - shape[-1] = hidden_size - recv_hidden_states_moe = torch.empty( - tuple(shape), dtype=torch.bfloat16, device=recv_hidden_states.device - ) - - # Mark combine operation for aggregated timing - with torch.profiler.record_function("combine"): - _ = backend.combine( - recv_hidden_states_moe, all_rank_max_num_tokens=max(all_rank_num_tokens) - ) - - _sync() - # if mpi_rank() == 0: - # print("########################################################") - # print(prof.key_averages()) - # print("########################################################") - return _parse_profiler_events(list(prof.events())) - - -def _parse_profiler_events( - events_list: list, -) -> Tuple[List[float], List[float], Dict[str, Any]]: - """Parse Kineto profiler events into per-iteration times and kernel breakdown. - - Expects the profiler to have been run with record_function("dispatch") and - record_function("combine") wrapping each operation (works for both eager - kernels and CUDA graph replays). - """ - # if mpi_rank() == 0: - # print("++++++++++++++++++++++++++++++++++++++++++++++++++++++++") - # for evt in events_list: - # print(evt) - # print("++++++++++++++++++++++++++++++++++++++++++++++++++++++++") - - def _is_gpu_event(evt) -> bool: - return getattr(evt, "device_type", None) == DeviceType.CUDA - - # Step 1: Collect GPU time ranges of "dispatch"/"combine" CUDA ranges. - gpu_dispatch_intervals: List[Tuple[int, int]] = [] - gpu_combine_intervals: List[Tuple[int, int]] = [] - - for evt in events_list: - if not _is_gpu_event(evt) or evt.name not in ("dispatch", "combine"): - continue - tr = getattr(evt, "time_range", None) - if tr is None: - continue - assert tr.end > tr.start - (gpu_dispatch_intervals if evt.name == "dispatch" else gpu_combine_intervals).append( - (tr.start, tr.end) - ) - gpu_dispatch_intervals.sort() - gpu_combine_intervals.sort() - - # Step 2: Scope resolver (GPU events only) --------------------------------- - def _find_scope(evt) -> Optional[str]: - """Return scope only when kernel range is strictly contained.""" - tr = getattr(evt, "time_range", None) - if tr is None: - return None - - # Be careful: Due to PDL, the end of dispatch and the start of combine may overlap, - # so we say a kernel is in dispatch/combine only if its range is strictly contained in a dispatch/combine range. - in_dispatch = any(s <= tr.start and tr.end <= e for s, e in gpu_dispatch_intervals) - in_combine = any(s <= tr.start and tr.end <= e for s, e in gpu_combine_intervals) - - assert not (in_dispatch and in_combine), ( - f"Kernel range is simultaneously inside dispatch and combine ranges: {evt.name}" - ) - - if in_dispatch: - return "dispatch" - if in_combine: - return "combine" - - # Neither in dispatch or combine (like the element-wise kernel for L2 cache flushing) -> uncategorized. - return None - - # Step 3: Iterate events and bucket by scope ---------------------------- - dispatch_kernel_times: Dict[str, List[float]] = {} - combine_kernel_times: Dict[str, List[float]] = {} - other_kernel_times: Dict[str, List[float]] = {} - - for evt in events_list: - if not _is_gpu_event(evt): - continue - if evt.device_time <= 0: - continue - if evt.name in ("dispatch", "combine"): - continue # skip record_function range markers - - scope = _find_scope(evt) - if scope == "dispatch": - dispatch_kernel_times.setdefault(evt.name, []).append(evt.device_time) - elif scope == "combine": - combine_kernel_times.setdefault(evt.name, []).append(evt.device_time) - else: - other_kernel_times.setdefault(evt.name, []).append(evt.device_time) - - # Step 4: Build per-kernel stats ---------------------------------------- - def _build_kernel_list(kernel_times: Dict[str, List[float]]) -> List[Dict[str, Any]]: - result = [] - for name, times in kernel_times.items(): - result.append( - { - "name": name, - "count": len(times), - "_times": times, # raw per-iteration times, gathered across ranks later - } - ) - return result - - dispatch_kernels = _build_kernel_list(dispatch_kernel_times) - combine_kernels = _build_kernel_list(combine_kernel_times) - other_kernels = _build_kernel_list(other_kernel_times) - - # Step 5: Collect per-iteration dispatch/combine times (us) --------------- - # Use the CUDA-side "dispatch"/"combine" range events (device_type=CUDA) - # for direct GPU time measurement. - dispatch_times_us: List[float] = [] - combine_times_us: List[float] = [] - for evt in events_list: - if not _is_gpu_event(evt): - continue - if evt.name == "dispatch": - dispatch_times_us.append(evt.device_time) - elif evt.name == "combine": - combine_times_us.append(evt.device_time) - - # Sort each category by mean time descending - dispatch_kernels.sort( - key=lambda x: sum(x["_times"]) / len(x["_times"]) if x["_times"] else 0, reverse=True + use_low_precision_moe_combine=use_low_precision_combine, ) - combine_kernels.sort( - key=lambda x: sum(x["_times"]) / len(x["_times"]) if x["_times"] else 0, reverse=True - ) - other_kernels.sort( - key=lambda x: sum(x["_times"]) / len(x["_times"]) if x["_times"] else 0, reverse=True - ) - - detailed_stats = { - "dispatch_kernels": dispatch_kernels, - "combine_kernels": combine_kernels, - "other_kernels": other_kernels, - } - - return dispatch_times_us, combine_times_us, detailed_stats def _demangle_names(names: List[str]) -> Dict[str, str]: @@ -469,153 +270,188 @@ def _demangle_names(names: List[str]) -> Dict[str, str]: return {n: n for n in names} -def _build_cuda_graph_kernel_stats_cupti( - cupti_kernels: List[Tuple[str, int, int]], # (name, start_ns, end_ns) - cupti_events: List[int], # device_timestamps of EXTERNAL events, sorted - iters: int, -) -> Optional[Dict[str, Any]]: - """Categorize GPU kernels from a CUDA graph replay into dispatch/combine/other. - - Uses CUPTI kernel timestamps and CUPTI CUDA_EVENT device_timestamps, all in the - same GPU nanosecond clock domain. - - The graph records 4 EXTERNAL events per timed iteration (no events during warmup): - event 4*i+0 → d_starts[i], 4*i+1 → d_ends[i] - event 4*i+2 → c_starts[i], 4*i+3 → c_ends[i] - - Each kernel is classified by whether its (k_start, k_end) falls within a - dispatch or combine window; everything else (including warmup kernels) is other. - - Returns None if CUPTI events are missing. - - The returned dict includes: - dispatch_times_us / combine_times_us: per-iter kernel-span times (ns → µs), - computed as (last_kernel_end − first_kernel_start) within each window. - None for iterations where no kernels were attributed (caller should fall back - to CUDA-event elapsed_time for those iterations). - """ - expected_events = 4 * iters - if len(cupti_events) != expected_events: - _maybe_warn_rank0( - f"[bench] CUPTI kernel breakdown skipped: expected {expected_events} CUDA_EVENT " - f"records ({iters} iters × 4) but got {len(cupti_events)}. " - "This usually means _try_init_cupti() was called after CUDA context creation." +def _build_kernel_stats_cupti( + cupti_kernels: List[Tuple[str, int, int]], + cupti_events: list[tuple[int, int]], + phase_event_ids: list[tuple[int, int, int, int]], +) -> Dict[str, Any]: + """Attribute kernels to event windows and measure each iteration's kernel span.""" + expected_ids = {event_id for iteration in phase_event_ids for event_id in iteration} + if len(expected_ids) != 4 * len(phase_event_ids): + raise RuntimeError("Each benchmark timing event must have a distinct CUPTI event ID.") + event_timestamps: dict[int, int] = {} + for event_id, timestamp in cupti_events: + if event_id not in expected_ids: + continue + if timestamp <= 0 or event_id in event_timestamps: + raise RuntimeError(f"CUPTI returned an invalid or duplicate timing event: {event_id}") + event_timestamps[event_id] = timestamp + missing_ids = expected_ids - event_timestamps.keys() + if missing_ids: + raise RuntimeError( + f"CUPTI is missing {len(missing_ids)} of {len(expected_ids)} timing events. " + "CUDA_EVENT tracking must be enabled before CUDA context creation." ) - return None if not cupti_kernels: - return None + raise RuntimeError("CUPTI captured no kernels for the timed run.") - d_starts_abs = [cupti_events[4 * i + 0] for i in range(iters)] - d_ends_abs = [cupti_events[4 * i + 1] for i in range(iters)] - c_starts_abs = [cupti_events[4 * i + 2] for i in range(iters)] - c_ends_abs = [cupti_events[4 * i + 3] for i in range(iters)] - - unique_names = list({name for name, _, _ in cupti_kernels}) - dm = _demangle_names(unique_names) - - dispatch_kernel_times: Dict[str, List[float]] = {} - combine_kernel_times: Dict[str, List[float]] = {} - other_kernel_times: Dict[str, List[float]] = {} + phase_windows = [] + for d_start_id, d_end_id, c_start_id, c_end_id in phase_event_ids: + d_start, d_end, c_start, c_end = ( + event_timestamps[event_id] for event_id in (d_start_id, d_end_id, c_start_id, c_end_id) + ) + if not d_start <= d_end <= c_start <= c_end: + raise RuntimeError("CUPTI timing events are not in dispatch/combine execution order.") + if phase_windows and d_start < phase_windows[-1][3]: + raise RuntimeError("CUPTI iteration timing windows overlap.") + phase_windows.append((d_start, d_end, c_start, c_end)) + + phase_bounds: dict[str, list[tuple[int, int] | None]] = { + phase: [None] * len(phase_windows) for phase in ("dispatch", "combine") + } - # Per-iteration [first_start_ns, last_end_ns] for kernel-span timing. - dispatch_iter_span: List[List[Optional[int]]] = [[None, None] for _ in range(iters)] - combine_iter_span: List[List[Optional[int]]] = [[None, None] for _ in range(iters)] + cupti_kernels.sort(key=lambda kernel: kernel[1]) - for name, k_start, k_end in cupti_kernels: - demangled = dm.get(name, name) - device_time_us = (k_end - k_start) / 1e3 # ns → µs + unique_names = list({name for name, _, _ in cupti_kernels}) + demangled_names = _demangle_names(unique_names) + kernel_times: Dict[str, Dict[str, List[float]]] = { + "dispatch": {}, + "combine": {}, + "other": {}, + } + for name, kernel_start, kernel_end in cupti_kernels: + if kernel_start <= 0 or kernel_end <= kernel_start: + raise RuntimeError(f"CUPTI returned invalid kernel timestamps for {name}.") category = "other" - iter_idx = -1 - for i in range(iters): - if k_start >= d_starts_abs[i] and k_end <= d_ends_abs[i]: - category = "dispatch" - iter_idx = i - break - if k_start >= c_starts_abs[i] and k_end <= c_ends_abs[i]: - category = "combine" - iter_idx = i + for iteration, (d_start, d_end, c_start, c_end) in enumerate(phase_windows): + for phase, start, end in (("dispatch", d_start, d_end), ("combine", c_start, c_end)): + if kernel_start >= start and kernel_end <= end: + category = phase + bounds = phase_bounds[phase][iteration] + phase_bounds[phase][iteration] = ( + (min(bounds[0], kernel_start), max(bounds[1], kernel_end)) + if bounds is not None + else (kernel_start, kernel_end) + ) + break + if kernel_start < end and kernel_end > start: + raise RuntimeError(f"CUPTI kernel {name} crosses a {phase} timing boundary.") + if category != "other": break - if category == "dispatch": - span = dispatch_iter_span[iter_idx] - span[0] = k_start if span[0] is None else min(span[0], k_start) - span[1] = k_end if span[1] is None else max(span[1], k_end) - dispatch_kernel_times.setdefault(demangled, []).append(device_time_us) - elif category == "combine": - span = combine_iter_span[iter_idx] - span[0] = k_start if span[0] is None else min(span[0], k_start) - span[1] = k_end if span[1] is None else max(span[1], k_end) - combine_kernel_times.setdefault(demangled, []).append(device_time_us) - else: - other_kernel_times.setdefault(demangled, []).append(device_time_us) - - def _build(ktimes: Dict[str, List[float]]) -> List[Dict[str, Any]]: - result = [{"name": n, "count": len(t), "_times": t} for n, t in ktimes.items()] - result.sort( - key=lambda x: sum(x["_times"]) / len(x["_times"]) if x["_times"] else 0, reverse=True + demangled_name = demangled_names.get(name, name) + kernel_times[category].setdefault(demangled_name, []).append( + (kernel_end - kernel_start) / 1e3 ) + + def _build(category: str) -> List[Dict[str, Any]]: + result = [ + {"name": name, "count": len(times), "_times": times} + for name, times in kernel_times[category].items() + ] + result.sort(key=lambda kernel: sum(kernel["_times"]) / len(kernel["_times"]), reverse=True) return result - dispatch_times_us = [ - (span[1] - span[0]) / 1e3 if span[0] is not None else None for span in dispatch_iter_span - ] - combine_times_us = [ - (span[1] - span[0]) / 1e3 if span[0] is not None else None for span in combine_iter_span - ] + spans = {} + for phase, bounds in phase_bounds.items(): + if any(bound is None for bound in bounds): + raise RuntimeError( + f"CUPTI captured no {phase} kernels in one or more timed iterations." + ) + # A span retains inter-kernel gaps but counts overlapping PDL kernels only once. + spans[f"{phase}_us_kernel_span"] = [ + (bound[1] - bound[0]) / 1e3 for bound in bounds if bound is not None + ] return { - "dispatch_kernels": _build(dispatch_kernel_times), - "combine_kernels": _build(combine_kernel_times), - "other_kernels": _build(other_kernel_times), - "dispatch_times_us": dispatch_times_us, - "combine_times_us": combine_times_us, + **spans, + "dispatch_kernels": _build("dispatch"), + "combine_kernels": _build("combine"), + "other_kernels": _build("other"), } -def _try_init_cupti(): - """Try to initialize CUPTI for CUDA-graph kernel breakdown. +def _init_cupti() -> tuple[Any, list[tuple[str, int, int]], list[tuple[int, int]]]: + """Enable kernel and CUDA-event tracking before CUDA context creation.""" + from cupti import cupti - MUST be called BEFORE the CUDA context is created (i.e. before any torch.cuda.* - call). CUPTI CUDA_EVENT activities are only delivered to subscribers registered - before the CUDA context is initialized; late registration silently drops them. + cupti_kernels: list[tuple[str, int, int]] = [] + cupti_events: list[tuple[int, int]] = [] - Also must be called before any NVLINK/NVLink backend creation: NVLINK_ONE_SIDED's - NVLink initialization changes CUDA profiling state in a way that prevents - CONCURRENT_KERNEL tracking if CUPTI is enabled afterwards. + def _buf_requested() -> tuple[int, int]: + return 8 * 1024 * 1024, 0 - Returns (cupti_module, kernels_list, event_timestamps_list, is_available). - """ + def _buf_completed(activities) -> None: + for activity in activities: + if activity.kind == cupti.ActivityKind.CONCURRENT_KERNEL: + cupti_kernels.append((activity.name, activity.start, activity.end)) + elif activity.kind == cupti.ActivityKind.CUDA_EVENT: + cupti_events.append((activity.event_id, activity.device_timestamp)) + + enabled = [] try: - from functools import partial as _partial + cupti.activity_register_callbacks(_buf_requested, _buf_completed) + for kind in (cupti.ActivityKind.CONCURRENT_KERNEL, cupti.ActivityKind.CUDA_EVENT): + cupti.activity_enable(kind) + enabled.append(kind) + cupti.activity_enable_cuda_event_device_timestamps(1) + except (cupti.cuptiError, AttributeError) as exc: + cleanup_errors = [] + for kind in reversed(enabled): + try: + cupti.activity_disable(kind) + except cupti.cuptiError as cleanup_exc: + cleanup_errors.append(str(cleanup_exc)) + raise RuntimeError( + f"Cannot enable CUPTI kernel/event tracing: {exc}; cleanup errors: {cleanup_errors}" + ) from exc + return cupti, cupti_kernels, cupti_events - from cupti import cupti as _cupti - _cupti_kernels: List[Tuple[str, int, int]] = [] - _cupti_events: List[int] = [] # device_timestamps of CUDA event records, in arrival order +_CUPTI_FALLBACK_WARNING_SHOWN = False - def _buf_requested(): - return 8 * 1024 * 1024, 0 - def _buf_completed(kernels, events, activities): - for act in activities: - if act.kind == _cupti.ActivityKind.CONCURRENT_KERNEL: - kernels.append((act.name, act.start, act.end)) - elif act.kind == _cupti.ActivityKind.CUDA_EVENT: - events.append(act.device_timestamp) +def _warn_kernel_span_unavailable(reason: str) -> str: + global _CUPTI_FALLBACK_WARNING_SHOWN + warning = ( + "CUPTI kernel-span timing is unavailable. " + "Falling back to CUDA-event timing, which may include non-kernel bubbles before " + "the first kernel and after the last kernel. Kernel-span timing excludes these " + "boundary bubbles and is generally more representative of communication-kernel " + f"execution time in E2E workloads. Reason: {reason}" + ) + if not _CUPTI_FALLBACK_WARNING_SHOWN: + _maybe_warn_rank0(f"[bench_moe_comm] WARNING: {warning}") + _CUPTI_FALLBACK_WARNING_SHOWN = True + return warning - _cupti.activity_enable(_cupti.ActivityKind.CONCURRENT_KERNEL) - _cupti.activity_enable(_cupti.ActivityKind.CUDA_EVENT) - _cupti.activity_enable_cuda_event_device_timestamps(1) - _cupti.activity_register_callbacks( - _buf_requested, _partial(_buf_completed, _cupti_kernels, _cupti_events) - ) - return _cupti, _cupti_kernels, _cupti_events, True - except Exception: - return None, [], [], False + +def _init_cupti_for_workers() -> tuple[Optional[Any], str | None]: + """Attempt tracing before CUDA context creation; keep MPI ranks on the same path.""" + ctx = None + error = None + try: + ctx = _init_cupti() + if ctx is None: + error = "CUPTI initialization returned no collector" + except (ImportError, OSError, RuntimeError) as exc: + error = f"{type(exc).__name__}: {exc}" + errors = mpi_allgather(error) + if any(errors): + reasons = [f"rank{rank}: {error}" for rank, error in enumerate(errors) if error] + if ctx is not None: + cupti = ctx[0] + for kind in (cupti.ActivityKind.CUDA_EVENT, cupti.ActivityKind.CONCURRENT_KERNEL): + try: + cupti.activity_disable(kind) + except cupti.cuptiError as exc: + reasons.append(f"Local CUPTI cleanup failed: {exc}") + return None, _warn_kernel_span_unavailable("; ".join(reasons)) + return ctx, None -def _time_dispatch_and_combine_cuda_graph( +def _time_dispatch_and_combine( backend: Communication, *, hidden_states: torch.Tensor, @@ -627,24 +463,25 @@ def _time_dispatch_and_combine_cuda_graph( warmup: int, iters: int, flush_l2: bool = True, + use_cuda_graph: bool = True, cupti_ctx: Optional[Any] = None, ) -> Tuple[List[float], List[float], Dict[str, Any]]: - """Time dispatch and combine using an unrolled CUDA graph + embedded CUDA events. + """Measure per-iteration dispatch/combine latency with CUDA events, in microseconds. + + After an eager shape-discovery run, execute warmup and timed iterations either + in one unrolled CUDA graph replay or eagerly. L2 flushing and simulated MoE + output initialization are outside each timed phase. Optional CUPTI activity + records are attributed using the IDs and GPU timestamps of those events. Order: - 1. One eager dispatch+combine to discover recv shape → allocate static_moe_out → sync. - 2. Capture a single big graph with `iters` iterations unrolled. - Each iteration: d_starts[i].record → dispatch → d_ends[i].record - → zero_ → c_starts[i].record → combine → c_ends[i].record - 3. Warmup: `warmup` eager iterations (no graph). - 4. Timed: one big_graph.replay() → GPU runs all iters back-to-back with zero CPU overhead. - 5. Sync, read per-iter timings from events. - 6. Profiler pass (two small graphs) for kernel breakdown. - - L2 cache is flushed before each iteration inside the graph (including warmup), - matching the eager-mode behaviour. - - Returns same types as _time_dispatch_and_combine. + 1. Discover the receive shape and allocate the static combine payload. + 2. Create timing events and obtain their CUPTI IDs when profiling is enabled. + 3. Define warmup and measured iterations; capture them together in graph mode. + 4. Clear setup profiling records, synchronize ranks, and execute the iterations. + 5. Read per-iteration dispatch/combine latency from CUDA events. + 6. Attribute CUPTI kernels to phases using the timing events' GPU timestamps. + + Returns event times and activity statistics containing per-iteration kernel spans. """ device = hidden_states.device max_tokens = max(all_rank_num_tokens) @@ -655,18 +492,10 @@ def _time_dispatch_and_combine_cuda_graph( l2_flush_size = (l2_size * 2) // 4 l2_buffer = torch.empty(l2_flush_size, dtype=torch.int32, device=device) - # ---- 0. CUPTI state ---- - # cupti_ctx is pre-initialized before backend creation (NVLINK_ONE_SIDED's NVLink - # init changes CUDA profiling state; CUPTI must be enabled before that call). if cupti_ctx is not None: - _cupti, _cupti_kernels, _cupti_events, _cupti_available = cupti_ctx - else: - _cupti_available = False - _cupti_kernels: List[Tuple[str, int, int]] = [] - _cupti_events: List[int] = [] - _cupti = None + cupti, cupti_kernels, cupti_events = cupti_ctx - # ---- 1. Shape discovery: one eager run ---- + # ---- 1. Discover receive shape and allocate the combine payload ---- backend.prepare_dispatch(token_selected_slots, all_rank_num_tokens) recv_hidden_states, _, _, _ = backend.dispatch( hidden_states, @@ -683,16 +512,24 @@ def _time_dispatch_and_combine_cuda_graph( backend.combine(static_moe_out, all_rank_max_num_tokens=max_tokens) torch.cuda.synchronize() - # ---- 2. Capture big graph (iters iterations unrolled) ---- + # ---- 2. Prepare timing events and CUPTI event IDs ---- # cudaEventRecordExternal (0x1, CUDA 11.2+) makes events recorded inside a # CUDA graph queryable via elapsed_time() after replay. Without this flag, # graph-internal events raise cudaErrorInvalidValue on elapsed_time(). - _cudart = ctypes.CDLL("libcudart.so") - _cudart.cudaEventRecordWithFlags.restype = ctypes.c_int - _cudart.cudaEventRecordWithFlags.argtypes = [ctypes.c_void_p, ctypes.c_void_p, ctypes.c_uint] - _CUDA_EVENT_RECORD_EXTERNAL = 0x1 + if use_cuda_graph: + _cudart = ctypes.CDLL("libcudart.so") + _cudart.cudaEventRecordWithFlags.restype = ctypes.c_int + _cudart.cudaEventRecordWithFlags.argtypes = [ + ctypes.c_void_p, + ctypes.c_void_p, + ctypes.c_uint, + ] + _CUDA_EVENT_RECORD_EXTERNAL = 0x1 - def _record_external(event: torch.cuda.Event) -> None: + def _record_event(event: torch.cuda.Event) -> None: + if not use_cuda_graph: + event.record() + return stream = torch.cuda.current_stream() ret = _cudart.cudaEventRecordWithFlags( event.cuda_event, stream.cuda_stream, _CUDA_EVENT_RECORD_EXTERNAL @@ -710,11 +547,20 @@ def _record_external(event: torch.cuda.Event) -> None: evt.record() torch.cuda.synchronize() - # Graph contains warmup + timed iters. Warmup iters have no events (unmeasured). - # Timed iters have 4 external events each. One replay() runs everything back-to-back, - # eliminating rank desync between warmup and timed sections. - big_graph = torch.cuda.CUDAGraph() - with torch.cuda.graph(big_graph): + phase_event_ids = [] + if cupti_ctx is not None: + for i in range(iters): + phase_event_ids.append( + ( + cupti.get_cuda_event_id(d_starts[i].cuda_event), + cupti.get_cuda_event_id(d_ends[i].cuda_event), + cupti.get_cuda_event_id(c_starts[i].cuda_event), + cupti.get_cuda_event_id(c_ends[i].cuda_event), + ) + ) + + # ---- 3. Define iterations and optionally capture the CUDA graph ---- + def _run_iterations() -> None: for _ in range(warmup): if l2_buffer is not None: l2_buffer.zero_() @@ -733,7 +579,7 @@ def _record_external(event: torch.cuda.Event) -> None: for i in range(iters): if l2_buffer is not None: l2_buffer.zero_() - _record_external(d_starts[i]) + _record_event(d_starts[i]) backend.prepare_dispatch( token_selected_slots, all_rank_num_tokens ) # For most ranks this is no-op except for NVLINK_TWO_SIDED @@ -744,56 +590,44 @@ def _record_external(event: torch.cuda.Event) -> None: token_final_scales, all_rank_num_tokens, ) - _record_external(d_ends[i]) + _record_event(d_ends[i]) static_moe_out.zero_() - _record_external(c_starts[i]) + _record_event(c_starts[i]) backend.combine(static_moe_out, all_rank_max_num_tokens=max_tokens) - _record_external(c_ends[i]) + _record_event(c_ends[i]) - # ---- 3. Timed replay + kernel breakdown via CUPTI ---- - if _cupti_available: - # Flush any activities captured before the replay (shape discovery, graph capture - # dry-run, etc.) and clear lists so only replay activities remain. - _cupti.activity_flush_all(0) - _cupti_kernels.clear() - _cupti_events.clear() + if use_cuda_graph: + # Keep warmup and timed iterations in one replay to avoid a host-side gap. + big_graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(big_graph): + _run_iterations() + # ---- 4. Clear setup records, synchronize, and execute ---- + if cupti_ctx is not None: + cupti.activity_flush_all(0) + cupti_kernels.clear() + cupti_events.clear() _sync() - big_graph.replay() - + if use_cuda_graph: + big_graph.replay() + else: + _run_iterations() _sync() + if cupti_ctx is not None: + cupti.activity_flush_all(0) - if _cupti_available: - # Flush AFTER _sync() (torch.cuda.synchronize + mpi_barrier) to ensure CUPTI - # delivers all pending graph-replay activities. flush_all(0) is non-blocking; - # the preceding synchronize gives CUPTI time to process the replay's records. - _cupti.activity_flush_all(0) - + # ---- 5. Read per-iteration CUDA-event timings ---- dispatch_times_us = [d_starts[i].elapsed_time(d_ends[i]) * 1e3 for i in range(iters)] combine_times_us = [c_starts[i].elapsed_time(c_ends[i]) * 1e3 for i in range(iters)] - if _cupti_available: - _cupti_kernels.sort(key=lambda k: k[1]) - _cupti_events.sort() # sort by device_timestamp; CUPTI may deliver out of order - - detailed_stats = _build_cuda_graph_kernel_stats_cupti(_cupti_kernels, _cupti_events, iters) - if detailed_stats is not None: - # Replace event-based times with tighter kernel-span times. - # Fall back per-iter to event timing if no kernels were attributed. - cupti_dispatch = detailed_stats.pop("dispatch_times_us") - cupti_combine = detailed_stats.pop("combine_times_us") - dispatch_times_us = [ - ct if ct is not None else et - for ct, et in zip(cupti_dispatch, dispatch_times_us, strict=True) - ] - combine_times_us = [ - ct if ct is not None else et - for ct, et in zip(cupti_combine, combine_times_us, strict=True) - ] - else: - detailed_stats = {"dispatch_kernels": [], "combine_kernels": [], "other_kernels": []} - else: - detailed_stats = {"dispatch_kernels": [], "combine_kernels": [], "other_kernels": []} + # ---- 6. Attribute CUPTI kernels to dispatch/combine phases ---- + detailed_stats = {"dispatch_kernels": [], "combine_kernels": [], "other_kernels": []} + if cupti_ctx is not None: + try: + detailed_stats = _build_kernel_stats_cupti(cupti_kernels, cupti_events, phase_event_ids) + except RuntimeError as exc: + # Let every MPI rank reach the reporting collectives even if one trace is incomplete. + detailed_stats["cupti_error"] = str(exc) return dispatch_times_us, combine_times_us, detailed_stats @@ -815,109 +649,19 @@ def _compute_stats(values: List[float]) -> Dict[str, float]: } -def _gather_per_rank(times_us: List[float], iter_stats: bool = False) -> Dict[str, Any]: +def _gather_per_rank(times_us: Optional[List[float]], iter_stats: bool = False) -> Dict[str, Any]: """Allgather per-iteration times from each rank, return per-rank results. If iter_stats=True, return full stats (mean/median/stdev/min/max). If iter_stats=False, return just the mean. """ all_times = mpi_allgather(times_us) - if iter_stats: - return {f"rank{i}": _compute_stats(t) for i, t in enumerate(all_times)} - return {f"rank{i}": (sum(t) / len(t) if t else 0.0) for i, t in enumerate(all_times)} - - -def _min_local_tokens_for_receiver_coverage(ep_size: int, top_k: int) -> int: - if top_k <= 0: - raise ValueError(f"top_k must be > 0, got {top_k}") - return (ep_size + top_k - 1) // top_k - - -def _scale_local_batch_sizes_for_receiver_coverage( - local_batch_sizes: List[int], ep_size: int, top_k: int -) -> List[int]: - min_tokens = _min_local_tokens_for_receiver_coverage(ep_size, top_k) - scaled: List[int] = [] - for local_num_tokens in local_batch_sizes: - value = max(int(local_num_tokens), min_tokens) - if not scaled or scaled[-1] != value: - scaled.append(value) - return scaled - - -def _verify_dispatch_sentinel( - backend: Communication, - *, - hidden_size: int, - top_k: int, - experts_per_rank: int, - ep_size: int, - act_dtype: torch.dtype, - device: torch.device, - local_num_tokens: Optional[int] = None, -) -> Dict[str, Any]: - """One dispatch+combine with sender-rank-tagged hidden_states. - - Each rank fills its hidden_states with the scalar ``rank + 1``. After - dispatch, each received row should be that integer cast to ``act_dtype``; - rows reading as 0 are either padding or a silently-broken peer read - (e.g. cross-rack MNNVL mapping that succeeded at construction but doesn't - actually back the peer's memory). Returns the per-rank decoded-sender - histogram for the caller to allgather and inspect. - """ - rank = mpi_rank() - min_tokens = _min_local_tokens_for_receiver_coverage(ep_size, top_k) - local_num_tokens = min_tokens if local_num_tokens is None else max(local_num_tokens, min_tokens) - all_rank_num_tokens = mpi_allgather(int(local_num_tokens)) - if not backend.is_workload_feasible(all_rank_num_tokens, num_chunks=1): - return {"rank": rank, "skipped": True} - - sentinel = float(rank + 1) - hidden_states = torch.full( - (local_num_tokens, hidden_size), - sentinel, - dtype=act_dtype, - device=device, - ) - flat_slots = torch.arange(local_num_tokens * top_k, device=device, dtype=torch.int64) - schedule = flat_slots + rank - target_rank = schedule % ep_size - local_expert = (schedule // ep_size) % experts_per_rank - token_selected_slots = ( - (target_rank * experts_per_rank + local_expert) - .view(local_num_tokens, top_k) - .to(torch.int32) - ) - token_final_scales = torch.ones( - local_num_tokens, - top_k, - dtype=torch.float32, - device=device, - ) - - backend.prepare_dispatch(token_selected_slots, all_rank_num_tokens) - recv_hs, _, _, _ = backend.dispatch( - hidden_states, - None, - token_selected_slots, - token_final_scales, - all_rank_num_tokens, - ) - # Pair dispatch with a combine so backend state mirrors the bench's - # warmup->timing call pattern (NCCL_EP especially relies on this). - shape = list(recv_hs.shape) - shape[-1] = hidden_size - moe_out = torch.zeros(tuple(shape), dtype=torch.bfloat16, device=recv_hs.device) - backend.combine(moe_out, all_rank_max_num_tokens=max(all_rank_num_tokens)) - torch.cuda.synchronize() - - first_col = recv_hs[:, 0].to(torch.float32) - decoded = first_col.round().to(torch.int64) - unique, counts = decoded.unique(return_counts=True) - histogram: Dict[int, int] = { - int(u) - 1: int(c) for u, c in zip(unique.tolist(), counts.tolist(), strict=True) + return { + f"rank{i}": None + if t is None + else (_compute_stats(t) if iter_stats else (sum(t) / len(t) if t else 0.0)) + for i, t in enumerate(all_times) } - return {"rank": rank, "histogram": histogram} def parse_args() -> argparse.Namespace: @@ -1012,7 +756,7 @@ def parse_args() -> argparse.Namespace: parser.add_argument( "--kernel_breakdown", action="store_true", - help="Show per-kernel timing breakdown.", + help="Also output per-kernel details. CUPTI kernel-span timing is attempted regardless of this flag.", ) parser.add_argument( "--iter_stats", @@ -1037,7 +781,7 @@ def parse_args() -> argparse.Namespace: help="Use deterministic balanced router assignments to avoid communication load imbalance.", ) parser.add_argument( - "--use_low_precision_moe_combine", + "--use_low_precision_combine", action="store_true", default=False, help="Enable low-precision (FP8) MoE combine path.", @@ -1045,17 +789,7 @@ def parse_args() -> argparse.Namespace: parser.add_argument( "--no_cuda_graph", action="store_true", - help="Disable CUDA graph mode. By default, dispatch and combine are captured into CUDA graphs for lower CPU overhead and more accurate timing.", - ) - parser.add_argument( - "--verify", - action="store_true", - help=( - "Run a single sentinel dispatch per backend before timing and print a " - "receiver/sender contribution matrix. Detects silent cross-rack " - "correctness failures where dispatch appears to succeed but produces " - "zeros or local-only data." - ), + help="Use eager execution with CUDA event timing instead of CUDA graph replay. Kernel breakdown still uses CUPTI.", ) parser.add_argument( "--pdl", @@ -1131,16 +865,8 @@ def _resolve_profile_args(args: argparse.Namespace) -> Tuple[int, int, int, Quan def _run_benchmark_worker_under_current_mpi( args: argparse.Namespace, launcher: str = "spawn" ) -> None: - # CUPTI MUST be initialized before the CUDA context is created. - # CUDA_EVENT activities are only delivered to CUPTI subscribers that were registered - # before the CUDA context was initialized; late registration captures CONCURRENT_KERNEL - # but silently drops CUDA_EVENT records. _set_device_from_local_rank() (below) is - # the first call that creates the CUDA context, so we init CUPTI here. - _early_cupti_ctx: Optional[Any] = None - if not args.no_cuda_graph: - _cupti_module, _cupti_kernels_list, _cupti_events_list, _cupti_ok = _try_init_cupti() - if _cupti_ok: - _early_cupti_ctx = (_cupti_module, _cupti_kernels_list, _cupti_events_list, True) + # Late CUPTI initialization captures kernels but misses CUDA_EVENT records. + cupti_ctx, cupti_warning = _init_cupti_for_workers() # Keep benchmark output clean. tllm.logger.set_level("error") @@ -1156,10 +882,6 @@ def _run_benchmark_worker_under_current_mpi( hidden_size, top_k, num_experts_total, quant_algo = _resolve_profile_args(args) local_batch_sizes = _iter_local_batch_sizes(args) - if args.verify: - local_batch_sizes = _scale_local_batch_sizes_for_receiver_coverage( - local_batch_sizes, ep_size, top_k - ) act_dtype = torch.bfloat16 quant_config = ( QuantConfig(quant_algo=None) @@ -1203,6 +925,8 @@ def _run_benchmark_worker_under_current_mpi( "device_count": torch.cuda.device_count(), "cuda_graph": not args.no_cuda_graph, "pdl": bool(args.pdl), + "cupti_enabled": cupti_ctx is not None, + "warning": cupti_warning, } if rank == 0: print(json.dumps(benchmark_metadata, indent=2), flush=True) @@ -1222,14 +946,6 @@ def _run_benchmark_worker_under_current_mpi( all_results: List[Dict[str, Any]] = [] - # CUPTI was initialized before the CUDA context at the top of this function. - # Reuse that early context; do not re-initialize here (too late for CUDA_EVENT delivery). - _cupti_ctx: Optional[Any] = _early_cupti_ctx - if not args.no_cuda_graph and _cupti_ctx is None: - _maybe_warn_rank0( - "[bench] CUPTI unavailable; dispatch_us/combine_us will use CUDA event elapsed_time." - ) - for backend_name in backends: try: model_config = _create_model_config( @@ -1238,7 +954,7 @@ def _run_benchmark_worker_under_current_mpi( act_dtype=act_dtype, max_num_tokens_per_rank=max_num_tokens_per_rank, quant_config=quant_config, - use_low_precision_moe_combine=args.use_low_precision_moe_combine, + use_low_precision_combine=args.use_low_precision_combine, ) backend = CommunicationFactory._create_forced_method( # pylint: disable=protected-access @@ -1280,58 +996,6 @@ def _run_benchmark_worker_under_current_mpi( # Ensure quantization params (e.g., NVFP4 global scale) live on CUDA. moe = moe.to(device) - if args.verify: - verify_local = _verify_dispatch_sentinel( - backend, - hidden_size=hidden_size, - top_k=top_k, - experts_per_rank=experts_per_rank, - ep_size=ep_size, - act_dtype=act_dtype, - device=device, - local_num_tokens=local_batch_sizes[0], - ) - all_verify = mpi_allgather(verify_local) - # Pass criterion: every receiver must have at least one token from - # every sender [0, ep_size). The verify local_num_tokens is scaled - # so local_num_tokens * top_k covers every receiver; - # any zero-column means the recv buffer was silently dropped from - # that sender. - verify_failed = False - for entry in all_verify: - if entry.get("skipped"): - verify_failed = True - break - hist = entry.get("histogram", {}) - if any(hist.get(s, 0) == 0 for s in range(ep_size)): - verify_failed = True - break - if rank == 0: - status = "FAIL" if verify_failed else "PASS" - print( - f"=== [verify] {backend_name} {status} -- sender->receiver " - f"contribution (rows=receiver, cols=sender; -1 col = " - f"padding/unmapped) ===", - flush=True, - ) - cols = [-1, *range(ep_size)] - header = "R\\S | " + " ".join(f"{c:>5}" for c in cols) + " | total" - print(header) - for entry in sorted(all_verify, key=lambda e: e.get("rank", -1)): - r = entry.get("rank") - if entry.get("skipped"): - print(f"{r:>3} | skipped (workload not feasible at verify size)") - continue - hist = entry.get("histogram", {}) - cells = " ".join(f"{hist.get(c, 0):>5}" for c in cols) - print(f"{r:>3} | {cells} | {sum(hist.values()):>5}") - sys.stdout.flush() - if verify_failed: - _maybe_warn_rank0( - f"[bench_moe_comm] Skipping timing for {backend_name}: verify FAILED." - ) - continue - for local_num_tokens in local_batch_sizes: all_rank_num_tokens = mpi_allgather(int(local_num_tokens)) if not backend.is_workload_feasible(all_rank_num_tokens, num_chunks=1): @@ -1357,12 +1021,8 @@ def _run_benchmark_worker_under_current_mpi( ) # Time dispatch and combine - _time_fn = ( - _time_dispatch_and_combine_cuda_graph - if not args.no_cuda_graph - else _time_dispatch_and_combine - ) - time_fn_kwargs: Dict[str, Any] = dict( + dispatch_times_us, combine_times_us, detailed_stats = _time_dispatch_and_combine( + backend, hidden_states=hidden_states, hidden_states_sf=hidden_states_sf, token_selected_slots=token_selected_slots, @@ -1372,23 +1032,34 @@ def _run_benchmark_worker_under_current_mpi( warmup=int(args.warmup), iters=int(args.iters), flush_l2=True, - ) - if not args.no_cuda_graph: - time_fn_kwargs["cupti_ctx"] = _cupti_ctx - dispatch_times_us, combine_times_us, detailed_stats = _time_fn( - backend, **time_fn_kwargs + use_cuda_graph=not args.no_cuda_graph, + cupti_ctx=cupti_ctx, ) iter_stats = bool(args.iter_stats) - dispatch_stats = _gather_per_rank(dispatch_times_us, iter_stats=iter_stats) - combine_stats = _gather_per_rank(combine_times_us, iter_stats=iter_stats) + cupti_errors = mpi_allgather(detailed_stats.get("cupti_error")) + if any(cupti_errors): + warning = _warn_kernel_span_unavailable( + f"{backend_name} @ local_batch_size={local_num_tokens}: " + + "; ".join( + f"rank{rank}: {error}" for rank, error in enumerate(cupti_errors) if error + ) + ) + previous_warning = benchmark_metadata["warning"] + benchmark_metadata["warning"] = ( + f"{previous_warning}\n{warning}" if previous_warning else warning + ) + elif cupti_ctx is not None: + # Use the same timing source on every rank for this measurement. + dispatch_times_us = detailed_stats["dispatch_us_kernel_span"] + combine_times_us = detailed_stats["combine_us_kernel_span"] # Prepare output output = { "backend": backend_name, "local_batch_size": int(local_num_tokens), - "dispatch_us": dispatch_stats, - "combine_us": combine_stats, + "dispatch_us": _gather_per_rank(dispatch_times_us, iter_stats=iter_stats), + "combine_us": _gather_per_rank(combine_times_us, iter_stats=iter_stats), } # Add kernel breakdown if requested and available diff --git a/tests/microbenchmarks/compare_moe_comm.py b/tests/microbenchmarks/compare_moe_comm.py index f12657c4dc1b..7f9ad32fa30b 100644 --- a/tests/microbenchmarks/compare_moe_comm.py +++ b/tests/microbenchmarks/compare_moe_comm.py @@ -142,6 +142,8 @@ def main(): ep = meta.get("ep_size", "?") backend = meta.get("backend", "?") print(f"[{tag}] {lbl} (ep={ep}, backend={backend})") + if meta.get("warning"): + print(f"[{tag}] WARNING: {meta['warning']}") print(f"Stat: {args.stat}, Rank: {args.rank}") print() @@ -158,8 +160,11 @@ def main(): sub_parts.append(f"{'B (us)':>{kernel_col_width}}") sub_parts.append(f"{'(A/B)':>{kernel_col_width}}") - # Also show total dispatch and total combine - for total_name in ["total_dispatch", "total_combine"]: + phase_metrics = [ + ("dispatch_us", "dispatch"), + ("combine_us", "combine"), + ] + for _, total_name in phase_metrics: header_parts.append(f"{total_name:>{kernel_col_width}}") header_parts.append(f"{total_name:>{kernel_col_width}}") header_parts.append(f"{'speedup':>{kernel_col_width}}") @@ -197,8 +202,7 @@ def main(): row.append(f"{'N/A':>{kernel_col_width}}") row.append(f"{'N/A':>{kernel_col_width}}") - # Total dispatch and total combine - for key in ["dispatch_us", "combine_us"]: + for key, _ in phase_metrics: ta = ra.get(key, {}).get(args.rank) tb = rb.get(key, {}).get(args.rank) if ta and tb: diff --git a/tests/unittest/_torch/distributed/test_mnnvl_memory_comm.py b/tests/unittest/_torch/distributed/test_mnnvl_memory_comm.py index e92c42879a86..d9c68865c914 100644 --- a/tests/unittest/_torch/distributed/test_mnnvl_memory_comm.py +++ b/tests/unittest/_torch/distributed/test_mnnvl_memory_comm.py @@ -32,8 +32,12 @@ import pytest import torch -from tensorrt_llm import _mnnvl_utils -from tensorrt_llm._mnnvl_utils import HelixCpMnnvlMemory, MnnvlMemory, ProcessGroupComm +from tensorrt_llm._torch.distributed import mnnvl_memory as _mnnvl_utils +from tensorrt_llm._torch.distributed.mnnvl_memory import ( + HelixCpMnnvlMemory, + MnnvlMemory, + ProcessGroupComm, +) from tensorrt_llm._torch.models.modeling_utils import MetaInitException, MetaInitMode diff --git a/tests/unittest/_torch/distributed/test_mnnvl_workspace_comm.py b/tests/unittest/_torch/distributed/test_mnnvl_workspace_comm.py index d64d69df7813..b4314d1657e6 100644 --- a/tests/unittest/_torch/distributed/test_mnnvl_workspace_comm.py +++ b/tests/unittest/_torch/distributed/test_mnnvl_workspace_comm.py @@ -241,7 +241,7 @@ def test_device_index_uses_local_rank_under_mpi(mpi_mode): @pytest.fixture def mnnvl_capable_hardware(monkeypatch): """Make every hardware-level precondition of is_mnnvl() pass.""" - import tensorrt_llm._mnnvl_utils as mnnvl_utils + import tensorrt_llm._torch.distributed.mnnvl_memory as mnnvl_utils monkeypatch.setattr(ops.platform, "machine", lambda: "aarch64") monkeypatch.setattr(mnnvl_utils.MnnvlMemory, "supports_mnnvl", staticmethod(lambda: True)) diff --git a/tests/unittest/_torch/misc/test_moe_a2a_warmup_timeout.py b/tests/unittest/_torch/misc/test_moe_a2a_warmup_timeout.py deleted file mode 100644 index 44817aa76173..000000000000 --- a/tests/unittest/_torch/misc/test_moe_a2a_warmup_timeout.py +++ /dev/null @@ -1,82 +0,0 @@ -import unittest -from unittest import mock - -from tensorrt_llm._torch.pyexecutor import model_engine -from tensorrt_llm._torch.pyexecutor.model_engine import PyTorchModelEngine - - -class _WarmupFlagStub: - """Minimal object that reuses the engine's is_warmup property. - - Building a real PyTorchModelEngine needs a model and a device; the property - itself only touches _is_warmup, the MoE all-to-all budget selector, and - moe_load_balancer_iter_info (a no-op when moe_load_balancer is None), so a - stub exercises the real code path without either. - - The stub borrows the property objects without inheriting, so it has to - declare moe_load_balancer itself. - """ - - is_warmup = PyTorchModelEngine.is_warmup - moe_load_balancer_iter_info = PyTorchModelEngine.moe_load_balancer_iter_info - moe_load_balancer = None - - -class TestMoeA2AWarmupBudget(unittest.TestCase): - """The MoE all-to-all completion-flag budget must track the warmup phase. - - The kernel-side deadline is only safe if it is raised for warmup *and* - lowered again afterwards; a budget that latches on would leave the hang - watchdog permanently relaxed in steady state. See nvbugs/6482566. - """ - - def test_set_warmup_forwards_value_to_op(self): - with mock.patch.object( - model_engine.torch.ops.trtllm, "moe_a2a_set_warmup", create=True - ) as op: - model_engine._set_moe_a2a_warmup(True) - model_engine._set_moe_a2a_warmup(False) - self.assertEqual([c.args[0] for c in op.call_args_list], [True, False]) - - def test_missing_op_is_tolerated(self): - """An older C++ build without the op must not break startup.""" - with mock.patch.object( - model_engine.torch.ops.trtllm, - "moe_a2a_set_warmup", - create=True, - side_effect=AttributeError("no such op"), - ): - model_engine._set_moe_a2a_warmup(True) # must not raise - - def test_capture_context_selects_steady_state_then_restores(self): - """CUDA graphs bake the budget in at capture time. - - Capture runs inside the warmup window, so the context manager must hand - the kernel the steady-state budget and restore warmup afterwards. - """ - seen = [] - with mock.patch.object(model_engine, "_set_moe_a2a_warmup", side_effect=seen.append): - with model_engine._moe_a2a_steady_state_budget_for_capture(): - self.assertEqual(seen, [False]) - self.assertEqual(seen, [False, True]) - - def test_is_warmup_setter_switches_budget_both_ways(self): - """Regression: the budget must not latch on after warmup. - - PyExecutor sets is_warmup=True before calling warmup() and False after, - both through this setter. Selecting the budget anywhere else (e.g. only - in set_warmup_flag) leaves the relaxed warmup budget in force for the - whole serving lifetime. - """ - stub = _WarmupFlagStub() - seen = [] - with mock.patch.object(model_engine, "_set_moe_a2a_warmup", side_effect=seen.append): - stub.is_warmup = True - stub.is_warmup = False - - self.assertEqual(seen, [True, False]) - self.assertFalse(stub.is_warmup) - - -if __name__ == "__main__": - unittest.main() diff --git a/tests/unittest/_torch/moe/multi_gpu/test_moe_a2a.py b/tests/unittest/_torch/moe/multi_gpu/test_moe_a2a.py deleted file mode 100644 index a433dc494898..000000000000 --- a/tests/unittest/_torch/moe/multi_gpu/test_moe_a2a.py +++ /dev/null @@ -1,923 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import pickle -import sys -import traceback - -import cloudpickle -import pytest -import torch -from mpi4py import MPI - -import tensorrt_llm as tllm -from tensorrt_llm._mnnvl_utils import MnnvlMemory -from tensorrt_llm._torch.moe.fused_moe.communication.moe_alltoall import \ - MoeAlltoAll -from tensorrt_llm.mapping import Mapping - -cloudpickle.register_pickle_by_value(sys.modules[__name__]) -MPI.pickle.__init__( - cloudpickle.dumps, - cloudpickle.loads, - pickle.HIGHEST_PROTOCOL, -) - - -@pytest.fixture(autouse=True) -def setup_test(): - torch.manual_seed(0x1234) - tllm.logger.set_level('error') - - -def compute_target_rank_id(expert_id, num_experts_per_rank): - """Compute the rank that owns a given expert using contiguous partitioning. - Experts are divided evenly across ranks: - - Rank 0: experts [0, num_experts_per_rank) - - Rank 1: experts [num_experts_per_rank, 2 * num_experts_per_rank) - - ... - For example, with 32 experts and 4 ranks (8 experts per rank): - - Rank 0: experts 0-7 - - Rank 1: experts 8-15 - - Rank 2: experts 16-23 - - Rank 3: experts 24-31 - """ - return expert_id // num_experts_per_rank - - -def generate_token_selected_experts(local_num_tokens: int, num_experts: int, - top_k: int) -> torch.Tensor: - """Generate global expert IDs tensor, aligned with single-GPU test semantics.""" - return torch.randint( - 0, - num_experts, - (local_num_tokens, top_k), - dtype=torch.int32, - device='cuda', - ) - - -def create_experts_per_rank(num_experts_per_rank, - hidden_size, - ep_rank, - device, - dtype=torch.bfloat16): - """ - Create a 3D tensor of expert weights for a given rank. - - Args: - num_experts_per_rank: Number of experts on this rank - hidden_size: Hidden dimension size - ep_rank: EP rank ID - device: Device to create experts on - - Returns: - experts: Tensor of shape [num_experts_per_rank, hidden_size, hidden_size] - """ - # For reproducibility, set the seed based on rank - experts = torch.empty((num_experts_per_rank, hidden_size, hidden_size), - dtype=dtype, - device=device) - for i in range(num_experts_per_rank): - torch.manual_seed(ep_rank * 1000 + i) - # Xavier uniform initialization for each expert - torch.nn.init.xavier_uniform_(experts[i]) - return experts - - -def fake_moe(hidden_states, - token_selected_experts, - token_final_scales, - experts, - is_ep=False, - ep_rank=None, - num_experts_per_rank=None): - """ - Emulate MoE computation by scaling tokens based on which experts belong to this rank. - - Args: - hidden_states: [num_tokens, hidden_size] - input hidden states - token_selected_experts: [num_tokens, top_k] - selected expert indices - token_final_scales: [num_tokens, top_k] - scaling factors for each expert - experts: [num_experts_per_rank, hidden_size, hidden_size] if is_ep, otherwise [num_experts, hidden_size, hidden_size] - expert weights - is_ep: If true, emulate MoE on a EP rank; otherwise, emulate MoE with all experts - ep_rank: EP rank ID - num_experts_per_rank: Number of experts per rank - - Returns: - processed_states: [num_tokens, hidden_size] - processed hidden states - """ - num_tokens, _ = hidden_states.shape - _, top_k = token_selected_experts.shape - - if is_ep: - assert ep_rank is not None and num_experts_per_rank is not None - - # Initialize output - processed_states = torch.zeros_like(hidden_states) - - # Process each token - for token_idx in range(num_tokens): - # For each expert selected for this token/ - for k in range(top_k): - expert_id = token_selected_experts[token_idx, k].item() - if is_ep: - if not (expert_id >= ep_rank * num_experts_per_rank - and expert_id < (ep_rank + 1) * num_experts_per_rank): - continue - # Convert global expert ID to local expert ID for this rank - local_expert_id = expert_id - ep_rank * num_experts_per_rank - expert = experts[local_expert_id] - else: - expert = experts[expert_id] - - scale = token_final_scales[token_idx, k] - processed_states[ - token_idx] += hidden_states[token_idx] @ expert * scale - - return processed_states - - -def make_nvfp4_payloads( - local_num_tokens: int, hidden_size: int, top_k: int, rank: int, - token_selected_experts: torch.Tensor) -> tuple[list, int]: - """Create the four NV FP4 payloads exactly as in single-GPU test.""" - payloads = [] - # Payload 0: Packed FP4 tokens (uint8) - packed_hidden_size = hidden_size // 2 - packed_hidden_states = torch.randint(0, - 256, - (local_num_tokens, packed_hidden_size), - dtype=torch.uint8, - device='cuda') - payloads.append(packed_hidden_states) - - # Payload 1: Scaling factors (fp8) - num_elts_per_sf = 16 - num_scaling_factors = hidden_size // num_elts_per_sf - scaling_factors = torch.randn( - local_num_tokens, - num_scaling_factors, - dtype=torch.float32, - device='cuda') # .to(torch.float8_e4m3fn) TODO: Test failed. - scaling_factors += rank - payloads.append(scaling_factors) - - # Payload 2: token_selected_experts - payloads.append(token_selected_experts) - - # Payload 3: token_final_scales (bfloat16) - token_final_scales = torch.rand(local_num_tokens, - top_k, - dtype=torch.bfloat16, - device='cuda') - - # Construct the data to contain info about send rank and local_token_idx, which is used for debugging - # token_final_scales[:, 0] = rank - # token_final_scales[:, 1] = torch.linspace(0, local_num_tokens - 1, local_num_tokens, dtype=torch.bfloat16, device='cuda') - - payloads.append(token_final_scales) - return payloads, 2 - - -def make_bfloat16_payloads( - local_num_tokens: int, hidden_size: int, top_k: int, rank: int, - token_selected_experts: torch.Tensor) -> tuple[list, int]: - """Create bfloat16 test payloads matching nvfp4 structure but without scaling factors.""" - payloads = [] - - # Payload 0: Hidden states (bfloat16) - hidden_states = torch.randn(local_num_tokens, - hidden_size, - dtype=torch.bfloat16, - device='cuda') - # Add rank-specific pattern for verification - hidden_states += rank - payloads.append(hidden_states) - - # Payload 1: token_selected_experts - payloads.append(token_selected_experts) - - # Payload 2: token_final_scales (bfloat16) - similar to nvfp4's payload 4 - token_final_scales = torch.rand(local_num_tokens, - top_k, - dtype=torch.bfloat16, - device='cuda') - - # Optional: Construct the data that is easier to debug - # token_final_scales[:, 0] = rank - # token_final_scales[:, 1] = torch.linspace(0, local_num_tokens - 1, local_num_tokens, dtype=torch.bfloat16, device='cuda') - - payloads.append(token_final_scales) - - return payloads, 1 - - -_CFT_COUNTER_STRIDE_BYTES = 256 -_CFT_COUNTER_STRIDE_U64 = _CFT_COUNTER_STRIDE_BYTES // 8 - - -def read_cft_dispatch_counters(moe_a2a, rank: int, - ep_size: int) -> torch.Tensor: - counter_offset = moe_a2a.metainfo[MoeAlltoAll._METAINFO_INDEX[ - "DISPATCH_COUNTED_WRITE_COUNTERS_OFFSET_INDEX"]].item() - counters = moe_a2a.workspace[rank, counter_offset:counter_offset + - ep_size * _CFT_COUNTER_STRIDE_BYTES].view( - torch.int64) - return counters[::_CFT_COUNTER_STRIDE_U64][:ep_size].cpu() - - -def run_moe_a2a_dispatch_single_rank(ep_size, all_num_tokens, top_k, - workspace_size_per_rank, num_experts, - hidden_size, invalid_token_expert_id, - enable_eplb): - """Worker function for MPIPoolExecutor.""" - rank = tllm.mpi_rank() - torch.cuda.set_device(rank) - - try: - mapping = Mapping( - rank=rank, - tp_size=ep_size, - moe_ep_size=ep_size, - world_size=ep_size, - ) - - # Create MoeAlltoAll manager - max_num_tokens = max(all_num_tokens) - - eplb_stats_num_experts = ( - num_experts // 2 if enable_eplb else None - ) # Use half of the experts for testing EPLB stats - moe_a2a = MoeAlltoAll( - mapping=mapping, - max_num_tokens=max_num_tokens, - top_k=top_k, - num_slots=num_experts, - workspace_size_per_rank=workspace_size_per_rank, - num_experts=eplb_stats_num_experts, - ) - - # Get the number of tokens for this specific rank (same as single-GPU) - rank_local_tokens = all_num_tokens[rank] - - # Generate data using helper functions - token_selected_experts = generate_token_selected_experts( - rank_local_tokens, num_experts, top_k) - payloads, expert_id_payload_index = make_nvfp4_payloads( - rank_local_tokens, hidden_size, top_k, rank, token_selected_experts) - - eplb_local_stats = None - if enable_eplb: - eplb_local_stats = (torch.arange( - eplb_stats_num_experts, dtype=torch.int32, device="cuda") + - rank * 1000) - - payload_bytes_per_token = [ - payload.shape[1] * payload.element_size() for payload in payloads - ] - actual_cft_dispatch = ( - moe_a2a.can_use_cft_counted_writes - and moe_a2a.use_cft_for_dispatch(max_num_tokens) - and all(bytes_per_token % 16 == 0 - for bytes_per_token in payload_bytes_per_token)) - - cft_counters_before = (read_cft_dispatch_counters( - moe_a2a, rank, ep_size) if actual_cft_dispatch else None) - if actual_cft_dispatch: - tllm.mpi_barrier() - - recv_tensors = moe_a2a.dispatch( - token_selected_experts, - payloads, - max_num_tokens, - invalid_token_expert_id=invalid_token_expert_id, - expert_id_payload_index=expert_id_payload_index, - eplb_local_stats=eplb_local_stats) - - if not actual_cft_dispatch: - completion_flags_offset = moe_a2a.metainfo[ - MoeAlltoAll._METAINFO_INDEX[ - "DISPATCH_COMPLETION_FLAGS_OFFSET_INDEX"]].item() - completion_flags = moe_a2a.workspace[ - rank, completion_flags_offset:completion_flags_offset + - ep_size * 4].view(torch.int32).cpu() - flag_val_offset = moe_a2a.metainfo[ - MoeAlltoAll._METAINFO_INDEX["FLAG_VAL_OFFSET_INDEX"]].item() - expected_flag_val = moe_a2a.workspace[ - rank, - flag_val_offset:flag_val_offset + 4].view(torch.int32).cpu() - assert torch.all(completion_flags == expected_flag_val), ( - f"Rank {rank} completion flags: {completion_flags}, expected flag val: {expected_flag_val}" - ) - - # Read counters and compact routing tensors from workspace - send_counters_offset = moe_a2a.metainfo[ - MoeAlltoAll._METAINFO_INDEX["SEND_COUNTERS_OFFSET_INDEX"]].item() - recv_counters_offset = moe_a2a.metainfo[ - MoeAlltoAll._METAINFO_INDEX["RECV_COUNTERS_OFFSET_INDEX"]].item() - topk_target_ranks_offset = moe_a2a.metainfo[MoeAlltoAll._METAINFO_INDEX[ - "TOPK_TARGET_RANKS_OFFSET_INDEX"]].item() - topk_send_indices_offset = moe_a2a.metainfo[MoeAlltoAll._METAINFO_INDEX[ - "TOPK_SEND_INDICES_OFFSET_INDEX"]].item() - - send_counters = moe_a2a.workspace[ - rank, send_counters_offset:send_counters_offset + ep_size * 4].view( - torch.int32).cpu() - recv_counters = moe_a2a.workspace[ - rank, recv_counters_offset:recv_counters_offset + ep_size * 4].view( - torch.int32).cpu() - topk_target_ranks = moe_a2a.workspace[ - rank, topk_target_ranks_offset:topk_target_ranks_offset + - max_num_tokens * top_k * 4].view(torch.int32).view( - max_num_tokens, top_k).cpu() - topk_send_indices = moe_a2a.workspace[ - rank, topk_send_indices_offset:topk_send_indices_offset + - max_num_tokens * top_k * 4].view(torch.int32).view( - max_num_tokens, top_k).cpu() - - if actual_cft_dispatch: - cft_counters_after = read_cft_dispatch_counters( - moe_a2a, rank, ep_size) - cft_counter_delta = cft_counters_after - cft_counters_before - expected_payload_bytes_per_token = sum(payload_bytes_per_token) - for peer_rank in range(ep_size): - if peer_rank == rank: - continue - expected_bytes = (recv_counters[peer_rank].item() * - expected_payload_bytes_per_token) - actual_bytes = cft_counter_delta[peer_rank].item() - assert actual_bytes >= expected_bytes, ( - f"Rank {rank} CFT dispatch counter from rank {peer_rank}: " - f"delta={actual_bytes}, expected at least {expected_bytes}") - - # Return results to be collected (move to CPU for MPI transfer) - eplb_gathered_stats = moe_a2a._state.eplb_gathered_stats - if eplb_gathered_stats is not None: - eplb_gathered_stats = eplb_gathered_stats.cpu() - if eplb_local_stats is not None: - eplb_local_stats = eplb_local_stats.cpu() - - return (token_selected_experts.cpu(), [p.cpu() for p in payloads], - [rt.cpu() for rt in recv_tensors], send_counters, - topk_send_indices, topk_target_ranks, recv_counters, - expert_id_payload_index, eplb_local_stats, eplb_gathered_stats) - except Exception: - traceback.print_exc() - raise - - -def verify_dispatch(all_token_selected_experts, all_payloads, all_recv_tensors, - all_send_counters, all_topk_send_indices, - all_topk_target_ranks, all_recv_counters, ep_size, - all_num_tokens, top_k, num_experts, expert_id_payload_index, - invalid_token_expert_id): - """Verify dispatch results including actual content verification""" - - max_num_tokens = max(all_num_tokens) - num_experts_per_rank = num_experts // ep_size - # Verify dimensions and dtypes - for send_rank in range(ep_size): - local_num_tokens = all_num_tokens[send_rank] - - token_selected_experts = all_token_selected_experts[send_rank] - assert len(token_selected_experts.shape - ) == 2, "token_selected_experts should be a 2D tensor" - assert token_selected_experts.dtype == torch.int32, "token_selected_experts should be a 32-bit integer tensor" - assert token_selected_experts.shape[ - 0] == local_num_tokens, "token_selected_experts.shape[0] should be local_num_tokens" - assert token_selected_experts.shape[ - 1] == top_k, "token_selected_experts.shape[1] should be top_k" - - payloads = all_payloads[send_rank] - recv_tensors = all_recv_tensors[send_rank] - num_payloads = len(payloads) - assert len( - recv_tensors - ) == num_payloads, "recv_tensors should have the same number of payloads as payloads" - for i in range(num_payloads): - payload = payloads[i] - assert len(payload.shape) == 2, "payload should be a 2D tensor" - assert payload.shape[ - 0] == local_num_tokens, "payload.shape[0] should be local_num_tokens" - - recv_tensor = recv_tensors[i] - assert len( - recv_tensor.shape) == 3, "recv_tensor should be a 3D tensor" - assert recv_tensor.shape[ - 0] == ep_size, "recv_tensor.shape[0] should be ep_size" - assert recv_tensor.shape[ - 1] == max_num_tokens, "recv_tensor.shape[1] should be max_num_tokens" - assert recv_tensor.shape[2] == payload.shape[ - 1], "recv_tensor.shape[2] should be payload.shape[1]" - assert recv_tensor.dtype == payload.dtype, "recv_tensor.dtype should be payload.dtype" - - # Verify counters and compact routing tensors - send_counters = all_send_counters[send_rank] - assert len( - send_counters.shape) == 1, "send_counters should be a 1D tensor" - assert send_counters.shape[0] == ep_size - assert send_counters.dtype == torch.int32 - - recv_counters = all_recv_counters[send_rank] - assert len( - recv_counters.shape) == 1, "recv_counters should be a 1D tensor" - assert recv_counters.shape[0] == ep_size - assert recv_counters.dtype == torch.int32 - - topk_send_indices = all_topk_send_indices[send_rank] - topk_target_ranks = all_topk_target_ranks[send_rank] - assert topk_send_indices.shape == (max_num_tokens, - top_k), "topk_send_indices shape" - assert topk_target_ranks.shape == (max_num_tokens, - top_k), "topk_target_ranks shape" - assert topk_send_indices.dtype == torch.int32 - assert topk_target_ranks.dtype == torch.int32 - - # Verify send_counters per (send_rank -> target_rank) - for send_rank in range(ep_size): - expected_sends = {} - token_experts = all_token_selected_experts[send_rank] - sent_to_rank = set() - - for token_idx in range(token_experts.shape[0]): - experts = token_experts[token_idx] - target_ranks = compute_target_rank_id(experts, num_experts_per_rank) - sent_to_rank.clear() - - for target_rank in target_ranks.tolist(): - if target_rank not in sent_to_rank: - if target_rank not in expected_sends: - expected_sends[target_rank] = 0 - expected_sends[target_rank] += 1 - sent_to_rank.add(target_rank) - - for target_rank in range(ep_size): - expected_to_rank = expected_sends.get(target_rank, 0) - actual_to_rank = all_send_counters[send_rank][target_rank].item() - assert actual_to_rank == expected_to_rank, ( - f"Rank {send_rank} sent {actual_to_rank} tokens to rank {target_rank}, expected {expected_to_rank}" - ) - - # Verify recv_counters match send_counters - for recv_rank in range(ep_size): - for send_rank in range(ep_size): - expected_recv = all_send_counters[send_rank][recv_rank].item() - actual_recv = all_recv_counters[recv_rank][send_rank].item() - assert actual_recv == expected_recv, ( - f"Rank {recv_rank} received {actual_recv} tokens from rank {send_rank}, expected {expected_recv}" - ) - - # Verify payload content using topk_send_indices and topk_target_ranks - for send_rank in range(ep_size): - token_selected_experts = all_token_selected_experts[send_rank] - payloads = all_payloads[send_rank] - topk_send_indices = all_topk_send_indices[send_rank] - topk_target_ranks = all_topk_target_ranks[send_rank] - local_num_tokens = all_num_tokens[send_rank] - - for token_idx in range(local_num_tokens): - experts = token_selected_experts[token_idx] - target_ranks = compute_target_rank_id(experts, num_experts_per_rank) - # Deduplicate target ranks per token - topk_target_ranks_ref = target_ranks.clone() - seen = set() - for kk in range(top_k): - tr = int(topk_target_ranks_ref[kk].item()) - if tr in seen: - topk_target_ranks_ref[kk] = -1 - else: - seen.add(tr) - - assert topk_target_ranks[ - token_idx, :].tolist() == topk_target_ranks_ref.tolist() - - for k in range(top_k): - dst_pos = topk_send_indices[token_idx, k].item() - target_rank = topk_target_ranks[token_idx, k].item() - if dst_pos == -1: - assert target_rank == -1 - continue - recv_tensors = all_recv_tensors[target_rank] - for payload_idx, payload in enumerate(payloads): - recv_tensor = recv_tensors[payload_idx] - source_data = payload[token_idx] - received_data = recv_tensor[send_rank, dst_pos] - torch.testing.assert_close(received_data, - source_data, - atol=0, - rtol=0) - - # Verify token_selected_experts of invalid tokens are correctly sanitized - for recv_rank in range(ep_size): - expert_ids_recv = all_recv_tensors[recv_rank][expert_id_payload_index] - for source_rank in range(ep_size): - valid = int(all_recv_counters[recv_rank][source_rank].item()) - for token_idx in range(max_num_tokens): - token_expert_ids = expert_ids_recv[source_rank, token_idx] - if token_idx >= valid: - assert torch.all( - token_expert_ids == invalid_token_expert_id) - - -class TestMoEAlltoAll: - - @pytest.mark.skipif(torch.cuda.device_count() < 8, - reason='needs at least 8 GPUs to run multi-GPU test') - @pytest.mark.threadleak( - enabled=False - ) # MPI pool executors have known thread cleanup timing issues - @pytest.mark.parametrize( - "mpi_pool_executor,all_num_tokens,top_k,enable_eplb", - [ - # (num_workers, all_num_tokens, top_k) - # Basic configurations - (4, [32, 32, 32, 32], 2, False - ), # Four ranks with uniform distribution - (4, [16, 32, 64, 48 - ], 2, False), # Four ranks with non-uniform distribution - (2, [100, 50], 2, False), # Two ranks with different loads - (8, [10, 20, 30, 40, 50, 60, 70, 80 - ], 2, False), # Eight ranks with increasing load - - # Different top_k values - (4, [32, 32, 32, 32], 4, False), # Four ranks with top_k = 4 - (4, [32, 32, 32, 32], 8, False), # Four ranks with top_k = 8 - - # Edge cases - (4, [1, 1, 1, 1], 2, False - ), # Four ranks with single token per rank - - # EPLB stats path - (4, [32, 32, 32, 32], 2, True), - ], - indirect=["mpi_pool_executor"]) - def test_dispatch(self, mpi_pool_executor, all_num_tokens, top_k, - enable_eplb): - """Test MoE A2A dispatch with MNNVL across multiple GPUs""" - - try: - MnnvlMemory.initialize() - assert MnnvlMemory.supports_mnnvl() - except Exception: - pytest.skip("MNNVL not supported on this system") - - ep_size = mpi_pool_executor.num_workers - assert ep_size == len( - all_num_tokens), "ep_size does not match all_num_tokens" - - assert torch.cuda.device_count( - ) >= ep_size, f"Need at least {ep_size} GPUs, found {torch.cuda.device_count()}" - - hidden_size = 1024 - num_experts = 32 - - # Large enough workspace - workspace_size_per_rank = 512 * 1024 * 1024 - - invalid_token_expert_id = -1 - - # Run dispatch on workers - each worker executes the same logic as single-GPU - # but on separate GPUs with MNNVL memory instead of regular CUDA memory - results = mpi_pool_executor.map( - run_moe_a2a_dispatch_single_rank, - *zip(*[(ep_size, all_num_tokens, top_k, workspace_size_per_rank, - num_experts, hidden_size, invalid_token_expert_id, - enable_eplb)] * ep_size), - ) - - # Collect results from all ranks (same as single-GPU collecting from emulated ranks) - all_results = list(results) - - # Extract results in same format as single-GPU test - all_token_selected_experts = [r[0] for r in all_results] - all_payloads = [r[1] for r in all_results] - all_recv_tensors = [r[2] for r in all_results] - all_send_counters = [r[3] for r in all_results] - all_topk_send_indices = [r[4] for r in all_results] - all_topk_target_ranks = [r[5] for r in all_results] - all_recv_counters = [r[6] for r in all_results] - all_expert_id_payload_index = [r[7] for r in all_results] - expert_id_payload_index = all_expert_id_payload_index[0] - all_eplb_local_stats = [r[8] for r in all_results] - all_eplb_gathered_stats = [r[9] for r in all_results] - - assert all(i == expert_id_payload_index - for i in all_expert_id_payload_index - ), "all_expert_id_payload_index should be the same" - - # Verify dispatch results with content verification - verify_dispatch(all_token_selected_experts, all_payloads, - all_recv_tensors, all_send_counters, - all_topk_send_indices, all_topk_target_ranks, - all_recv_counters, ep_size, all_num_tokens, top_k, - num_experts, expert_id_payload_index, - invalid_token_expert_id) - - if enable_eplb: - expected_stats = torch.stack(all_eplb_local_stats, dim=0) - for rank in range(ep_size): - gathered_stats = all_eplb_gathered_stats[rank] - assert gathered_stats is not None - assert torch.equal( - gathered_stats, - expected_stats), (f"Rank {rank} gathered_stats mismatch") - - @pytest.mark.threadleak(enabled=False) - @pytest.mark.parametrize( - "mpi_pool_executor,all_num_tokens,top_k,payload_in_workspace,use_fp8_combine", - [ - # (num_workers, all_num_tokens, top_k, payload_in_workspace, use_fp8_combine) - (4, [32, 32, 32, 32], 2, False, False), - (4, [16, 32, 64, 48], 2, False, False), - (2, [100, 50], 2, False, False), - (4, [32, 32, 32, 32], 4, False, False), - (4, [32, 32, 32, 32 - ], 10, False, False), # top_k=10 used by Qwen3-next - (4, [1, 1, 1, 1], 2, False, False), - (8, [640, 640, 640, 640, 640, 640, 640, 640], 4, False, False), - (4, [32, 0, 16, 0], 2, False, False), - # payload_in_workspace=True - (4, [32, 32, 32, 32], 4, True, False), - (4, [32, 0, 16, 0], 4, True, False), - (4, [16, 32, 64, 48], 4, True, False), # non-uniform tokens - (4, [32, 32, 32, 32], 10, True, False), - # use_fp8_combine=True: staged quantization (external payload) - (4, [32, 32, 32, 32], 4, False, True), - (4, [32, 0, 16, 0], 4, False, True), - (4, [16, 32, 64, 48], 4, False, True), # non-uniform tokens - (4, [32, 32, 32, 32], 10, False, True), - # use_fp8_combine=True, payload_in_workspace=True: in-place quantization - (4, [32, 32, 32, 32], 4, True, True), - (4, [32, 0, 16, 0], 4, True, True), - (4, [16, 32, 64, 48], 4, True, True), # non-uniform tokens - (4, [32, 32, 32, 32], 10, True, True), - ], - indirect=["mpi_pool_executor"]) - def test_combine(self, mpi_pool_executor, all_num_tokens, top_k, - payload_in_workspace, use_fp8_combine): - """Test MoE A2A combine with MNNVL across multiple GPUs. - - When use_fp8_combine=True, runs two back-to-back rounds (BF16 reference then FP8) - and compares within FP8 rounding tolerance. When False, verifies against a - ground-truth fake-MoE computation. - """ - try: - MnnvlMemory.initialize() - assert MnnvlMemory.supports_mnnvl() - except Exception: - pytest.skip("MNNVL not supported on this system") - - ep_size = mpi_pool_executor.num_workers - if ep_size > torch.cuda.device_count(): - pytest.skip( - f"Need at least {ep_size} GPUs to run this test, but only {torch.cuda.device_count()} are available" - ) - assert ep_size == len( - all_num_tokens), "ep_size does not match all_num_tokens" - - # gpt-oss-20b - hidden_size = 2880 - num_experts = 32 - - # Large enough workspace - workspace_size_per_rank = 512 * 1024 * 1024 - - invalid_token_expert_id = -1 - results = mpi_pool_executor.map( - run_moe_a2a_dispatch_moe_combine_single_rank, - *zip(*[(ep_size, all_num_tokens, top_k, workspace_size_per_rank, - num_experts, hidden_size, invalid_token_expert_id, - payload_in_workspace, use_fp8_combine)] * ep_size), - ) - - try: - all_results = list(results) - except Exception: - traceback.print_exc() - raise - - if use_fp8_combine: - verify_combine(all_results, ep_size, rtol=0.13, atol=1.0) - else: - verify_combine(all_results, ep_size, rtol=0.1, atol=0.5) - - -def run_moe_a2a_dispatch_moe_combine_single_rank( - ep_size, - all_num_tokens, - top_k, - workspace_size_per_rank, - num_experts, - hidden_size, - invalid_token_expert_id, - payload_in_workspace=False, - use_low_precision_combine=False): - """Worker function for dispatch and combine test. - - Runs one dispatch+combine round and returns - (token_selected_experts, payloads, combined_output, rank_experts) for - ground-truth verification via verify_combine. - """ - rank = tllm.mpi_rank() - torch.cuda.set_device(rank) - device = torch.cuda.current_device() - max_num_tokens = max(all_num_tokens) - rank_local_tokens = all_num_tokens[rank] - - try: - mapping = Mapping(rank=rank, - tp_size=ep_size, - moe_ep_size=ep_size, - world_size=ep_size) - - moe_a2a = MoeAlltoAll( - mapping=mapping, - max_num_tokens=max_num_tokens, - top_k=top_k, - num_slots=num_experts, - workspace_size_per_rank=workspace_size_per_rank, - ) - - token_selected_experts = generate_token_selected_experts( - rank_local_tokens, num_experts, top_k) - payloads, expert_id_payload_index = make_bfloat16_payloads( - rank_local_tokens, hidden_size, top_k, rank, token_selected_experts) - - num_experts_per_rank = num_experts // ep_size - rank_experts = create_experts_per_rank(num_experts_per_rank, - hidden_size, - rank, - device, - dtype=torch.bfloat16) - - def dispatch_and_fake_moe(): - """Run one dispatch round and return fake-MoE output [ep_size, max_tokens, hidden].""" - recv_tensors = moe_a2a.dispatch( - token_selected_experts, - payloads, - max_num_tokens, - invalid_token_expert_id=invalid_token_expert_id, - expert_id_payload_index=expert_id_payload_index) - hs, tse, tfs = recv_tensors[0], recv_tensors[1], recv_tensors[2] - moe_out = fake_moe( - hs.view(ep_size * max_num_tokens, hs.shape[-1]), - tse.view(ep_size * max_num_tokens, tse.shape[-1]), - tfs.view(ep_size * max_num_tokens, tfs.shape[-1]), - rank_experts, - is_ep=True, - ep_rank=rank, - num_experts_per_rank=num_experts_per_rank, - ) - return moe_out.view(ep_size, max_num_tokens, hs.shape[-1]) - - def _combine(moe_out, use_low_precision): - """Call combine, optionally staging moe_out via workspace buffer.""" - if payload_in_workspace: - ws = moe_a2a.get_combine_payload_tensor_in_workspace( - max_num_tokens, hidden_size, torch.bfloat16) - ws.copy_(moe_out.view(-1, hidden_size)) - return moe_a2a.combine( - ws.view(ep_size, max_num_tokens, hidden_size), - max_num_tokens, - payload_in_workspace=True, - use_low_precision_combine=use_low_precision, - ) - return moe_a2a.combine(moe_out, - max_num_tokens, - use_low_precision_combine=use_low_precision) - - moe_out = dispatch_and_fake_moe() - combined_output = _combine(moe_out, - use_low_precision=use_low_precision_combine) - - wire_bytes_per_token = hidden_size * (1 if use_low_precision_combine - else moe_out.element_size()) - actual_cft_combine = (moe_a2a.can_use_cft_counted_writes - and moe_a2a.use_cft_for_combine(max_num_tokens) - and wire_bytes_per_token % 16 == 0) - - if not actual_cft_combine: - completion_flags_offset = moe_a2a.metainfo[ - MoeAlltoAll._METAINFO_INDEX[ - "COMBINE_COMPLETION_FLAGS_OFFSET_INDEX"]].item() - completion_flags = moe_a2a.workspace[ - rank, completion_flags_offset:completion_flags_offset + - ep_size * 4].view(torch.int32).cpu() - flag_val_offset = moe_a2a.metainfo[ - MoeAlltoAll._METAINFO_INDEX["FLAG_VAL_OFFSET_INDEX"]].item() - expected_flag_val = moe_a2a.workspace[ - rank, - flag_val_offset:flag_val_offset + 4].view(torch.int32).cpu() - assert torch.all(completion_flags == expected_flag_val), ( - f"Rank {rank} completion flags: {completion_flags}, expected flag val: {expected_flag_val}" - ) - - return ( - token_selected_experts.cpu(), - [p.cpu() for p in payloads], - combined_output.cpu(), - rank_experts.cpu(), - ) - except Exception: - traceback.print_exc() - raise - - -def verify_combine(all_results, ep_size, rtol, atol): - """Verify that combine correctly sums the dispatched tokens.""" - - # Extract results - all_token_selected_experts = [r[0] for r in all_results] - all_original_payloads = [r[1] for r in all_results] - all_combined_outputs = [r[2] for r in all_results] - all_rank_experts = [r[3] - for r in all_results] # Extract experts from each rank - - # For each rank, verify the combined output - for rank in range(ep_size): - # print("### Verify rank %d ###" % rank) - token_selected_experts = all_token_selected_experts[rank] - original_payloads = all_original_payloads[rank] - hidden_states = original_payloads[0] - token_final_scales = original_payloads[2] - - combined_output = all_combined_outputs[rank] - - # Check the following are equal: - # expected: Directly emulate MoE with all experts as if EP is not used. - # actual: Tokens are dispatched to target ranks, MoE is performed on target ranks, and then the results from all target ranks are summed up (combine). - - # Gather all experts from all ranks for non-EP emulation - all_experts = torch.cat(all_rank_experts, dim=0) - expected_combined_output = fake_moe(hidden_states, - token_selected_experts, - token_final_scales, - all_experts, - is_ep=False) - - # Custom assertion with detailed error message - try: - torch.testing.assert_close(combined_output, - expected_combined_output, - rtol=rtol, - atol=atol) - except AssertionError as e: - # Find the first mismatch location - abs_diff = (combined_output - expected_combined_output).abs() - rel_diff = abs_diff / (expected_combined_output.abs() + 1e-8) - - # Check both absolute and relative tolerance - mask = (abs_diff > atol) & (rel_diff > rtol) - if mask.any(): - # Get the first mismatch - mismatch_indices = torch.nonzero(mask)[0].tolist() - token_idx, elem_idx = mismatch_indices - - # Build context visualization - context_values_expected = [] - context_values_actual = [] - - for offset in [-2, -1, 0, 1, 2]: - idx = elem_idx + offset - if 0 <= idx < combined_output.shape[1]: - context_values_expected.append( - f"{expected_combined_output[token_idx, idx].item():.4f}" - ) - context_values_actual.append( - f"{combined_output[token_idx, idx].item():.4f}") - else: - context_values_expected.append("-") - context_values_actual.append("-") - - # Add ... to indicate continuation - expected_str = ' '.join(context_values_expected) - actual_str = ' '.join(context_values_actual) - - # Add ... on left if not at beginning - if elem_idx > 2: - expected_str = "... " + expected_str - actual_str = "... " + actual_str - - # Add ... on right if not at end - if elem_idx < combined_output.shape[1] - 3: - expected_str = expected_str + " ..." - actual_str = actual_str + " ..." - - error_msg = f"\nexpected: [{expected_str}]\n" - error_msg += f"actual: [{actual_str}]\n" - error_msg += f"\n{str(e)}" - - raise AssertionError(error_msg) diff --git a/tests/unittest/_torch/moe/multi_gpu/test_moe_a2a_workspace.py b/tests/unittest/_torch/moe/multi_gpu/test_moe_a2a_workspace.py index a857fb587c76..604a288d60fa 100644 --- a/tests/unittest/_torch/moe/multi_gpu/test_moe_a2a_workspace.py +++ b/tests/unittest/_torch/moe/multi_gpu/test_moe_a2a_workspace.py @@ -15,7 +15,7 @@ from mpi4py.futures import MPIPoolExecutor import tensorrt_llm as tllm -from tensorrt_llm._mnnvl_utils import MnnvlMemory +from tensorrt_llm._torch.distributed.mnnvl_memory import MnnvlMemory from tensorrt_llm._torch.moe.fused_moe.communication.nvlink_one_sided import NVLinkOneSided from tensorrt_llm.mapping import Mapping @@ -93,10 +93,10 @@ def _run_worker(capture, in_workspace, use_cft, low_precision): with pytest.MonkeyPatch.context() as patch: # Pin worker-side policy rather than inheriting user/CI overrides. for name in ( - "TRTLLM_MOE_A2A_FORCE_CFT", - "TRTLLM_MOE_A2A_CFT_MAX_BATCH_FOR_DISPATCH", - "TRTLLM_MOE_A2A_CFT_MAX_BATCH_FOR_COMBINE", - "TRTLLM_MOE_A2A_WORKSPACE_MB", + "TRTLLM_NVLINK_ONE_SIDED_A2A_FORCE_CFT", + "TRTLLM_NVLINK_ONE_SIDED_A2A_CFT_MAX_BATCH_FOR_DISPATCH", + "TRTLLM_NVLINK_ONE_SIDED_A2A_CFT_MAX_BATCH_FOR_COMBINE", + "TRTLLM_NVLINK_ONE_SIDED_A2A_WORKSPACE_MB", ): patch.delenv(name, raising=False) try: diff --git a/tests/unittest/_torch/moe/multi_gpu/test_nvlink_one_sided.py b/tests/unittest/_torch/moe/multi_gpu/test_nvlink_one_sided.py new file mode 100644 index 000000000000..66ee73f0472e --- /dev/null +++ b/tests/unittest/_torch/moe/multi_gpu/test_nvlink_one_sided.py @@ -0,0 +1,837 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +"""NVLinkOneSided path selection and model-shaped round trips with independent references.""" + +from __future__ import annotations + +import pickle +import sys +from collections.abc import Iterator +from dataclasses import dataclass, replace +from enum import Enum +from typing import Literal +from unittest.mock import patch + +import cloudpickle +import pytest +import torch +from mpi4py import MPI +from mpi4py.futures import MPIPoolExecutor + +from tensorrt_llm._torch.distributed.mnnvl_memory import MnnvlMemory +from tensorrt_llm._torch.moe.fused_moe.communication.nvlink_one_sided import ( + FORCE_CFT_ENV, + NVLinkOneSided, + _cft_device_support_reason, + _get_nvidia_driver_version, + cft_driver_is_supported, +) +from tensorrt_llm.mapping import Mapping + +cloudpickle.register_pickle_by_value(sys.modules[__name__]) +MPI.pickle.__init__(cloudpickle.dumps, cloudpickle.loads, pickle.HIGHEST_PROTOCOL) +pytestmark = pytest.mark.threadleak(enabled=False) + + +@dataclass(frozen=True) +class ModelShape: + name: str + hidden_size: int + num_experts: int + top_k: int + # activation format of MoE for post-quant dispatch + dispatch_dtype: Literal["bf16", "blockwise_fp8", "mxfp8", "nvfp4"] + + +MODELS = { + model.name: model + for model in ( + ModelShape("gpt_oss", 2880, 128, 4, "mxfp8"), + ModelShape("deepseek_v3", 7168, 256, 8, "blockwise_fp8"), + ModelShape("deepseek_r1_nvfp4", 7168, 256, 8, "nvfp4"), + ModelShape("qwen3_5_397b_a17b", 4096, 512, 10, "bf16"), + ModelShape("deepseek_v4_flash", 4096, 256, 6, "mxfp8"), + ModelShape("deepseek_v4_pro", 7168, 384, 6, "mxfp8"), + ModelShape( + "kimi_k3", 3584, 896, 16, "mxfp8" + ), # K3 communicates its latent MoE width, not the model's 7168-wide states. + ) +} + + +def _expert_bounds(num_experts: int, ep_size: int, rank: int) -> tuple[int, int]: + base, remainder = divmod(num_experts, ep_size) + start = rank * base + min(rank, remainder) + return start, start + base + int(rank < remainder) + + +class Routing(Enum): + """Synthetic expert routing patterns for communication tests. + + SPREAD distributes routes across all ranks, LOCAL selects the source + rank's experts, and HOTSPOT selects only rank 0's experts. + """ + + SPREAD = "spread" + LOCAL = "local" + HOTSPOT = "hotspot" + + def make_expert_ids( + self, + model: ModelShape, + tokens: tuple[int, ...], + rank: int, + round_index: int, + device: torch.device | str = "cuda", + ) -> torch.Tensor: + """Build int32 global expert IDs shaped [tokens[rank], model.top_k].""" + ep_size = len(tokens) + count = tokens[rank] + token = torch.arange(count, device=device, dtype=torch.int64)[:, None] + choice = torch.arange(model.top_k, device=device, dtype=torch.int64)[None, :] + if self is Routing.LOCAL: + owners = torch.full((count, model.top_k), rank, device=device, dtype=torch.int64) + elif self is Routing.HOTSPOT: + owners = torch.zeros((count, model.top_k), device=device, dtype=torch.int64) + else: + owners = (token + choice + rank + round_index) % ep_size + bounds = [_expert_bounds(model.num_experts, ep_size, r) for r in range(ep_size)] + starts = torch.tensor([b[0] for b in bounds], device=device) + sizes = torch.tensor([b[1] - b[0] for b in bounds], device=device) + return (starts[owners] + (token + choice + 3 * round_index) % sizes[owners]).int() + + +@dataclass(frozen=True) +class Round: + """One dispatch → simulated expert computation → combine iteration. + + Attributes: + tokens: Local token counts indexed by EP rank; zero-token ranks still + participate in communication. + routing: Strategy used to construct expert IDs; see Routing. + delay_rank: Rank to delay on the GPU after dispatch, before expert + computation, to exercise rank skew. -1 disables the delay. + """ + + tokens: tuple[int, ...] + routing: Routing + delay_rank: int + + +@dataclass(frozen=True) +class Case: + """A workload and execution settings sharing one communicator across rounds. + + Attributes: + model: Expert layout, hidden width, and dispatch payload format. + rounds: Ordered iterations with no host synchronization between them. + All token-count tuples must describe the same EP group size. + mode: Automatic CFT selection, forced fence, or requested CFT; platform + and payload restrictions still apply to CFT. + payload_in_workspace: Copy simulated expert outputs into the communication workspace + before combine, instead of passing an external tensor. + fp8_combine: Use FP8 on the combine wire; the returned output remains BF16. + graph: Capture and replay the entire round sequence in one CUDA graph. + pdl: Enable programmatic dependent launch for communication kernels. + eplb: Gather and verify expert-load statistics alongside dispatch. + """ + + model: ModelShape + rounds: tuple[Round, ...] + mode: Literal["auto", "fence", "cft"] + payload_in_workspace: bool + fp8_combine: bool + graph: bool + pdl: bool + eplb: bool + + @property + def ep_size(self) -> int: + """Number of participating EP workers.""" + return len(self.rounds[0].tokens) + + @property + def runtime_max_num_tokens_per_rank(self) -> int: + """Allocation-time per-rank token limit covering every round.""" + return max(max(r.tokens) for r in self.rounds) + + +def _make_inputs(case: Case, rank: int, round_index: int) -> tuple[torch.Tensor, ...]: + import torch + + model = case.model + spec = case.rounds[round_index] + count = spec.tokens[rank] + generator = torch.Generator(device="cuda").manual_seed(1234 + rank * 97 + round_index) + # Positive, exactly representable values keep cancellation and saturation out + # of the communication reference. Scales vary by row and channel block. + x = torch.randint(1, 17, (max(count, 1), model.hidden_size), generator=generator, device="cuda") + row_scale = 2.0 ** ((torch.arange(max(count, 1), device="cuda") + round_index) % 3) + block_size = 128 if model.dispatch_dtype == "blockwise_fp8" else 32 + block_scale = 2.0 ** ((torch.arange(model.hidden_size, device="cuda") // block_size) % 3) + x = (x.float() / 16 * row_scale[:, None] * block_scale[None, :]).to(torch.bfloat16) + if model.dispatch_dtype == "blockwise_fp8": + from tensorrt_llm.quantization.utils.fp8_utils import fp8_quantize_1x128_sf_transpose + + payload, sf = fp8_quantize_1x128_sf_transpose(x, use_ue8m0=False) + # Dispatch requires contiguous token-major scales, not GEMM's column-major layout. + sf = sf[:count].contiguous() + assert payload.dtype == torch.float8_e4m3fn + assert sf.dtype == torch.float32 + assert sf.shape == (count, model.hidden_size // 128) + elif model.dispatch_dtype == "mxfp8": + payload, sf = torch.ops.trtllm.mxfp8_quantize(x, False, 32) + sf = sf.view(x.shape[0], -1)[:count] + elif model.dispatch_dtype == "nvfp4": + global_scale = torch.ones((), dtype=torch.float32, device="cuda") + payload, sf = torch.ops.trtllm.fp4_quantize(x, global_scale, 16, False, False) + sf = sf.view(x.shape[0], -1)[:count] + else: + payload, sf = x, None + payload = payload[:count] + + slots = spec.routing.make_expert_ids(model, spec.tokens, rank, round_index) + token = torch.arange(count, device="cuda", dtype=torch.int64)[:, None] + choice = torch.arange(model.top_k, device="cuda", dtype=torch.int64)[None, :] + # The first routing weight also identifies the source token independently + # of the implementation's compact send indices and receive counters. + identity = rank * case.runtime_max_num_tokens_per_rank + token + 1 + weights = (identity * (choice + 1)).float() / ( + case.runtime_max_num_tokens_per_rank * case.ep_size * 32 + ) + return payload, sf, slots, weights + + +def _dequantize(payload: torch.Tensor, sf: torch.Tensor | None, mode: str) -> torch.Tensor: + if mode == "bf16": + return payload.float() + if mode == "blockwise_fp8": + values = payload.view(torch.float8_e4m3fn).float() + blocks = values.reshape(values.shape[0], values.shape[1] // 128, 128) + return (blocks * sf[..., None]).flatten(1) + if mode == "mxfp8": + values = payload.view(torch.float8_e4m3fn).float() + exponents = sf.view(torch.uint8).int() - 127 + return torch.ldexp( + values.reshape(values.shape[0], values.shape[1] // 32, 32), exponents[..., None] + ).flatten(1) + packed = payload.view(torch.uint8) + codes = torch.stack((packed & 15, packed >> 4), dim=-1).flatten(1).long() + # Decode E2M1 on-device without a host lookup-table copy during graph capture. + mantissa = (codes & 1).float() + exponent = (codes >> 1) & 3 + magnitude = torch.where( + exponent == 0, mantissa * 0.5, torch.ldexp(1.0 + mantissa * 0.5, exponent - 1) + ) + values = torch.where(codes < 8, magnitude, -magnitude) + scales = sf.view(torch.float8_e4m3fn).float().repeat_interleave(16, dim=-1) + return values * scales + + +def _expert_output( + payload: torch.Tensor, + sf: torch.Tensor | None, + slots: torch.Tensor, + weights: torch.Tensor, + case: Case, + rank: int, +) -> torch.Tensor: + start, end = _expert_bounds(case.model.num_experts, case.ep_size, rank) + owned = (slots >= start) & (slots < end) + gain = torch.where(owned, weights * (1.0 + (slots % 7).float() / 8), 0).sum(dim=-1) + values = _dequantize(payload, sf, case.model.dispatch_dtype) + # Padded receive slots have invalid expert IDs and unspecified payload bytes. + values = torch.where(owned.any(dim=-1)[:, None], values, 0) + return (values * gain[:, None]).to(torch.bfloat16) + + +def _cpu(tensor: torch.Tensor | None) -> torch.Tensor | None: + if tensor is None: + return None + # Transport FP8 as raw bytes: its storage is not supported by the legacy + # torch serialization used by MPI pickle. References reinterpret the bytes. + if tensor.dtype == torch.float8_e4m3fn: + tensor = tensor.view(torch.uint8) + return tensor.cpu() + + +def _run_worker(case: Case) -> dict: + # Import locally so cloudpickle does not serialize torch's dynamic ops namespace. + import os + + import torch + + # Selection is read by each communicator constructor, not cached by the MPI pool. + os.environ[FORCE_CFT_ENV] = {"auto": "", "fence": "0", "cft": "1"}[case.mode] + + rank = MPI.COMM_WORLD.Get_rank() + torch.cuda.set_device(rank) + MnnvlMemory.initialize() + supported = MnnvlMemory.supports_mnnvl() + cft_reason = None + if case.mode != "fence": + if not cft_driver_is_supported(_get_nvidia_driver_version()): + cft_reason = "CFT requires driver 615 or newer" + else: + cft_reason = _cft_device_support_reason() + sm_major = torch.cuda.get_device_capability()[0] + quant_supported = ( + case.model.dispatch_dtype == "bf16" + or (case.model.dispatch_dtype == "blockwise_fp8" and sm_major >= 9) + or sm_major >= 10 + ) + reasons = MPI.COMM_WORLD.allgather((supported, cft_reason, quant_supported)) + if not all(item[0] for item in reasons): + return {"skip": "NVLink one-sided is not supported on every participating GPU"} + if case.mode == "cft" and any(item[1] for item in reasons): + return {"skip": str(reasons)} + if not all(item[2] for item in reasons): + return { + "skip": f"{case.model.dispatch_dtype} payload generation is unsupported on a participating GPU" + } + if case.mode == "auto" and len({item[1] is None for item in reasons}) != 1: + return {"skip": "automatic CFT selection requires consistent capability across ranks"} + + mapping = Mapping( + rank=rank, world_size=case.ep_size, tp_size=case.ep_size, moe_ep_size=case.ep_size + ) + comm = NVLinkOneSided( + mapping=mapping, + num_slots=case.model.num_experts, + top_k=case.model.top_k, + max_num_tokens_per_rank=case.runtime_max_num_tokens_per_rank, + hidden_size=case.model.hidden_size, + dtype=torch.bfloat16, + payload_in_workspace=case.payload_in_workspace, + use_low_precision_combine=case.fp8_combine, + num_experts=case.model.num_experts // 2 if case.eplb else None, + ) + try: + inputs = [_make_inputs(case, rank, i) for i in range(len(case.rounds))] + stats = None + if case.eplb: + stats = ( + torch.arange(case.model.num_experts // 2, dtype=torch.int32, device="cuda") + + rank * 1000 + ) + dispatch_op = torch.ops.trtllm.moe_a2a_dispatch + combine_op = torch.ops.trtllm.moe_a2a_combine + launches = [] + combine_offsets = [] + + def record_dispatch(*args, **kwargs): + # Record the actual op argument, not just the wrapper's capability. + launches.append(("dispatch", bool(args[10]))) + result = dispatch_op(*args, **kwargs) + combine_offsets.append(result[1]) + return result + + def record_combine(*args, **kwargs): + launches.append(("combine", bool(args[11]))) + return combine_op(*args, **kwargs) + + def sequence() -> list[dict]: + outputs = [] + for i, (spec, tensors) in enumerate(zip(case.rounds, inputs, strict=True)): + payload, sf, slots, weights = tensors + counts = list(spec.tokens) + comm.prepare_dispatch(slots, counts) + recv = comm.dispatch(payload, sf, slots, weights, counts, eplb_local_stats=stats) + # Combine clears dispatch state; consume statistics at the same + # point as the MoE scheduler, before starting expert computation. + eplb = comm.get_eplb_gathered_statistics().clone() if case.eplb else None + if rank == spec.delay_rank: + torch.cuda._sleep(200_000) + # Detailed dispatch snapshots are omitted in the race sequence; + # even device copies can perturb the overlap being exercised. + snapshot = ( + tuple(None if t is None else t.clone() for t in recv) + if len(inputs) == 1 + else None + ) + expert_out = _expert_output(*recv, case, rank) + if case.payload_in_workspace: + workspace_out = comm.get_combine_payload_tensor_in_workspace( + max(counts), case.model.hidden_size, torch.bfloat16 + ) + workspace_out.view_as(expert_out).copy_(expert_out) + expert_out = workspace_out + combined = comm.combine(expert_out, all_rank_max_num_tokens=max(counts)) + outputs.append({"combined": combined.clone(), "dispatch": snapshot, "eplb": eplb}) + return outputs + + with ( + patch.object(torch.ops.trtllm, "moe_a2a_dispatch", record_dispatch), + patch.object(torch.ops.trtllm, "moe_a2a_combine", record_combine), + ): + # No host reads, barriers, or synchronizations between rounds. + MPI.COMM_WORLD.Barrier() + if case.graph: + sequence() + torch.cuda.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + outputs = sequence() + launches = launches[-2 * len(inputs) :] + for _ in range(3): + graph.replay() + else: + outputs = sequence() + torch.cuda.synchronize() + + # Region boundaries must stay fixed across runtime token counts and paths. + assert len(set(combine_offsets)) == 1, combine_offsets + result = { + "rank": rank, + "inputs": [tuple(_cpu(t) for t in tensors) for tensors in inputs], + "outputs": [ + { + "combined": _cpu(o["combined"]), + "dispatch": None + if o["dispatch"] is None + else tuple(_cpu(t) for t in o["dispatch"]), + "eplb": _cpu(o["eplb"]), + } + for o in outputs + ], + "launches": launches, + "cft_capable": comm.can_use_cft_counted_writes, + } + MPI.COMM_WORLD.Barrier() + return result + finally: + comm.destroy() + + +def verify_dispatch(case: Case, results: list[dict], receiver: int, round_index: int) -> None: + """Check received tokens, payload bytes, routing metadata, and padding against original inputs.""" + payload, sf, slots, weights = results[receiver]["outputs"][round_index]["dispatch"] + maximum = max(case.rounds[round_index].tokens) + payload = payload.view(case.ep_size, maximum, -1) + sf = None if sf is None else sf.view(case.ep_size, maximum, -1) + slots = slots.view(case.ep_size, maximum, -1) + weights = weights.view(case.ep_size, maximum, -1) + begin, end = _expert_bounds(case.model.num_experts, case.ep_size, receiver) + for source in range(case.ep_size): + original, original_sf, original_slots, original_weights = results[source]["inputs"][ + round_index + ] + wanted = ((original_slots >= begin) & (original_slots < end)).any(dim=1) + expected_ids = torch.where(wanted)[0] + valid = (slots[source] >= 0).any(dim=1) + observed_ids = ( + (weights[source, valid, 0] * (case.runtime_max_num_tokens_per_rank * case.ep_size * 32)) + .round() + .long() + - source * case.runtime_max_num_tokens_per_rank + - 1 + ) + order = observed_ids.argsort() + torch.testing.assert_close(observed_ids[order], expected_ids, rtol=0, atol=0) + torch.testing.assert_close( + slots[source, ~valid], torch.full_like(slots[source, ~valid], -1) + ) + torch.testing.assert_close( + slots[source, valid][order], original_slots[expected_ids], rtol=0, atol=0 + ) + torch.testing.assert_close( + weights[source, valid][order], original_weights[expected_ids], rtol=0, atol=0 + ) + torch.testing.assert_close( + payload[source].view(torch.uint8)[valid][order], + original.view(torch.uint8)[expected_ids], + rtol=0, + atol=0, + ) + if sf is not None: + torch.testing.assert_close( + sf[source].view(torch.uint8)[valid][order], + original_sf.view(torch.uint8)[expected_ids], + rtol=0, + atol=0, + ) + + +def verify_combine(case: Case, results: list[dict], rank: int, round_index: int) -> None: + """Check combined output against original inputs, including per-rank BF16/FP8 rounding.""" + payload, sf, slots, weights = results[rank]["inputs"][round_index] + spec = case.rounds[round_index] + reference = torch.zeros((spec.tokens[rank], case.model.hidden_size), dtype=torch.float32) + for expert_rank in range(case.ep_size): + contribution = _expert_output(payload, sf, slots, weights, case, expert_rank) + if case.fp8_combine: + contribution = contribution.float().clamp(-448, 448).to(torch.float8_e4m3fn) + reference += contribution.float() + output = results[rank]["outputs"][round_index]["combined"] + assert output.dtype == torch.bfloat16 + torch.testing.assert_close( + output, + reference.to(torch.bfloat16), + rtol=0.016, + atol=0.002, + msg=lambda detail: ( + f"rank={rank}, round={round_index}, tokens={spec.tokens}, " + f"mode={case.mode}, graph={case.graph}\n{detail}" + ), + ) + + +def _assert_results(case: Case, results: list[dict]) -> None: + for rank in range(case.ep_size): + result = results[rank] + assert len(result["outputs"]) == len(case.rounds) + expected_launches = [] + for i, spec in enumerate(case.rounds): + payload, sf, slots, weights = result["inputs"][i] + verify_combine(case, results, rank, i) + if result["outputs"][i]["dispatch"] is not None: + verify_dispatch(case, results, rank, i) + requested = ( + result["cft_capable"] + and case.mode != "fence" + and (case.mode == "cft" or max(spec.tokens) <= 128) + ) + # Calculate wire eligibility independently from the production helper. + payloads = (payload, slots, weights) if sf is None else (payload, sf, slots, weights) + dispatch_cft = requested and all( + t.shape[1] * t.element_size() % 16 == 0 for t in payloads + ) + combine_cft = ( + requested and case.model.hidden_size * (1 if case.fp8_combine else 2) % 16 == 0 + ) + expected_launches.extend((("dispatch", dispatch_cft), ("combine", combine_cft))) + if case.eplb: + expected = torch.stack( + [ + torch.arange(case.model.num_experts // 2, dtype=torch.int32) + r * 1000 + for r in range(case.ep_size) + ] + ) + torch.testing.assert_close(result["outputs"][i]["eplb"], expected, rtol=0, atol=0) + assert result["launches"] == expected_launches + + +@pytest.fixture(scope="module") +def mpi_pools() -> Iterator[dict[tuple[int, bool], MPIPoolExecutor]]: + """Reuse imported workers; each case still creates and destroys its communicator.""" + pools: dict[tuple[int, bool], MPIPoolExecutor] = {} + try: + yield pools + finally: + for pool in pools.values(): + pool.shutdown(wait=True, cancel_futures=True) + + +def _run(case: Case, pools: dict[tuple[int, bool], MPIPoolExecutor]) -> None: + if torch.cuda.device_count() < case.ep_size: + pytest.skip(f"requires {case.ep_size} GPUs") + # PDL may be cached in native code, so different settings need separate workers. + # CFT mode is Python-side and is set on every worker before constructing the comm. + env = { + "TRTLLM_NVLINK_ONE_SIDED_A2A_CFT_MAX_BATCH_FOR_DISPATCH": "128", + "TRTLLM_NVLINK_ONE_SIDED_A2A_CFT_MAX_BATCH_FOR_COMBINE": "128", + "TRTLLM_NVLINK_ONE_SIDED_A2A_WORKSPACE_MB": "", + "TRTLLM_ENABLE_PDL": "1" if case.pdl else "0", + "TRTLLM_NVLINK_ONE_SIDED_A2A_TIMEOUT_SEC": "30", + "TRTLLM_NVLINK_ONE_SIDED_A2A_WARMUP_TIMEOUT_SEC": "60", + } + key = (case.ep_size, case.pdl) + if key not in pools: + pools[key] = MPIPoolExecutor(case.ep_size, env=env) + executor = pools[key] + healthy = False + try: + results = list(executor.map(_run_worker, [case] * case.ep_size)) + skipped = [r["skip"] for r in results if "skip" in r] + if skipped: + assert len(skipped) == case.ep_size, "inconsistent platform support across ranks" + pytest.skip(skipped[0]) + # Task submission order is independent of the MPI rank that executes it. + results.sort(key=lambda result: result["rank"]) + assert [result["rank"] for result in results] == list(range(case.ep_size)) + _assert_results(case, results) + healthy = True + finally: + # A failed collective or reference check must not contaminate another case. + if not healthy: + pools.pop(key).shutdown(wait=True, cancel_futures=True) + + +CASES = [ + # Model coverage uses equal token counts on every rank. + *[ + pytest.param( + Case( + model=model, + rounds=( + Round(tokens=(num_tokens,) * ep_size, routing=Routing.SPREAD, delay_rank=-1), + ), + mode="auto", + payload_in_workspace=False, + fp8_combine=False, + graph=False, + pdl=True, + eplb=False, + ), + id=f"model-{model.name}-ep{ep_size}-tokens{num_tokens}", + ) + for model, ep_size in ( + (MODELS["gpt_oss"], 4), + (MODELS["deepseek_v3"], 8), + (MODELS["deepseek_r1_nvfp4"], 8), + (MODELS["qwen3_5_397b_a17b"], 8), + (MODELS["deepseek_v4_flash"], 4), + (MODELS["deepseek_v4_pro"], 8), + (MODELS["kimi_k3"], 8), + ) + for num_tokens in (1, 128, 1024) + ], + # Cover BF16/FP8 combine with external/workspace-resident input payloads. + *[ + pytest.param( + Case( + model=MODELS["deepseek_v3"], + rounds=(Round(tokens=(9, 5), routing=Routing.SPREAD, delay_rank=-1),), + mode="auto", + payload_in_workspace=payload_in_workspace, + fp8_combine=fp8_combine, + graph=False, + pdl=True, + eplb=False, + ), + id=( + f"combine-{'fp8' if fp8_combine else 'bf16'}-" + f"{'workspace' if payload_in_workspace else 'external'}" + ), + ) + for fp8_combine in (False, True) + for payload_in_workspace in (False, True) + ], + # Reuse one four-rank communicator across changing token counts and peer dependencies, + # including zero-token ranks and delayed ranks. Counts 128/129 straddle the + # automatic CFT threshold; graph cases replay the entire sequence. + *[ + pytest.param( + Case( + model=MODELS["deepseek_v3"], + rounds=( + Round((3, 0, 2, 1), Routing.LOCAL, 0), + Round((1, 129, 0, 5), Routing.SPREAD, 1), + Round((0, 5, 3, 1), Routing.HOTSPOT, 2), + Round((128, 3, 1, 0), Routing.SPREAD, 3), + Round((5, 2, 0, 3), Routing.LOCAL, 0), + Round((1, 0, 4, 2), Routing.SPREAD, 1), + # Local-only routing gives rank 0 no token dependency on rank 1. + # Combine on rank 0 must still wait for rank 1 until rank 1 consumes its dispatched inputs. + # Otherwise, rank 0's next dispatch would abrupt the data that rank 1 is consuming. + Round((1, 128, 0, 0), Routing.LOCAL, 1), + Round((129, 1, 0, 0), Routing.SPREAD, 0), + ), + mode=mode, + payload_in_workspace=True, + fp8_combine=False, + graph=graph, + pdl=True, + eplb=False, + ), + id=f"round-sequence-ep4-{mode}-{'graph' if graph else 'eager'}", + ) + for mode, graph in ( + ("auto", False), + ("auto", True), + ("cft", True), + ("fence", True), + ) + ], + # Gather rank-distinct EPLB statistics from all four ranks, including the rank + # with zero input tokens. This checks statistics transport, not expert migration. + pytest.param( + Case( + model=MODELS["deepseek_v3"], + rounds=(Round(tokens=(9, 0, 3, 1), routing=Routing.SPREAD, delay_rank=-1),), + mode="auto", + payload_in_workspace=False, + fp8_combine=False, + graph=False, + pdl=True, + eplb=True, + ), + id="eplb-statistics", + ), + # Use 129 experts with the DeepSeek V3 payload shape, split across two ranks (65/64), + # to exercise remainder-aware ownership in dispatch and the combine reference. + pytest.param( + Case( + model=replace(MODELS["deepseek_v3"], num_experts=129), + rounds=(Round(tokens=(9, 3), routing=Routing.SPREAD, delay_rank=-1),), + mode="auto", + payload_in_workspace=False, + fp8_combine=False, + graph=False, + pdl=True, + eplb=False, + ), + id="non-divisible-experts", + ), +] + + +@pytest.mark.parametrize("case", CASES) +def test_nvlink_one_sided(case: Case, mpi_pools: dict[tuple[int, bool], MPIPoolExecutor]) -> None: + _run(case, mpi_pools) + + +# ============================================================================ +# Allocation-time workspace layout (no MPI workers) +# ============================================================================ + + +@pytest.mark.parametrize("cft,fp8_combine", [(False, False), (True, False), (True, True)]) +def test_workspace_layout(cft: bool, fp8_combine: bool) -> None: + from tensorrt_llm.bindings import internal as _tllm_internal + + thop = _tllm_internal.thop + + def field(layout: torch.Tensor, name: str) -> int: + return int(layout[getattr(thop, f"MOE_A2A_{name}")]) + + ep_size, top_k, capacity, hidden_size = 8, 6, 128, 7168 + layout = NVLinkOneSided._make_workspace_layout( + ep_size, top_k, capacity, hidden_size, torch.bfloat16, None, 0, cft, fp8_combine + ) + dispatch_start = field(layout, "DISPATCH_PAYLOAD_OFFSET_INDEX") + dispatch_end = dispatch_start + field(layout, "DISPATCH_PAYLOAD_SIZE_INDEX") + combine_start = field(layout, "COMBINE_INPUT_OFFSET_INDEX") + combine_end = combine_start + field(layout, "COMBINE_INPUT_SIZE_INDEX") + total = field(layout, "WORKSPACE_SIZE_INDEX") + assert field(layout, "TOPK_TARGET_INDICES_OFFSET_INDEX") < dispatch_start + assert dispatch_end <= field(layout, "COMBINE_COMPLETION_FLAGS_OFFSET_INDEX") < combine_start + assert field(layout, "COMBINE_INPUT_SIZE_INDEX") == ep_size * capacity * hidden_size * 2 + assert dispatch_start % 256 == combine_start % 256 == total % 256 == 0 + assert total == NVLinkOneSided.calculate_required_workspace_size( + ep_size, + top_k, + capacity, + hidden_size, + torch.bfloat16, + can_use_cft_counted_writes=cft, + use_low_precision_combine=fp8_combine, + ) + if cft: + recv_start = field(layout, "COMBINE_RECV_OFFSET_INDEX") + recv_bytes = field(layout, "COMBINE_RECV_SIZE_INDEX") + assert combine_end <= recv_start and recv_start % 256 == 0 + assert recv_bytes == ep_size * capacity * hidden_size * (1 if fp8_combine else 2) + assert recv_start + recv_bytes == total + assert ( + dispatch_end + <= field(layout, "COMBINE_COUNTED_WRITE_COUNTERS_OFFSET_INDEX") + < combine_start + ) + else: + assert field(layout, "COMBINE_RECV_OFFSET_INDEX") == 0 + assert field(layout, "COMBINE_RECV_SIZE_INDEX") == 0 + assert field(layout, "COMBINE_COUNTED_WRITE_COUNTERS_OFFSET_INDEX") == 0 + assert combine_end == total + + # Routing metadata reserves the configured top-k, not the maximum supported top-k. + routes_start = field(layout, "TOPK_TARGET_RANKS_OFFSET_INDEX") + indices_start = field(layout, "TOPK_TARGET_INDICES_OFFSET_INDEX") + assert indices_start - routes_start == capacity * top_k * 4 + with pytest.raises(RuntimeError, match="capacity"): + torch.ops.trtllm.moe_a2a_get_workspace_layout(ep_size, capacity, top_k, -1, 256, 0) + with pytest.raises(RuntimeError, match="overflows"): + torch.ops.trtllm.moe_a2a_get_workspace_layout( + ep_size, capacity, top_k, (1 << 63) - 512, 1024, 0 + ) + + +# ============================================================================ +# CFT path selection and capability checks +# ============================================================================ + + +@pytest.mark.parametrize( + ("force_env", "driver_version", "num_tokens", "device_supported", "expected"), + [ + pytest.param(None, "615.00", 128, True, True, id="auto-at-threshold"), + pytest.param(None, "615.00", 129, True, False, id="auto-above-threshold"), + pytest.param("0", "615.00", 128, True, False, id="force-fence"), + pytest.param("1", "615.00", 129, True, True, id="force-cft-above-threshold"), + pytest.param("1", b"614.99", 128, True, False, id="force-cft-old-driver"), + pytest.param("1", None, 128, True, False, id="force-cft-nvml-error"), + pytest.param("1", "615.00", 128, False, False, id="force-cft-unsupported-device"), + pytest.param("invalid", "615.00", 129, True, False, id="invalid-env-uses-auto"), + ], +) +def test_cft_selection( + monkeypatch: pytest.MonkeyPatch, + force_env: str | None, + driver_version: str | bytes | None, + num_tokens: int, + device_supported: bool, + expected: bool, +) -> None: + from tensorrt_llm._torch.moe.fused_moe.communication import nvlink_one_sided + + if force_env is None: + monkeypatch.delenv(FORCE_CFT_ENV, raising=False) + else: + monkeypatch.setenv(FORCE_CFT_ENV, force_env) + + def query_driver_version() -> str | bytes: + if driver_version is None: + raise nvlink_one_sided.pynvml.NVMLError(nvlink_one_sided.pynvml.NVML_ERROR_UNKNOWN) + return driver_version + + monkeypatch.setattr(nvlink_one_sided.pynvml, "nvmlDeviceGetCount", lambda: 1) + monkeypatch.setattr(nvlink_one_sided.pynvml, "nvmlSystemGetDriverVersion", query_driver_version) + force_cft = nvlink_one_sided.get_force_cft() + version = nvlink_one_sided._get_nvidia_driver_version() + assert version == ( + driver_version.decode() if isinstance(driver_version, bytes) else driver_version + ) + can_use_cft = ( + nvlink_one_sided.resolve_cft_counted_writes(force_cft, version) and device_supported + ) + assert nvlink_one_sided.should_use_cft(can_use_cft, force_cft, 128, num_tokens) is expected + + +@pytest.mark.parametrize( + ("capability", "unsupported_index"), + [ + pytest.param((9, 0), None, id="hopper"), + pytest.param((10, 3), None, id="supported"), + pytest.param((10, 3), 0, id="no-fabric-handle"), + pytest.param((10, 3), 1, id="no-unicast-endpoint"), + pytest.param((10, 3), 2, id="no-counted-ops"), + ], +) +def test_cft_device_support( + monkeypatch: pytest.MonkeyPatch, + capability: tuple[int, int], + unsupported_index: int | None, +) -> None: + from tensorrt_llm._torch.moe.fused_moe.communication import nvlink_one_sided + + cuda = nvlink_one_sided.cuda + attributes = ( + cuda.CUdevice_attribute.CU_DEVICE_ATTRIBUTE_HANDLE_TYPE_FABRIC_SUPPORTED, + cuda.CUdevice_attribute.CU_DEVICE_ATTRIBUTE_LOGICAL_ENDPOINT_UNICAST_SUPPORTED, + cuda.CUdevice_attribute.CU_DEVICE_ATTRIBUTE_LOGICAL_ENDPOINT_COUNTED_OPS_SUPPORTED, + ) + unsupported = None if unsupported_index is None else attributes[unsupported_index] + monkeypatch.setattr(torch.cuda, "get_device_capability", lambda: capability) + monkeypatch.setattr(torch.cuda, "current_device", lambda: 0) + monkeypatch.setattr( + cuda, + "cuDeviceGetAttribute", + lambda attribute, device: (cuda.CUresult.CUDA_SUCCESS, int(attribute != unsupported)), + ) + reason = nvlink_one_sided._cft_device_support_reason() + if capability[0] < 10: + assert "SM90" in reason + elif unsupported is None: + assert reason is None + else: + assert unsupported.name in reason diff --git a/tests/unittest/_torch/moe/test_moe_a2a_cft.py b/tests/unittest/_torch/moe/test_moe_a2a_cft.py deleted file mode 100644 index 0d9bf38dbac3..000000000000 --- a/tests/unittest/_torch/moe/test_moe_a2a_cft.py +++ /dev/null @@ -1,78 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import pytest - -from tensorrt_llm._torch.moe.fused_moe.communication.moe_alltoall import ( - get_force_cft as get_force_cft_standalone, -) -from tensorrt_llm._torch.moe.fused_moe.communication.moe_alltoall import ( - should_use_cft as should_use_cft_standalone, -) -from tensorrt_llm._torch.moe.fused_moe.communication.nvlink_one_sided import ( - FORCE_CFT_ENV, - get_force_cft, - should_use_cft, -) - -# Environment-variable parsing only. The marker is also what makes the file -# reachable: the CPU stage collects only files that carry it. -pytestmark = pytest.mark.cpu_only - - -@pytest.mark.parametrize( - ("value", "expected"), - [ - (None, None), - ("", None), - ("2", None), - ("true", None), - (" 1 ", None), - ("0", False), - ("1", True), - ], -) -def test_get_force_cft(monkeypatch: pytest.MonkeyPatch, value: str | None, expected: bool | None): - if value is None: - monkeypatch.delenv(FORCE_CFT_ENV, raising=False) - else: - monkeypatch.setenv(FORCE_CFT_ENV, value) - - assert get_force_cft() is expected - assert get_force_cft_standalone() is expected - - -@pytest.mark.parametrize( - ("can_use_cft", "force_cft", "runtime_max_tokens_per_rank", "expected"), - [ - (True, None, 128, True), - (True, None, 129, False), - (True, False, 1, False), - (True, True, 129, True), - (False, True, 1, False), - (False, None, 1, False), - ], -) -def test_should_use_cft( - can_use_cft: bool, - force_cft: bool | None, - runtime_max_tokens_per_rank: int, - expected: bool, -): - assert should_use_cft(can_use_cft, force_cft, 128, runtime_max_tokens_per_rank) is expected - assert ( - should_use_cft_standalone(can_use_cft, force_cft, 128, runtime_max_tokens_per_rank) - is expected - ) diff --git a/tests/unittest/_torch/moe/test_moe_a2a_workspace.py b/tests/unittest/_torch/moe/test_moe_a2a_workspace.py deleted file mode 100644 index b2637f5103d9..000000000000 --- a/tests/unittest/_torch/moe/test_moe_a2a_workspace.py +++ /dev/null @@ -1,178 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: Apache-2.0 -"""Capacity checks for stable one-sided dispatch/combine workspace regions.""" - -from types import SimpleNamespace -from unittest.mock import MagicMock - -import pytest -import torch - -from tensorrt_llm._torch.moe.fused_moe.communication import nvlink_one_sided as a2a -from tensorrt_llm._torch.moe.fused_moe.communication.nvlink_one_sided import NVLinkOneSided - -pytestmark = pytest.mark.cpu_only - - -@pytest.fixture(params=[False, True], ids=["fence", "cft"]) -def comm(request): - instance = NVLinkOneSided.__new__(NVLinkOneSided) - instance._destroyed = True - instance._workspace_state = {} - instance._dispatch_state = {"phase": "idle"} - instance.ep_size = 4 - instance.max_num_tokens_per_rank = 32 - instance.can_use_cft_counted_writes = request.param - instance.workspace_size_per_rank = ( - 128 + (2 + int(request.param)) * 4 * 32 * 128 * 2 + 4 * 32 * 8 * 8 - ) - instance.PAYLOAD_DATA_OFFSET_INDEX = 0 - instance._workspace_lifecycle = SimpleNamespace(metainfo=torch.tensor([128])) - return instance - - -def test_fixed_offset_is_shared_and_independent_of_dispatch_layout(comm): - offset = comm._reserve_combine_region(128, torch.bfloat16) - for count in (1, 9, 32, 2): - for quantized in (False, True): - payloads = [ - torch.empty( - (0, 64 if quantized else 128), - dtype=torch.uint8 if quantized else torch.bfloat16, - ) - ] - if quantized: - payloads.append(torch.empty((0, 8), dtype=torch.uint8)) - payloads += [torch.empty((0, 8), dtype=torch.int32), torch.empty((0, 8))] - assert comm._check_dispatch_region(payloads, count) <= offset - assert comm._reserve_combine_region(128, torch.bfloat16) == offset - assert offset % 128 == 0 - regions = 2 if comm.can_use_cft_counted_writes else 1 - assert offset + regions * 4 * 32 * 128 * 2 <= comm.workspace_size_per_rank - - -@pytest.mark.parametrize("count", [0, -1, 33]) -def test_rejects_runtime_count_outside_capacity(comm, count): - with pytest.raises(ValueError, match="capacity"): - comm._check_dispatch_region([torch.empty((0, 128))], count) - - -def test_rejects_dispatch_overlap_before_native_write(comm): - comm._reserve_combine_region(128, torch.bfloat16) - with pytest.raises(ValueError, match="overlaps"): - comm._check_dispatch_region([torch.empty((0, 1024))], 32) - - -@pytest.mark.parametrize("hidden,dtype", [(64, torch.bfloat16), (128, torch.float32)]) -def test_rejects_shared_combine_layout_change(comm, hidden, dtype): - comm._reserve_combine_region(128, torch.bfloat16) - with pytest.raises(ValueError, match="stable combine"): - comm._reserve_combine_region(hidden, dtype) - - -def test_rejects_insufficient_combine_capacity(comm): - comm.workspace_size_per_rank = 256 - with pytest.raises(ValueError, match="too small"): - comm._reserve_combine_region(128, torch.bfloat16) - assert not comm._workspace_state - - -def test_lazy_reservation_checks_first_dispatch(comm): - comm._dispatch_state["dispatch_payload_end"] = comm.workspace_size_per_rank - with pytest.raises(ValueError, match="overlaps"): - comm._reserve_combine_region(128, torch.bfloat16) - assert not comm._workspace_state - - -def test_lazy_reservation_guards_following_dispatch(comm): - payload = torch.empty((0, 64), dtype=torch.uint8) - comm._dispatch_state["dispatch_payload_end"] = comm._check_dispatch_region([payload], 9) - offset = comm._reserve_combine_region(128, torch.bfloat16) - assert comm._check_dispatch_region([payload], 32) <= offset - with pytest.raises(ValueError, match="overlaps"): - comm._check_dispatch_region([torch.empty((0, 1024))], 32) - - -@pytest.mark.parametrize("offset_delta", [-128, 128]) -def test_dispatch_rejects_native_offset(comm, monkeypatch, offset_delta): - """A native/Python layout disagreement must not publish a dispatched phase.""" - comm._reserve_combine_region(128, torch.bfloat16) - comm.mnnvl_mem = SimpleNamespace(mapped=True) - comm.workspace = torch.empty(0, dtype=torch.uint8) - comm.ep_rank, comm.top_k, comm.num_experts = 0, 8, 256 - comm._rank_mask_enabled = False - comm._force_cft = None - comm.cft_max_batch_for_dispatch = 128 - comm.invalid_token_expert_id = -1 - comm._workspace_lifecycle.coordinator = MagicMock() - comm._workspace_lifecycle.watchdog_for = MagicMock(return_value=None) - monkeypatch.setattr(a2a, "reject_rank_mask_cuda_graph_capture", MagicMock()) - hidden = torch.empty((1, 128), dtype=torch.bfloat16) - slots = torch.zeros((1, 8), dtype=torch.int32) - scales = torch.ones((1, 8)) - payloads = [hidden, slots, scales] - expected_end = comm._check_dispatch_region(payloads, 1) - native_dispatch = MagicMock( - return_value=(payloads, expected_end + offset_delta, torch.empty(0)) - ) - monkeypatch.setattr(torch.ops.trtllm, "moe_a2a_dispatch", native_dispatch, raising=False) - - with pytest.raises(RuntimeError, match="native A2A dispatch layout disagrees"): - comm.dispatch(hidden, None, slots, scales, [1] * comm.ep_size) - - native_dispatch.assert_called_once() - assert comm._dispatch_state["phase"] == "idle" - assert "dispatch_payload_end" not in comm._dispatch_state - assert "combine_payload_offset" not in comm._dispatch_state - - -@pytest.mark.parametrize("use_cft", [False, True]) -def test_failed_reservation_does_not_publish_workspace(monkeypatch, use_cft): - def init_base(self, mapping): - self.mapping = mapping - self.ep_size = mapping.world_size - self.ep_rank = mapping.rank - - monkeypatch.setattr(a2a.Communication, "__init__", init_base) - monkeypatch.setattr(NVLinkOneSided, "_init_constants", classmethod(lambda cls: None)) - monkeypatch.setattr(NVLinkOneSided, "get_aux_data_size", staticmethod(lambda *args: 128)) - for name in ( - "PAYLOAD_DATA_OFFSET_INDEX", - "FLAG_VAL_OFFSET_INDEX", - "DISPATCH_COMPLETION_FLAGS_OFFSET_INDEX", - "COMBINE_COMPLETION_FLAGS_OFFSET_INDEX", - ): - monkeypatch.setattr(NVLinkOneSided, name, 0) - monkeypatch.setattr(NVLinkOneSided, "_WORKSPACES", {}) - monkeypatch.setattr(NVLinkOneSided, "_WORKSPACE_REFCOUNTS", {}) - monkeypatch.setattr(NVLinkOneSided, "_WORKSPACE", None) - for name in ("MnnvlMemory", "CftMnnvlMemory"): - monkeypatch.setattr(a2a, name, MagicMock()) - monkeypatch.setattr( - torch.ops.trtllm, - "moe_a2a_initialize", - MagicMock(return_value=torch.tensor([128])), - raising=False, - ) - lifecycle = SimpleNamespace(metainfo=torch.tensor([128]), register=MagicMock()) - monkeypatch.setattr( - a2a, - "_MnnvlAlltoAllWorkspaceLifecycle", - SimpleNamespace(get_or_create=MagicMock(return_value=lifecycle)), - ) - monkeypatch.setenv("TRTLLM_MOE_A2A_WORKSPACE_MB", "1") - monkeypatch.delenv("TRTLLM_MOE_A2A_FORCE_CFT", raising=False) - with pytest.raises(ValueError, match="too small"): - NVLinkOneSided( - SimpleNamespace(world_size=4, rank=0, has_cp_helix=lambda: False), - 256, - 8, - 1024, - hidden_size=1024, - dtype=torch.bfloat16, - can_use_cft_counted_writes=use_cft, - ) - lifecycle.register.assert_not_called() - assert not NVLinkOneSided._WORKSPACES - assert not NVLinkOneSided._WORKSPACE_REFCOUNTS - assert NVLinkOneSided._WORKSPACE is None diff --git a/tests/unittest/_torch/moe/test_moe_comm.py b/tests/unittest/_torch/moe/test_moe_comm.py index 1bc5651d7f77..1da7c69afa46 100644 --- a/tests/unittest/_torch/moe/test_moe_comm.py +++ b/tests/unittest/_torch/moe/test_moe_comm.py @@ -64,8 +64,8 @@ from mpi4py import MPI import tensorrt_llm as tllm -import tensorrt_llm._mnnvl_utils as mnnvl -from tensorrt_llm._mnnvl_utils import MnnvlMemory, MnnvlMoe +import tensorrt_llm._torch.distributed.mnnvl_memory as mnnvl +from tensorrt_llm._torch.distributed.mnnvl_memory import MnnvlMemory from tensorrt_llm._torch.moe.fused_moe.communication.allgather_reducescatter import ( AllGatherReduceScatter, ) @@ -73,7 +73,10 @@ from tensorrt_llm._torch.moe.fused_moe.communication.deep_ep_low_latency import DeepEPLowLatency from tensorrt_llm._torch.moe.fused_moe.communication.nccl_ep import NcclEP from tensorrt_llm._torch.moe.fused_moe.communication.nvlink_one_sided import NVLinkOneSided -from tensorrt_llm._torch.moe.fused_moe.communication.nvlink_two_sided import NVLinkTwoSided +from tensorrt_llm._torch.moe.fused_moe.communication.nvlink_two_sided import ( + MnnvlMoe, + NVLinkTwoSided, +) from tensorrt_llm._torch.moe.fused_moe.communication.nvlink_two_sided_flashinfer import ( NVLinkTwoSidedFlashinfer, ) @@ -257,15 +260,15 @@ def _read_nvlink_topk_target_ranks( return raw.view(torch.int32).view(max_num_tokens, top_k).cpu() -def _read_nvlink_topk_send_indices( +def _read_nvlink_topk_target_indices( comm: NVLinkOneSided, max_num_tokens: int, top_k: int, ) -> torch.Tensor: - """Read topk_send_indices[max_num_tokens, top_k] from NVLinkOneSided workspace.""" + """Read topk_target_indices[max_num_tokens, top_k] from NVLinkOneSided workspace.""" from tensorrt_llm.bindings import internal as _tllm_internal - offset_index = int(_tllm_internal.thop.MOE_A2A_TOPK_SEND_INDICES_OFFSET_INDEX) + offset_index = int(_tllm_internal.thop.MOE_A2A_TOPK_TARGET_INDICES_OFFSET_INDEX) offset = comm.moe_a2a_metainfo[offset_index].item() raw = comm.workspace[ comm.ep_rank, @@ -303,12 +306,12 @@ def _run_nvlink_rank_mask_dispatch( runtime_max_tokens_per_rank, comm.top_k, ) - topk_send_indices = _read_nvlink_topk_send_indices( + topk_target_indices = _read_nvlink_topk_target_indices( comm, runtime_max_tokens_per_rank, comm.top_k, ) - return recv_tensors, int(combine_payload_offset), topk_target_ranks, topk_send_indices + return recv_tensors, int(combine_payload_offset), topk_target_ranks, topk_target_indices def _run_nvlink_rank_mask_combine( @@ -347,7 +350,7 @@ def _run_nvlink_rank_mask_dispatch_combine( active_rank_mask: Optional[torch.Tensor], ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """Run raw NVLink one-sided dispatch/combine with an optional active rank mask.""" - recv_tensors, combine_payload_offset, topk_target_ranks, topk_send_indices = ( + recv_tensors, combine_payload_offset, topk_target_ranks, topk_target_indices = ( _run_nvlink_rank_mask_dispatch( comm, token_selected_experts, @@ -366,14 +369,14 @@ def _run_nvlink_rank_mask_dispatch_combine( enable_rank_mask, active_rank_mask, ) - return combined.cpu(), topk_target_ranks, topk_send_indices + return combined.cpu(), topk_target_ranks, topk_target_indices def _expected_nvlink_rank_mask_combine_output( comm: NVLinkOneSided, payload: torch.Tensor, topk_target_ranks: torch.Tensor, - topk_send_indices: torch.Tensor, + topk_target_indices: torch.Tensor, local_num_tokens: int, runtime_max_tokens_per_rank: int, ) -> torch.Tensor: @@ -386,7 +389,7 @@ def _expected_nvlink_rank_mask_combine_output( dtype=torch.float32, device=payload.device, ) - payload_offset_index = int(_tllm_internal.thop.MOE_A2A_PAYLOAD_DATA_OFFSET_INDEX) + payload_offset_index = int(_tllm_internal.thop.MOE_A2A_DISPATCH_PAYLOAD_OFFSET_INDEX) payload_offset = comm.moe_a2a_metainfo[payload_offset_index].item() bytes_per_rank = ( comm.ep_size * runtime_max_tokens_per_rank * hidden_size * payload.element_size() @@ -395,7 +398,7 @@ def _expected_nvlink_rank_mask_combine_output( for token_idx in range(local_num_tokens): for k in range(comm.top_k): target_rank = int(topk_target_ranks[token_idx, k].item()) - dst_idx = int(topk_send_indices[token_idx, k].item()) + dst_idx = int(topk_target_indices[token_idx, k].item()) if dst_idx < 0: continue raw = comm.workspace[target_rank, payload_offset : payload_offset + bytes_per_rank] @@ -580,7 +583,7 @@ def create_comm_object( # Reset class-level singleton to avoid assertion failures when # test params change across MPI process reuse. NVLinkOneSided._WORKSPACE = None - os.environ["TRTLLM_MOE_A2A_WORKSPACE_MB"] = NVLINK_WORKSPACE_MB + os.environ["TRTLLM_NVLINK_ONE_SIDED_A2A_WORKSPACE_MB"] = NVLINK_WORKSPACE_MB return NVLinkOneSided( mapping=mapping, @@ -1515,7 +1518,7 @@ def _worker_rank_mask_one_rank_masked( comm, payload, topk_target_ranks, - topk_send_indices, + topk_target_indices, local_num_tokens, local_num_tokens, ) @@ -1950,7 +1953,7 @@ def _build_combine_reference( # scaling: per-row global fp32 scale + per-group-of-16 fp8 scale, # with E2M1 quantization. After NVLink transfer, # dequantize_nvfp4_sharedmem reverses the process. The top_k - # reduction is then done in bf16 by torch.sum in _mnnvl_utils.py. The + # reduction is then done in bf16 by torch.sum in nvlink_two_sided.py. The # NVFP4 round-trip is precomputed on the worker GPU. for proc_result in all_results: nvfp4_out = proc_result["moe_output_for_ref"] diff --git a/tests/unittest/_torch/moe/test_moe_module.py b/tests/unittest/_torch/moe/test_moe_module.py index 5d2f504f4d8c..d211fc41ef47 100644 --- a/tests/unittest/_torch/moe/test_moe_module.py +++ b/tests/unittest/_torch/moe/test_moe_module.py @@ -69,8 +69,8 @@ from transformers.configuration_utils import PretrainedConfig import tensorrt_llm.bindings.internal.runtime as _tbr -from tensorrt_llm._mnnvl_utils import MnnvlMemory from tensorrt_llm._torch.autotuner import AutoTuner, autotune +from tensorrt_llm._torch.distributed.mnnvl_memory import MnnvlMemory from tensorrt_llm._torch.model_config import ModelConfig from tensorrt_llm._torch.moe.fused_moe import ( DEFAULT_MOE_ACTIVATION, @@ -801,7 +801,7 @@ def run_forward(): # thread is spawned lazily on first submit and persists by design, so the # multi-GPU tests disable pytest-threadleak via @pytest.mark.threadleak( # enabled=False) (same convention as the conftest mpi_pool_executor users -# test_moe_a2a / test_autotuner), rather than excluding it in pytest.ini. +# test_nvlink_one_sided / test_autotuner), rather than excluding it in pytest.ini. # --------------------------------------------------------------------------- @@ -863,7 +863,7 @@ def moe_multi_gpu_executor(): pool's manager thread is spawned lazily on first submit and persists by design, so the multi-GPU tests disable pytest-threadleak via @pytest.mark.threadleak(enabled=False) (same convention as the other - mpi_pool_executor users, test_moe_a2a / test_autotuner). world_size is 4. + mpi_pool_executor users, test_nvlink_one_sided / test_autotuner). world_size is 4. """ world_size = 4 with MPIPoolExecutor( diff --git a/tests/unittest/_torch/multi_gpu/test_mnnvl_allreduce.py b/tests/unittest/_torch/multi_gpu/test_mnnvl_allreduce.py index f0abca2492bf..be89e165682f 100644 --- a/tests/unittest/_torch/multi_gpu/test_mnnvl_allreduce.py +++ b/tests/unittest/_torch/multi_gpu/test_mnnvl_allreduce.py @@ -26,9 +26,9 @@ from utils.util import skip_pre_blackwell import tensorrt_llm -from tensorrt_llm._mnnvl_utils import MnnvlMemory from tensorrt_llm._torch.distributed import (AllReduce, AllReduceFusionOp, AllReduceParams) +from tensorrt_llm._torch.distributed.mnnvl_memory import MnnvlMemory from tensorrt_llm._torch.distributed.ops import MNNVLAllReduce from tensorrt_llm.functional import AllReduceStrategy from tensorrt_llm.mapping import Mapping diff --git a/tests/unittest/_torch/multi_gpu/test_mnnvl_memory.py b/tests/unittest/_torch/multi_gpu/test_mnnvl_memory.py index 3b0995a9052e..83a7489231fe 100644 --- a/tests/unittest/_torch/multi_gpu/test_mnnvl_memory.py +++ b/tests/unittest/_torch/multi_gpu/test_mnnvl_memory.py @@ -1,4 +1,4 @@ -# SPDX-FileCopyrightText: Copyright (c) 2022-2024 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-FileCopyrightText: Copyright (c) 2022-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 # # Licensed under the Apache License, Version 2.0 (the "License"); @@ -20,6 +20,7 @@ import tensorrt_llm as tllm from tensorrt_llm import Mapping +from tensorrt_llm._torch.distributed.mnnvl_memory import MnnvlMemory class TestMnnvlMemory(unittest.TestCase): @@ -42,7 +43,7 @@ def setUp(self): local_dev_count = torch.cuda.device_count() assert self.local_world_size <= local_dev_count, "ntasks_per_node should be less than local device count" torch.cuda.set_device(self.local_rank) - tllm.MnnvlMemory.initialize() + MnnvlMemory.initialize() # MnnvlMemory splits its communicator per MoE expert-parallel group, so # an allocation is shared across ranks only when moe_ep_size spans them. self.mapping = Mapping(self.world_size, @@ -56,15 +57,15 @@ def align_memory(size: int): align_size = 2 * 1024 * 1024 return (size + align_size - 1) // align_size * align_size - @pytest.mark.skipif(not tllm.MnnvlMemory.supports_mnnvl(), + @pytest.mark.skipif(not MnnvlMemory.supports_mnnvl(), reason="Mnnvl memory is not supported on this platform" ) # Skip tests on unsupported platform def test_mnnvl_memory(self): # allocate un-aligned memory allocate0_size = 4 * 1024 * 1024 - 3 * 1024 - mnnvl_memory0 = tllm.MnnvlMemory(self.mapping, allocate0_size) + mnnvl_memory0 = MnnvlMemory(self.mapping, allocate0_size) allocate0_size_aligned = TestMnnvlMemory.align_memory(allocate0_size) - assert tllm.MnnvlMemory.current_mem_offset == allocate0_size_aligned + assert MnnvlMemory.current_mem_offset == allocate0_size_aligned tensor0 = mnnvl_memory0.as_torch_strided_tensor(torch.int32) numel_per_rank = allocate0_size // 4 @@ -79,9 +80,9 @@ def test_mnnvl_memory(self): device='cuda')), f"segment written by rank {r} mismatched" allocate1_size = 30 * 1024 * 1024 - 2 * 1024 - mnnvl_memory1 = tllm.MnnvlMemory(self.mapping, allocate1_size) + mnnvl_memory1 = MnnvlMemory(self.mapping, allocate1_size) allocate1_size_aligned = TestMnnvlMemory.align_memory(allocate1_size) - assert tllm.MnnvlMemory.current_mem_offset == allocate0_size_aligned + allocate1_size_aligned + assert MnnvlMemory.current_mem_offset == allocate0_size_aligned + allocate1_size_aligned tensor1 = mnnvl_memory1.as_torch_strided_tensor(torch.float32) numel_per_rank = allocate1_size // 4 tensor1[(self.rank + 5) % self.world_size] = torch.arange( @@ -103,11 +104,10 @@ def test_mnnvl_memory(self): tllm.mpi_barrier() large_allocation2_size = 768 * 1024 * 1024 - large_mnnvl_memory2 = tllm.MnnvlMemory(self.mapping, - large_allocation2_size) + large_mnnvl_memory2 = MnnvlMemory(self.mapping, large_allocation2_size) allocate2_size_aligned = TestMnnvlMemory.align_memory( large_allocation2_size) - assert tllm.MnnvlMemory.current_mem_offset == allocate2_size_aligned + assert MnnvlMemory.current_mem_offset == allocate2_size_aligned assert large_mnnvl_memory2.rank_stride == (1 << 30) del tensor1 diff --git a/tests/unittest/_torch/ray_orchestrator/multi_gpu/test_mnnvl_allreduce.py b/tests/unittest/_torch/ray_orchestrator/multi_gpu/test_mnnvl_allreduce.py index 6be9c9d20d70..ac9da13a1dda 100644 --- a/tests/unittest/_torch/ray_orchestrator/multi_gpu/test_mnnvl_allreduce.py +++ b/tests/unittest/_torch/ray_orchestrator/multi_gpu/test_mnnvl_allreduce.py @@ -102,7 +102,7 @@ def mnnvl_supported(self) -> bool: Only the hardware capability, not is_mnnvl(): off aarch64 the test bypasses that policy gate below, so skipping has to key off what the machine can actually do. """ - from tensorrt_llm._mnnvl_utils import MnnvlMemory + from tensorrt_llm._torch.distributed.mnnvl_memory import MnnvlMemory MnnvlMemory.initialize() return MnnvlMemory.supports_mnnvl() diff --git a/tests/unittest/_torch/test_mnnvl_alltoall_workspace.py b/tests/unittest/_torch/test_mnnvl_alltoall_workspace.py index 2caaf152c267..7f5c4df8b0b4 100644 --- a/tests/unittest/_torch/test_mnnvl_alltoall_workspace.py +++ b/tests/unittest/_torch/test_mnnvl_alltoall_workspace.py @@ -21,10 +21,9 @@ import pytest import torch -import tensorrt_llm._mnnvl_utils as mnnvl +import tensorrt_llm._torch.distributed.mnnvl_memory as mnnvl import tensorrt_llm._torch.moe.fused_moe.communication.nvlink_one_sided as one_sided_module from tensorrt_llm._torch.mnnvl_alltoall_workspace import _MnnvlAlltoAllWorkspaceLifecycle -from tensorrt_llm._torch.moe.fused_moe.communication.moe_alltoall import MoeAlltoAll from tensorrt_llm._torch.moe.fused_moe.communication.nvlink_one_sided import NVLinkOneSided from tensorrt_llm._torch.moe.fused_moe.communication.nvlink_two_sided import NVLinkTwoSided @@ -659,11 +658,8 @@ def test_two_sided_checkpoint_restore_noop_preserves_shared_owner_state( assert second._dispatch_state -@pytest.mark.parametrize("wrapper_type", [MoeAlltoAll, NVLinkOneSided]) -def test_frontend_checkpoint_delegates_to_shared_lifecycle( - wrapper_type: type[MoeAlltoAll] | type[NVLinkOneSided], -) -> None: - wrapper = wrapper_type.__new__(wrapper_type) +def test_frontend_checkpoint_delegates_to_shared_lifecycle() -> None: + wrapper = NVLinkOneSided.__new__(NVLinkOneSided) wrapper.can_use_cft_counted_writes = False wrapper._workspace_lifecycle = Mock() comm = Mock() @@ -676,17 +672,13 @@ def test_frontend_checkpoint_delegates_to_shared_lifecycle( assert wrapper._workspace_lifecycle.checkpoint_restore.call_args.args[0] is comm -@pytest.mark.parametrize("wrapper_type", [MoeAlltoAll, NVLinkOneSided]) -def test_frontend_destroy_unregisters_from_shared_lifecycle( - wrapper_type: type[MoeAlltoAll] | type[NVLinkOneSided], -) -> None: - wrapper = wrapper_type.__new__(wrapper_type) +def test_frontend_destroy_unregisters_from_shared_lifecycle() -> None: + wrapper = NVLinkOneSided.__new__(NVLinkOneSided) wrapper._destroyed = False wrapper._workspace_registered = True lifecycle = Mock() wrapper._workspace_lifecycle = lifecycle - if wrapper_type is NVLinkOneSided: - wrapper._workspace_key = None + wrapper._workspace_key = None wrapper.destroy() wrapper.destroy() @@ -694,53 +686,6 @@ def test_frontend_destroy_unregisters_from_shared_lifecycle( lifecycle.unregister.assert_called_once_with(wrapper) -def test_moe_alltoall_aborted_registration_does_not_unregister( - monkeypatch: pytest.MonkeyPatch, -) -> None: - lifecycle = Mock() - lifecycle.register.side_effect = RuntimeError("registration failed") - monkeypatch.setattr(MoeAlltoAll, "_WORKSPACES", {}) - monkeypatch.setattr(MoeAlltoAll, "_init_constants", Mock()) - monkeypatch.setattr( - MoeAlltoAll, - "_METAINFO_INDEX", - { - "FLAG_VAL_OFFSET_INDEX": 0, - "DISPATCH_COMPLETION_FLAGS_OFFSET_INDEX": 0, - "COMBINE_COMPLETION_FLAGS_OFFSET_INDEX": 0, - }, - ) - monkeypatch.setattr(mnnvl.MnnvlMemory, "initialize", Mock()) - memory = Mock() - memory.as_torch_strided_tensor.return_value = torch.zeros(1, dtype=torch.uint8) - monkeypatch.setattr( - "tensorrt_llm._torch.moe.fused_moe.communication.moe_alltoall.MnnvlMemory", - Mock(return_value=memory), - ) - monkeypatch.setattr( - _MnnvlAlltoAllWorkspaceLifecycle, - "get_or_create", - Mock(return_value=lifecycle), - ) - monkeypatch.setattr( - torch.ops.trtllm, - "moe_a2a_initialize", - Mock(return_value=torch.tensor([1])), - ) - mapping = SimpleNamespace(moe_ep_size=2, moe_ep_rank=0) - - with pytest.raises(RuntimeError, match="registration failed"): - MoeAlltoAll( - mapping=mapping, - max_num_tokens=1, - top_k=1, - num_slots=2, - workspace_size_per_rank=1, - ) - - lifecycle.unregister.assert_not_called() - - def test_one_sided_checkpoint_rejects_destroyed_workspace( monkeypatch: pytest.MonkeyPatch, ) -> None: diff --git a/tests/unittest/_torch/test_mnnvl_memory_lifecycle.py b/tests/unittest/_torch/test_mnnvl_memory_lifecycle.py index b42605ebf8d4..42e9cf6eea8f 100644 --- a/tests/unittest/_torch/test_mnnvl_memory_lifecycle.py +++ b/tests/unittest/_torch/test_mnnvl_memory_lifecycle.py @@ -19,8 +19,8 @@ import pytest import torch -import tensorrt_llm._mnnvl_utils as mnnvl -from tensorrt_llm._torch.moe.fused_moe.communication.moe_alltoall import MoeAlltoAll +import tensorrt_llm._torch.distributed.mnnvl_memory as mnnvl +from tensorrt_llm._torch.moe.fused_moe.communication.nvlink_one_sided import NVLinkOneSided from tensorrt_llm._torch.moe.fused_moe.communication.nvlink_two_sided import NVLinkTwoSided from tensorrt_llm.mapping import Mapping @@ -721,7 +721,7 @@ def test_create_and_map_handles_close_failure_does_not_mask_original_error(monke def _make_moe_alltoall_for_lifecycle(): - obj = MoeAlltoAll.__new__(MoeAlltoAll) + obj = NVLinkOneSided.__new__(NVLinkOneSided) obj._destroyed = True obj.mnnvl_mem = Mock(mapped=True) return obj diff --git a/tests/unittest/_torch/test_mnnvl_utils.py b/tests/unittest/_torch/test_mnnvl_utils.py index 6984fac51a16..8e039dde0905 100644 --- a/tests/unittest/_torch/test_mnnvl_utils.py +++ b/tests/unittest/_torch/test_mnnvl_utils.py @@ -17,7 +17,7 @@ import pynvml -from tensorrt_llm._mnnvl_utils import MnnvlMemory +from tensorrt_llm._torch.distributed.mnnvl_memory import MnnvlMemory from tensorrt_llm._torch.moe.fused_moe.communication.deep_ep_low_latency import DeepEPLowLatency @@ -31,7 +31,10 @@ def teardown_function() -> None: MnnvlMemory.support_nvlink.cache_clear() -@patch("tensorrt_llm._mnnvl_utils.torch.cuda.get_device_name", return_value="NVIDIA H200 NVL") +@patch( + "tensorrt_llm._torch.distributed.mnnvl_memory.torch.cuda.get_device_name", + return_value="NVIDIA H200 NVL", +) def test_pcie_nvl_sku_detected_by_name(mock_get_device_name) -> None: with patch.object(MnnvlMemory, "_ensure_nvml_initialized") as mock_initialize: assert MnnvlMemory._is_pcie_nvl_sku(0) @@ -39,15 +42,19 @@ def test_pcie_nvl_sku_detected_by_name(mock_get_device_name) -> None: mock_initialize.assert_not_called() -@patch("tensorrt_llm._mnnvl_utils.torch.cuda.get_device_name", return_value="NVIDIA H200") +@patch( + "tensorrt_llm._torch.distributed.mnnvl_memory.torch.cuda.get_device_name", + return_value="NVIDIA H200", +) @patch.object(MnnvlMemory, "_ensure_nvml_initialized") -@patch("tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetCount", return_value=8) +@patch("tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetCount", return_value=8) @patch( - "tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetHandleByIndex", side_effect=lambda index: index + "tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetHandleByIndex", + side_effect=lambda index: index, ) @patch.object(MnnvlMemory, "support_nvlink", return_value=True) @patch( - "tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetP2PStatus", + "tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetP2PStatus", return_value=pynvml.NVML_P2P_STATUS_NOT_SUPPORTED, ) def test_split_nvlink_topology_detected( @@ -64,7 +71,7 @@ def common_ancestor(_self_handle: int, peer_handle: int) -> int: return pynvml.NVML_TOPOLOGY_NODE with patch( - "tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetTopologyCommonAncestor", + "tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetTopologyCommonAncestor", side_effect=common_ancestor, ): assert MnnvlMemory._is_pcie_nvl_sku(0) @@ -73,19 +80,23 @@ def common_ancestor(_self_handle: int, peer_handle: int) -> int: mock_support_nvlink.assert_called_once_with(0, need_all_up=False) -@patch("tensorrt_llm._mnnvl_utils.torch.cuda.get_device_name", return_value="NVIDIA H200") +@patch( + "tensorrt_llm._torch.distributed.mnnvl_memory.torch.cuda.get_device_name", + return_value="NVIDIA H200", +) @patch.object(MnnvlMemory, "_ensure_nvml_initialized") -@patch("tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetCount", return_value=2) +@patch("tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetCount", return_value=2) @patch( - "tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetHandleByIndex", side_effect=lambda index: index + "tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetHandleByIndex", + side_effect=lambda index: index, ) @patch.object(MnnvlMemory, "support_nvlink", return_value=False) @patch( - "tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetP2PStatus", + "tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetP2PStatus", return_value=pynvml.NVML_P2P_STATUS_NOT_SUPPORTED, ) @patch( - "tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetTopologyCommonAncestor", + "tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetTopologyCommonAncestor", return_value=pynvml.NVML_TOPOLOGY_SYSTEM, ) def test_pcie_hopper_with_system_peers_is_not_split_nvlink( @@ -103,20 +114,23 @@ def test_pcie_hopper_with_system_peers_is_not_split_nvlink( mock_common_ancestor.assert_called_once_with(0, 1) -@patch("tensorrt_llm._mnnvl_utils.torch.cuda.get_device_name", return_value="NVIDIA H200") +@patch( + "tensorrt_llm._torch.distributed.mnnvl_memory.torch.cuda.get_device_name", + return_value="NVIDIA H200", +) @patch.object(MnnvlMemory, "_ensure_nvml_initialized") -@patch("tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetCount", return_value=2) +@patch("tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetCount", return_value=2) @patch( - "tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetHandleByIndex", + "tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetHandleByIndex", side_effect=lambda index: index, ) @patch.object(MnnvlMemory, "support_nvlink") @patch( - "tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetP2PStatus", + "tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetP2PStatus", return_value=pynvml.NVML_P2P_STATUS_OK, ) @patch( - "tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetTopologyCommonAncestor", + "tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetTopologyCommonAncestor", return_value=pynvml.NVML_TOPOLOGY_SYSTEM, ) def test_dual_socket_hgx_with_system_peers_is_not_split_nvlink( @@ -134,14 +148,18 @@ def test_dual_socket_hgx_with_system_peers_is_not_split_nvlink( mock_common_ancestor.assert_called_once_with(0, 1) -@patch("tensorrt_llm._mnnvl_utils.torch.cuda.get_device_name", return_value="NVIDIA H200") +@patch( + "tensorrt_llm._torch.distributed.mnnvl_memory.torch.cuda.get_device_name", + return_value="NVIDIA H200", +) @patch.object(MnnvlMemory, "_ensure_nvml_initialized") -@patch("tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetCount", return_value=8) +@patch("tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetCount", return_value=8) @patch( - "tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetHandleByIndex", side_effect=lambda index: index + "tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetHandleByIndex", + side_effect=lambda index: index, ) @patch( - "tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetTopologyCommonAncestor", + "tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetTopologyCommonAncestor", return_value=pynvml.NVML_TOPOLOGY_NODE, ) @patch.object(MnnvlMemory, "support_nvlink", return_value=True) @@ -157,7 +175,10 @@ def test_nvswitch_topology_remains_supported( mock_support_nvlink.assert_not_called() -@patch("tensorrt_llm._mnnvl_utils.torch.cuda.get_device_name", return_value="NVIDIA B200 NVL") +@patch( + "tensorrt_llm._torch.distributed.mnnvl_memory.torch.cuda.get_device_name", + return_value="NVIDIA B200 NVL", +) @patch.object(MnnvlMemory, "_ensure_nvml_initialized") def test_b200_does_not_use_hopper_topology_fallback(mock_initialize, mock_get_device_name) -> None: assert not MnnvlMemory._is_pcie_nvl_sku(0) @@ -167,23 +188,26 @@ def test_b200_does_not_use_hopper_topology_fallback(mock_initialize, mock_get_de def test_topology_probe_initializes_nvml() -> None: with ( patch( - "tensorrt_llm._mnnvl_utils.torch.cuda.get_device_name", + "tensorrt_llm._torch.distributed.mnnvl_memory.torch.cuda.get_device_name", return_value="NVIDIA H200", ), patch( - "tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetCount", + "tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetCount", side_effect=[pynvml.NVMLError_Uninitialized(), 1], ), - patch("tensorrt_llm._mnnvl_utils.pynvml.nvmlInit") as mock_nvml_init, - patch("tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetHandleByIndex", return_value=0), + patch("tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlInit") as mock_nvml_init, + patch( + "tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetHandleByIndex", + return_value=0, + ), ): assert not MnnvlMemory._is_pcie_nvl_sku(0) mock_nvml_init.assert_called_once_with() -@patch("tensorrt_llm._mnnvl_utils.get_sm_version", return_value=90) -@patch("tensorrt_llm._mnnvl_utils.torch.cuda.current_device", return_value=0) +@patch("tensorrt_llm._torch.distributed.mnnvl_memory.get_sm_version", return_value=90) +@patch("tensorrt_llm._torch.distributed.mnnvl_memory.torch.cuda.current_device", return_value=0) @patch.object(MnnvlMemory, "_is_pcie_nvl_sku", return_value=True) @patch.object(MnnvlMemory, "support_nvlink") def test_supports_mnnvl_rejects_split_topology( @@ -193,8 +217,8 @@ def test_supports_mnnvl_rejects_split_topology( mock_support_nvlink.assert_not_called() -@patch("tensorrt_llm._mnnvl_utils.get_sm_version", return_value=90) -@patch("tensorrt_llm._mnnvl_utils.torch.cuda.current_device", return_value=0) +@patch("tensorrt_llm._torch.distributed.mnnvl_memory.get_sm_version", return_value=90) +@patch("tensorrt_llm._torch.distributed.mnnvl_memory.torch.cuda.current_device", return_value=0) @patch.object(MnnvlMemory, "_is_pcie_nvl_sku", return_value=False) @patch.object(MnnvlMemory, "support_nvlink", return_value=True) def test_supports_mnnvl_accepts_full_fabric( @@ -206,10 +230,14 @@ def test_supports_mnnvl_accepts_full_fabric( @patch.object(MnnvlMemory, "_ensure_nvml_initialized") @patch( - "tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetHandleByIndex", side_effect=lambda index: index + "tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetHandleByIndex", + side_effect=lambda index: index, +) +@patch("tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.NVML_NVLINK_MAX_LINKS", 36) +@patch( + "tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetNvLinkCapability", + return_value=True, ) -@patch("tensorrt_llm._mnnvl_utils.pynvml.NVML_NVLINK_MAX_LINKS", 36) -@patch("tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetNvLinkCapability", return_value=True) def test_support_nvlink_ignores_indices_past_the_gpu_link_count( mock_capability, mock_get_handle, mock_initialize ) -> None: @@ -226,16 +254,23 @@ def link_state(handle, link_idx): raise pynvml.NVMLError_NotSupported() return True - with patch("tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetNvLinkState", side_effect=link_state): + with patch( + "tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetNvLinkState", + side_effect=link_state, + ): assert MnnvlMemory.support_nvlink(0, True) @patch.object(MnnvlMemory, "_ensure_nvml_initialized") @patch( - "tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetHandleByIndex", side_effect=lambda index: index + "tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetHandleByIndex", + side_effect=lambda index: index, +) +@patch("tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.NVML_NVLINK_MAX_LINKS", 36) +@patch( + "tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetNvLinkCapability", + return_value=True, ) -@patch("tensorrt_llm._mnnvl_utils.pynvml.NVML_NVLINK_MAX_LINKS", 36) -@patch("tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetNvLinkCapability", return_value=True) def test_support_nvlink_rejects_a_down_link_inside_the_gpu_range( mock_capability, mock_get_handle, mock_initialize ) -> None: @@ -247,16 +282,23 @@ def link_state(handle, link_idx): raise pynvml.NVMLError_NotSupported() return link_idx != 3 - with patch("tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetNvLinkState", side_effect=link_state): + with patch( + "tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetNvLinkState", + side_effect=link_state, + ): assert not MnnvlMemory.support_nvlink(0, True) @patch.object(MnnvlMemory, "_ensure_nvml_initialized") @patch( - "tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetHandleByIndex", side_effect=lambda index: index + "tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetHandleByIndex", + side_effect=lambda index: index, +) +@patch("tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.NVML_NVLINK_MAX_LINKS", 36) +@patch( + "tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetNvLinkCapability", + return_value=True, ) -@patch("tensorrt_llm._mnnvl_utils.pynvml.NVML_NVLINK_MAX_LINKS", 36) -@patch("tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetNvLinkCapability", return_value=True) def test_support_nvlink_keeps_probing_after_a_rejected_index( mock_capability, mock_get_handle, mock_initialize ) -> None: @@ -270,7 +312,10 @@ def link_state(handle, link_idx): raise pynvml.NVMLError_NotSupported() return link_idx != down_link - with patch("tensorrt_llm._mnnvl_utils.pynvml.nvmlDeviceGetNvLinkState", side_effect=link_state): + with patch( + "tensorrt_llm._torch.distributed.mnnvl_memory.pynvml.nvmlDeviceGetNvLinkState", + side_effect=link_state, + ): assert not MnnvlMemory.support_nvlink(0, True)