diff --git a/cpp/tensorrt_llm/kernels/communicationKernels/moeAlltoAllKernels.cu b/cpp/tensorrt_llm/kernels/communicationKernels/moeAlltoAllKernels.cu index 91cb5725fede..0c4431b8eaef 100644 --- a/cpp/tensorrt_llm/kernels/communicationKernels/moeAlltoAllKernels.cu +++ b/cpp/tensorrt_llm/kernels/communicationKernels/moeAlltoAllKernels.cu @@ -210,6 +210,14 @@ __device__ __forceinline__ int compute_target_rank_id(int expert_id, int base, i return remainder + (expert_id - split) / base; } +// Test bit `rank` in a kRankMaskWords-wide little-endian uint64 bitmask. +// Word 0 covers ranks 0..63, word 1 covers ranks 64..127, etc. +// `rank >> 6` and `rank & 63` divide / modulo by 64. +__device__ __forceinline__ bool is_rank_active(uint64_t const* mask, int rank) +{ + return (mask[rank >> 6] >> (rank & 63)) & 1ULL; +} + // ============================================================================ // Helper Functions for Vectorized Memory Operations // ============================================================================ @@ -416,7 +424,7 @@ __global__ void moeA2ADispatchKernel(int32_t const* token_selected_experts, // [ int* smem_topk_target_ranks = smem; int* smem_topk_send_indices = smem + TOP_K; - uint64_t already_copied = 0; + uint64_t already_copied[kRankMaskWords] = {}; // Precompute the ceil/floor partition parameters once per thread, outside the // per-token TOP_K loop. The fast path (remainder == 0) then collapses to a single // integer divide per call, matching the pre-PR uniform-partition cost exactly. @@ -432,7 +440,15 @@ __global__ void moeA2ADispatchKernel(int32_t const* token_selected_experts, // [ // Supports the non-divisible case where num_experts % ep_size != 0. int target_rank = compute_target_rank_id(expert_id, ep_base, ep_remainder); - if (already_copied & (1ULL << target_rank)) + // Skip duplicates AND dead ranks: both produce the same -1 sentinel that combine + // checks via topk_send_indices[k] < 0. A token whose only target is dead is dropped + // from this collective; higher-layer logic (EPLB redistribution) is responsible + // for re-routing such tokens on subsequent iterations. + int const mask_word = target_rank >> 6; + uint64_t const mask_bit = 1ULL << (target_rank & 63); + bool const target_already_copied = already_copied[mask_word] & mask_bit; + bool const target_dead = !is_rank_active(ptrs.active_rank_mask, target_rank); + if (target_already_copied || target_dead) { if (thread_idx == 0) { @@ -457,7 +473,7 @@ __global__ void moeA2ADispatchKernel(int32_t const* token_selected_experts, // [ smem_topk_target_ranks[k] = target_rank; smem_topk_send_indices[k] = dst_token_idx; } - already_copied |= 1ULL << target_rank; + already_copied[mask_word] |= mask_bit; } // Sync before dispatching data ThreadingPolicy::sync(); @@ -511,10 +527,13 @@ __global__ void moeA2ADispatchKernel(int32_t const* token_selected_experts, // [ if (is_last_token) { -// Store send_counters to recv_counters +// Store send_counters to recv_counters. +// Skip masked target ranks: their symmetric memory may be inaccessible. #pragma unroll 1 // No unroll as one iter is typically enough for (int target_rank = lane_id; target_rank < ep_size; target_rank += warpSize) { + if (!is_rank_active(ptrs.active_rank_mask, target_rank)) + continue; int send_count = ptrs.send_counters[target_rank]; ptrs.recv_counters[target_rank][rank_id] = send_count; } @@ -522,9 +541,12 @@ __global__ void moeA2ADispatchKernel(int32_t const* token_selected_experts, // [ if constexpr (ENABLE_EPLB) { // Write local stats into peer buffers before the release fence below. + // Skip masked target ranks for the same reason as above. #pragma unroll 1 for (int target_rank = 0; target_rank < ep_size; ++target_rank) { + if (!is_rank_active(ptrs.active_rank_mask, target_rank)) + continue; int* target_stats = ptrs.eplb_gathered_stats[target_rank]; for (int expert_id = lane_id; expert_id < eplb_stats_num_experts; expert_id += warpSize) { @@ -543,9 +565,13 @@ __global__ void moeA2ADispatchKernel(int32_t const* token_selected_experts, // [ #else asm volatile("fence.acq_rel.sys;"); #endif + // Signal completion to all active peers; skip dead ranks (their symmetric memory + // is unreachable). #pragma unroll 1 // No unroll as one iter is typically enough for (int target_rank = lane_id; target_rank < ep_size; target_rank += warpSize) { + if (!is_rank_active(ptrs.active_rank_mask, target_rank)) + 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)); @@ -555,9 +581,13 @@ __global__ void moeA2ADispatchKernel(int32_t const* token_selected_experts, // [ #endif } + // Wait for all active peers to signal; skip dead ranks (otherwise we would + // spin forever — this is the bug the rank-mask is here to prevent). #pragma unroll 1 // No unroll for (int peer_rank = lane_id; peer_rank < ep_size; peer_rank += warpSize) { + if (!is_rank_active(ptrs.active_rank_mask, peer_rank)) + continue; bool flag_set = false; auto s = clock64(); do @@ -603,8 +633,13 @@ void moe_a2a_dispatch_launch(MoeA2ADispatchParams const& params) // Validate parameters TLLM_CHECK(params.top_k > 0 && params.top_k <= kMaxTopK); TLLM_CHECK(params.ep_size > 0 && params.ep_size <= kMaxRanks); + TLLM_CHECK(params.ep_rank >= 0 && params.ep_rank < params.ep_size); TLLM_CHECK(params.local_num_tokens >= 0); TLLM_CHECK(params.num_payloads > 0 && params.num_payloads <= kMaxPayloads); + // The local rank must always be marked active in its own view of the mask; + // otherwise the kernel itself would be running on a "dead" rank. + TLLM_CHECK_WITH_INFO((params.active_rank_mask[params.ep_rank >> 6] >> (params.ep_rank & 63)) & 1ULL, + "active_rank_mask must mark the local ep_rank (%d) as active", params.ep_rank); // Prepare kernel pointers struct DispatchKernelPointers kernel_ptrs = {}; @@ -642,6 +677,12 @@ void moe_a2a_dispatch_launch(MoeA2ADispatchParams const& params) kernel_ptrs.topk_send_indices = params.topk_send_indices; kernel_ptrs.eplb_local_stats = params.eplb_local_stats; + // Copy active-rank bitmask into the kernel pointers struct + for (int w = 0; w < kRankMaskWords; ++w) + { + kernel_ptrs.active_rank_mask[w] = params.active_rank_mask[w]; + } + int const kBlockSize = tensorrt_llm::common::getEnvMoeA2ADispatchBlockSize(); // One block per token: grid_size == local_num_tokens. If 0, launch a single block to @@ -1153,9 +1194,13 @@ __global__ void moeA2ACombineKernel( if (blockIdx.x == 0) { + // Signal readiness to all active peers; skip dead ranks (their symmetric memory + // is unreachable). #pragma unroll 1 // No unroll for (int peer_rank = lane_id; peer_rank < ep_size; peer_rank += warpSize) { + if (!is_rank_active(ptrs.active_rank_mask, peer_rank)) + 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)); #if ENABLE_DEBUG_PRINT @@ -1165,9 +1210,13 @@ __global__ void moeA2ACombineKernel( } } + // Wait for all active peers to signal; skip dead ranks (otherwise we would spin + // forever — this is the bug the rank-mask is here to prevent). #pragma unroll 1 // No unroll for (int peer_rank = lane_id; peer_rank < ep_size; peer_rank += warpSize) { + if (!is_rank_active(ptrs.active_rank_mask, peer_rank)) + continue; bool flag_set = false; auto s = clock64(); do @@ -1271,8 +1320,13 @@ void moe_a2a_combine_launch(MoeA2ACombineParams const& params) // Validate parameters TLLM_CHECK(params.top_k > 0 && params.top_k <= kMaxTopK); TLLM_CHECK(params.ep_size > 0 && params.ep_size <= kMaxRanks); + TLLM_CHECK(params.ep_rank >= 0 && params.ep_rank < params.ep_size); TLLM_CHECK(params.local_num_tokens >= 0); TLLM_CHECK(params.elements_per_token > 0); + // The local rank must always be marked active in its own view of the mask; + // otherwise the kernel itself would be running on a "dead" rank. + TLLM_CHECK_WITH_INFO((params.active_rank_mask[params.ep_rank >> 6] >> (params.ep_rank & 63)) & 1ULL, + "active_rank_mask must mark the local ep_rank (%d) as active", params.ep_rank); // Configure kernel launch (one block per token). int const kBlockSize = tensorrt_llm::common::getEnvMoeA2ACombineBlockSize(); @@ -1306,6 +1360,12 @@ void moe_a2a_combine_launch(MoeA2ACombineParams const& params) kernel_ptrs.topk_target_ranks = params.topk_target_ranks; kernel_ptrs.topk_send_indices = params.topk_send_indices; + // Copy active-rank bitmask into the kernel pointers struct + for (int w = 0; w < kRankMaskWords; ++w) + { + kernel_ptrs.active_rank_mask[w] = params.active_rank_mask[w]; + } + // stride_per_token: byte distance between tokens in the recv buffer. // FP8 external payload: EPT × 1 (compact FP8 layout) // FP8 in-place / non-FP8: EPT × sizeof(PayloadT) (payload-dtype stride) diff --git a/cpp/tensorrt_llm/kernels/communicationKernels/moeAlltoAllKernels.h b/cpp/tensorrt_llm/kernels/communicationKernels/moeAlltoAllKernels.h index 317ff4d2240c..9a6f3904c501 100644 --- a/cpp/tensorrt_llm/kernels/communicationKernels/moeAlltoAllKernels.h +++ b/cpp/tensorrt_llm/kernels/communicationKernels/moeAlltoAllKernels.h @@ -26,9 +26,12 @@ namespace kernels::moe_comm { // Configuration constants -static constexpr int kMaxTopK = 22; // Maximum top-k experts per token -static constexpr int kMaxPayloads = 4; // Maximum number of different payload types -static constexpr int kMaxRanks = 64; // Maximum supported EP size +static constexpr int kMaxTopK = 22; // Maximum top-k experts per token +static constexpr int kMaxPayloads = 4; // Maximum number of different payload types +static constexpr int kMaxRanks = 128; // Maximum supported EP size (covers NVL72 with headroom) +static constexpr int kRankMaskWords = 2; // uint64 words to hold the active-rank bitmask + // (kRankMaskWords * 64 must be >= kMaxRanks) +static_assert(kRankMaskWords * 64 >= kMaxRanks, "active_rank_mask too small for kMaxRanks"); // Describes a single payload type to be communicated struct PayloadDescriptor @@ -65,6 +68,12 @@ struct DispatchKernelPointers // Optional: Statistics for EPLB int const* eplb_local_stats; // [eplb_stats_num_experts] int* eplb_gathered_stats[kMaxRanks]; // [ep_size, eplb_stats_num_experts] per rank + + // Active-rank bitmask: bit i set => rank i is alive and participates in this collective. + // Word 0 covers ranks 0..63; word 1 covers ranks 64..127. Tokens routed to a masked + // rank are dropped (topk_*[k] = -1); flag writes/waits to/from masked peers are skipped. + // The local rank's own bit must always be set; this is checked at launch time. + uint64_t active_rank_mask[kRankMaskWords]; }; // Combine kernel pointers - non-const output in src_data_ptrs[0], const recv buffers @@ -82,6 +91,11 @@ struct CombineKernelPointers // 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 duplicates int const* topk_send_indices; // dst index per k, -1 for duplicates + + // Active-rank bitmask: see DispatchKernelPointers::active_rank_mask. Combine skips flag + // writes/waits to/from masked peers; per-token accumulation uses topk_send_indices[k] < 0 + // (set by dispatch) to skip dead-targeted slots, so no explicit mask check is needed there. + uint64_t active_rank_mask[kRankMaskWords]; }; // Dispatch phase parameters @@ -125,6 +139,11 @@ struct MoeA2ADispatchParams int const* eplb_local_stats; // [eplb_stats_num_experts] int* eplb_gathered_stats[kMaxRanks]; // [ep_size, eplb_stats_num_experts] per rank + // Active-rank bitmask: see DispatchKernelPointers::active_rank_mask. The launch function + // copies these words into the kernel pointers struct. Defaults to all-ones for + // backwards-compatible "no masking" behavior. + uint64_t active_rank_mask[kRankMaskWords] = {~uint64_t{0}, ~uint64_t{0}}; + // CUDA stream cudaStream_t stream; }; @@ -170,6 +189,11 @@ struct MoeA2ACombineParams // rank has signaled the target rank void const* recv_buffers[kMaxRanks]; // Per-rank receive buffers (only for single payload) + // Active-rank bitmask: see DispatchKernelPointers::active_rank_mask. The launch function + // copies these words into the kernel pointers struct. Defaults to all-ones for + // backwards-compatible "no masking" behavior. + uint64_t active_rank_mask[kRankMaskWords] = {~uint64_t{0}, ~uint64_t{0}}; + // CUDA stream cudaStream_t stream; }; diff --git a/cpp/tensorrt_llm/thop/moeAlltoAllOp.cpp b/cpp/tensorrt_llm/thop/moeAlltoAllOp.cpp index 7a767976dabd..fc45afd792bb 100644 --- a/cpp/tensorrt_llm/thop/moeAlltoAllOp.cpp +++ b/cpp/tensorrt_llm/thop/moeAlltoAllOp.cpp @@ -42,6 +42,43 @@ inline size_t alignOffset(size_t offset, size_t alignment) return (offset + alignment - 1) & ~(alignment - 1); } +// Resolve an optional rank-mask tensor into a fixed-width uint64 array. +// If the caller did not provide a mask, default to "all ranks active" (all bits set), which +// reproduces the pre-fault-tolerance behavior bit-for-bit. +// +// On failure (wrong dtype / device / shape), throws via TORCH_CHECK so the error surfaces +// at the Python op boundary rather than the kernel launch. +inline void resolveActiveRankMask(torch::optional const& maskTensor, int64_t epRank, + uint64_t (&out)[tensorrt_llm::kernels::moe_comm::kRankMaskWords]) +{ + using tensorrt_llm::kernels::moe_comm::kRankMaskWords; + using tensorrt_llm::kernels::moe_comm::kMaxRanks; + TORCH_CHECK( + epRank >= 0 && epRank < kMaxRanks, "epRank must be in the range [0, ", kMaxRanks, ") for active_rank_mask"); + if (!maskTensor.has_value() || !maskTensor.value().defined()) + { + for (int w = 0; w < kRankMaskWords; ++w) + { + out[w] = ~uint64_t{0}; + } + return; + } + torch::Tensor const& t = maskTensor.value(); + TORCH_CHECK(t.is_cpu(), "active_rank_mask must be a CPU tensor"); + TORCH_CHECK(t.scalar_type() == torch::kUInt64, "active_rank_mask must have dtype uint64"); + TORCH_CHECK(t.dim() == 1, "active_rank_mask must be a 1D tensor"); + TORCH_CHECK(t.numel() == kRankMaskWords, "active_rank_mask must have exactly ", kRankMaskWords, " uint64 elements"); + TORCH_CHECK(t.is_contiguous(), "active_rank_mask must be contiguous"); + auto const* src = static_cast(t.const_data_ptr()); + for (int w = 0; w < kRankMaskWords; ++w) + { + out[w] = src[w]; + } + // Local rank's bit must be set; otherwise the kernel would be running on a "dead" rank. + TORCH_CHECK((out[epRank >> 6] >> (epRank & 63)) & 1ULL, "active_rank_mask must mark the local ep_rank (", epRank, + ") as active"); +} + // Calculate auxiliary data offsets MoeA2ADataOffsets calculateOffsets(int epSize, int maxNumTokens, int eplbStatsNumExperts) { @@ -117,11 +154,14 @@ MoeA2ADataOffsets calculateOffsets(int epSize, int maxNumTokens, int eplbStatsNu torch::Tensor moeA2AInitializeOp(torch::Tensor const& workspace, int64_t epRank, int64_t epSize, int64_t maxNumTokens, torch::optional eplbStatsNumExperts) { + using tensorrt_llm::kernels::moe_comm::kMaxRanks; + // Validate inputs 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)"); // Initialize workspace to zero @@ -181,13 +221,15 @@ torch::Tensor moeA2AInitializeOp(torch::Tensor const& workspace, int64_t epRank, std::tuple, int64_t, torch::Tensor> moeA2ADispatchOp( torch::Tensor const& tokenSelectedExperts, std::vector const& inputPayloads, torch::Tensor const& workspace, torch::Tensor const& metainfo, int64_t runtimeMaxTokensPerRank, int64_t epRank, - int64_t epSize, int64_t topK, int64_t numExperts, torch::optional eplbLocalStats) + int64_t epSize, int64_t topK, int64_t numExperts, torch::optional eplbLocalStats, + torch::optional activeRankMask) { using tensorrt_llm::kernels::moe_comm::PayloadDescriptor; using tensorrt_llm::kernels::moe_comm::MoeA2ADispatchParams; using tensorrt_llm::kernels::moe_comm::moe_a2a_dispatch_launch; using tensorrt_llm::kernels::moe_comm::kMaxTopK; using tensorrt_llm::kernels::moe_comm::kMaxPayloads; + using tensorrt_llm::kernels::moe_comm::kMaxRanks; // Validate inputs CHECK_INPUT(tokenSelectedExperts, torch::kInt32); @@ -203,6 +245,7 @@ std::tuple, int64_t, torch::Tensor> moeA2ADispatchOp( int64_t localNumTokens = tokenSelectedExperts.size(0); TORCH_CHECK(runtimeMaxTokensPerRank > 0, "runtimeMaxTokensPerRank must be positive"); + 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]"); TORCH_CHECK(!inputPayloads.empty(), "inputPayloads must not be empty"); @@ -360,6 +403,10 @@ std::tuple, int64_t, torch::Tensor> moeA2ADispatchOp( params.eplb_local_stats = nullptr; } + // Resolve the optional active-rank mask. Default (no mask) = all bits set, which + // exactly reproduces the pre-fault-tolerance kernel behavior. + resolveActiveRankMask(activeRankMask, epRank, params.active_rank_mask); + params.stream = at::cuda::getCurrentCUDAStream(); // Prepare for dispatch (zero counters/indices and increment flag_val) @@ -413,11 +460,13 @@ std::tuple, int64_t, torch::Tensor> moeA2ADispatchOp( // In both cases, the combine kernel reads from the workspace at 'combinePayloadOffset'. torch::Tensor moeA2ACombineOp(torch::Tensor const& payload, int64_t localNumTokens, torch::Tensor const& workspace, torch::Tensor const& metainfo, int64_t runtimeMaxTokensPerRank, int64_t epRank, int64_t epSize, int64_t topK, - int64_t combinePayloadOffset, bool payloadInWorkspace, bool useLowPrecision = false) + int64_t combinePayloadOffset, bool payloadInWorkspace, bool useLowPrecision = false, + torch::optional activeRankMask = torch::nullopt) { using tensorrt_llm::kernels::moe_comm::MoeA2ACombineParams; using tensorrt_llm::kernels::moe_comm::moe_a2a_combine_launch; using tensorrt_llm::kernels::moe_comm::kMaxTopK; + using tensorrt_llm::kernels::moe_comm::kMaxRanks; // Validate inputs CHECK_TH_CUDA(payload); @@ -431,6 +480,7 @@ torch::Tensor moeA2ACombineOp(torch::Tensor const& payload, int64_t localNumToke TORCH_CHECK(reinterpret_cast(payload.data_ptr()) % 16 == 0, "payload must be 16-byte aligned"); int64_t elementsPerToken = payload.size(2); TORCH_CHECK(elementsPerToken > 0, "elementsPerToken must be positive"); + 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]"); @@ -520,6 +570,9 @@ torch::Tensor moeA2ACombineOp(torch::Tensor const& payload, int64_t localNumToke params.recv_buffers[target_rank] = target_workspace_ptr + combinePayloadOffset; } + // Resolve the optional active-rank mask. Default (no mask) = all bits set. + resolveActiveRankMask(activeRankMask, epRank, params.active_rank_mask); + params.stream = at::cuda::getCurrentCUDAStream(); moe_a2a_prepare_combine_launch(params); @@ -613,12 +666,14 @@ TORCH_LIBRARY_FRAGMENT(trtllm, module) "moe_a2a_dispatch(Tensor token_selected_experts, Tensor[] input_payloads, " "Tensor(a!->*) workspace, Tensor metainfo, int runtime_max_tokens_per_rank, " "int ep_rank, int ep_size, int top_k, int num_experts, " - "Tensor? eplb_local_stats=None) -> (Tensor(a!)[], int, Tensor(a!))"); + "Tensor? eplb_local_stats=None, " + "Tensor? active_rank_mask=None) -> (Tensor(a!)[], int, Tensor(a!))"); module.def( "moe_a2a_combine(Tensor(a) payload, int local_num_tokens," "Tensor(a!) workspace, Tensor metainfo, int runtime_max_tokens_per_rank, " "int ep_rank, int ep_size, int top_k, int combine_payload_offset, " - "bool payload_in_workspace, bool use_low_precision=False) -> Tensor"); + "bool payload_in_workspace, bool use_low_precision=False, " + "Tensor? active_rank_mask=None) -> Tensor"); 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) -> Tensor"); diff --git a/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py b/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py index a8e113278e75..35d6f69fd626 100644 --- a/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py +++ b/tensorrt_llm/_torch/custom_ops/cpp_custom_ops.py @@ -480,6 +480,7 @@ def _( top_k: int, num_experts: int, eplb_local_stats: Optional[torch.Tensor] = None, + active_rank_mask: Optional[torch.Tensor] = None, ) -> Tuple[List[torch.Tensor], int, torch.Tensor]: recv_tensors: List[torch.Tensor] = [] for payload in input_payloads: @@ -510,6 +511,7 @@ def _( combine_payload_offset: int, payload_in_workspace: bool, use_low_precision: bool = False, + active_rank_mask: Optional[torch.Tensor] = None, ) -> torch.Tensor: return payload.new_empty((local_num_tokens, payload.shape[2])) diff --git a/tensorrt_llm/_torch/modules/fused_moe/communication/nvlink_one_sided.py b/tensorrt_llm/_torch/modules/fused_moe/communication/nvlink_one_sided.py index 3b634dd7072c..f9068f71d717 100644 --- a/tensorrt_llm/_torch/modules/fused_moe/communication/nvlink_one_sided.py +++ b/tensorrt_llm/_torch/modules/fused_moe/communication/nvlink_one_sided.py @@ -51,7 +51,7 @@ class NVLinkOneSided(Communication): """ # Constants from C++ (must match moeAlltoAllKernels.h) - MAX_RANKS = 64 + MAX_RANKS = 128 MAX_TOP_K = 8 MAX_PAYLOADS = 8 diff --git a/tests/unittest/_torch/modules/moe/test_moe_comm.py b/tests/unittest/_torch/modules/moe/test_moe_comm.py index c59751f42025..057f87a95f7a 100644 --- a/tests/unittest/_torch/modules/moe/test_moe_comm.py +++ b/tests/unittest/_torch/modules/moe/test_moe_comm.py @@ -73,6 +73,7 @@ NVLinkTwoSidedFlashinfer, ) from tensorrt_llm._torch.modules.fused_moe.deep_ep_utils import deep_ep_installed +from tensorrt_llm._torch.modules.fused_moe.ep_group_health import EPGroupHealth from tensorrt_llm.deep_ep.buffer import Buffer from tensorrt_llm.mapping import Mapping @@ -114,7 +115,6 @@ # to avoid _WORKSPACE singleton assertion failures. NVLINK_WORKSPACE_MB = "512" - # ============================================================================ # Test Configuration # ============================================================================ @@ -192,6 +192,84 @@ def _safe_cpu(t: Optional[torch.Tensor]) -> Optional[torch.Tensor]: return t.cpu() +def _ep_mask_words(ep_size: int, dead_ranks: Set[int]) -> torch.Tensor: + """Build the CPU active-rank mask tensor expected by moe_a2a ops.""" + health = EPGroupHealth(ep_size) + for rank in dead_ranks: + health.mark_failed(rank) + return torch.tensor(health.get_mask_words(), dtype=torch.uint64, device="cpu") + + +def _make_rank_mask_payload(local_num_tokens: int, hidden_size: int, rank: int) -> torch.Tensor: + """Make deterministic per-rank payloads for exact equality assertions.""" + base = torch.arange(local_num_tokens * hidden_size, dtype=torch.bfloat16, device="cuda").view( + local_num_tokens, hidden_size + ) + return base + (rank * 1000.0) + + +def _read_nvlink_topk_target_ranks( + comm: NVLinkOneSided, + max_num_tokens: int, + top_k: int, +) -> torch.Tensor: + """Read topk_target_ranks[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_TARGET_RANKS_OFFSET_INDEX) + offset = comm.moe_a2a_metainfo[offset_index].item() + raw = comm.workspace[ + comm.ep_rank, + offset : offset + max_num_tokens * top_k * 4, + ] + return raw.view(torch.int32).view(max_num_tokens, top_k).cpu() + + +def _run_nvlink_rank_mask_dispatch_combine( + comm: NVLinkOneSided, + token_selected_experts: torch.Tensor, + payload: torch.Tensor, + runtime_max_tokens_per_rank: int, + active_rank_mask: Optional[torch.Tensor], +) -> Tuple[torch.Tensor, torch.Tensor]: + """Run raw NVLink one-sided dispatch/combine with an optional active rank mask.""" + recv_tensors, combine_payload_offset, _ = torch.ops.trtllm.moe_a2a_dispatch( + token_selected_experts, + [payload], + comm.workspace, + comm.moe_a2a_metainfo, + runtime_max_tokens_per_rank, + comm.ep_rank, + comm.ep_size, + comm.top_k, + comm.num_experts, + None, # eplb_local_stats + active_rank_mask, + ) + + topk_target_ranks = _read_nvlink_topk_target_ranks( + comm, + runtime_max_tokens_per_rank, + comm.top_k, + ) + + combined = torch.ops.trtllm.moe_a2a_combine( + recv_tensors[0], + token_selected_experts.size(0), + comm.workspace, + comm.moe_a2a_metainfo, + runtime_max_tokens_per_rank, + comm.ep_rank, + comm.ep_size, + comm.top_k, + int(combine_payload_offset), + False, # payload_in_workspace + False, # use_low_precision + active_rank_mask, + ) + return combined.cpu(), topk_target_ranks + + # ============================================================================ # Source Encoding Utilities # ============================================================================ @@ -856,6 +934,154 @@ def _worker_full_pipeline(config: CommTestConfig) -> dict: comm.destroy() +def _make_rank_mask_config( + ep_size: int, + local_num_tokens: int, + top_k: int, +) -> CommTestConfig: + """Build the small NVLinkOneSided config used by active-rank-mask tests.""" + return CommTestConfig( + comm_type=COMM_NVLINK_ONE_SIDED, + ep_size=ep_size, + num_experts=FIXED_NUM_EXPERTS, + top_k=top_k, + hidden_size=1024, + all_num_tokens=[local_num_tokens] * ep_size, + ) + + +def _worker_rank_mask_all_active_matches_no_mask(config: CommTestConfig) -> dict: + """Check that all-active active_rank_mask is bit-identical to no mask.""" + rank = tllm.mpi_rank() + torch.cuda.set_device(rank) + + comm = None + try: + mapping = Mapping( + rank=rank, + tp_size=config.ep_size, + moe_ep_size=config.ep_size, + world_size=config.ep_size, + ) + comm = create_comm_object(config.comm_type, mapping, config) + + local_num_tokens = config.all_num_tokens[rank] + torch.manual_seed(0xA2A + rank) + token_selected_experts = torch.randint( + 0, + config.num_experts, + (local_num_tokens, config.top_k), + dtype=torch.int32, + device="cuda", + ) + payload = _make_rank_mask_payload(local_num_tokens, config.hidden_size, rank) + + out_no_mask, topk_no_mask = _run_nvlink_rank_mask_dispatch_combine( + comm, + token_selected_experts, + payload, + local_num_tokens, + active_rank_mask=None, + ) + out_all_active, topk_all_active = _run_nvlink_rank_mask_dispatch_combine( + comm, + token_selected_experts, + payload, + local_num_tokens, + active_rank_mask=_ep_mask_words(config.ep_size, dead_ranks=set()), + ) + + return { + "output_eq": torch.equal(out_no_mask, out_all_active), + "topk_eq": torch.equal(topk_no_mask, topk_all_active), + } + except Exception: + traceback.print_exc() + raise + finally: + if comm is not None and hasattr(comm, "destroy"): + comm.destroy() + + +def _expected_target_ranks( + token_selected_experts: torch.Tensor, + num_experts: int, + ep_size: int, +) -> torch.Tensor: + """Map each selected expert to its target EP rank using the kernel partition rule.""" + token_selected_experts_cpu = token_selected_experts.cpu() + expected = torch.empty_like(token_selected_experts_cpu) + for token_idx in range(token_selected_experts_cpu.shape[0]): + for k in range(token_selected_experts_cpu.shape[1]): + expert_id = int(token_selected_experts_cpu[token_idx, k].item()) + expected[token_idx, k] = _expert_id_to_rank(expert_id, num_experts, ep_size) + return expected + + +def _worker_rank_mask_one_rank_masked( + config: CommTestConfig, + dead_rank: int, +) -> dict: + """Run dispatch/combine with one EP rank omitted from active_rank_mask.""" + rank = tllm.mpi_rank() + torch.cuda.set_device(rank) + + comm = None + try: + mapping = Mapping( + rank=rank, + tp_size=config.ep_size, + moe_ep_size=config.ep_size, + world_size=config.ep_size, + ) + # All ranks must initialize the symmetric workspace before the dead rank + # stops participating in dispatch/combine. + comm = create_comm_object(config.comm_type, mapping, config) + + if rank == dead_rank: + MPI.COMM_WORLD.barrier() + return {"status": "dead"} + + local_num_tokens = config.all_num_tokens[rank] + torch.manual_seed(0xA2A + rank) + token_selected_experts = torch.randint( + 0, + config.num_experts, + (local_num_tokens, config.top_k), + dtype=torch.int32, + device="cuda", + ) + payload = _make_rank_mask_payload(local_num_tokens, config.hidden_size, rank) + mask = _ep_mask_words(config.ep_size, dead_ranks={dead_rank}) + + combined, topk_target_ranks = _run_nvlink_rank_mask_dispatch_combine( + comm, + token_selected_experts, + payload, + local_num_tokens, + active_rank_mask=mask, + ) + expected_target_ranks = _expected_target_ranks( + token_selected_experts, + config.num_experts, + config.ep_size, + ) + + MPI.COMM_WORLD.barrier() + return { + "status": "alive", + "combined": combined, + "topk_target_ranks": topk_target_ranks, + "expected_target_ranks": expected_target_ranks, + } + except Exception: + traceback.print_exc() + raise + finally: + if comm is not None and hasattr(comm, "destroy"): + comm.destroy() + + # ============================================================================ # Verification Functions # ============================================================================ @@ -1630,6 +1856,102 @@ def _run_full_test(mpi_pool_executor, config: CommTestConfig): verify_combine_results(all_results, config, rtol=0.02, atol=0.15) +def _skip_if_rank_mask_config_unsupported(config: CommTestConfig) -> None: + """Skip active-rank-mask tests when NVLinkOneSided cannot run locally.""" + skip_reason = check_platform_support(config.comm_type) + if skip_reason: + pytest.skip(skip_reason) + + skip_reason = check_feasibility(config.comm_type, config) + if skip_reason: + pytest.skip(skip_reason) + + if config.ep_size > torch.cuda.device_count(): + pytest.skip(f"Need {config.ep_size} GPUs but only {torch.cuda.device_count()} available") + + +def _run_rank_mask_all_active_test( + mpi_pool_executor, + local_num_tokens: int, + top_k: int, +) -> None: + ep_size = mpi_pool_executor.num_workers + config = _make_rank_mask_config(ep_size, local_num_tokens, top_k) + _skip_if_rank_mask_config_unsupported(config) + + results = list( + mpi_pool_executor.map( + _worker_rank_mask_all_active_matches_no_mask, + *zip(*[(config,)] * config.ep_size), + ) + ) + + for rank, result in enumerate(results): + assert result["output_eq"], ( + f"rank {rank}: combine output differs between no-mask and all-active mask" + ) + assert result["topk_eq"], ( + f"rank {rank}: topk_target_ranks differ between no-mask and all-active mask" + ) + + +def _run_rank_mask_one_rank_masked_test( + mpi_pool_executor, + dead_rank: int, + local_num_tokens: int, + top_k: int, +) -> None: + ep_size = mpi_pool_executor.num_workers + config = _make_rank_mask_config(ep_size, local_num_tokens, top_k) + _skip_if_rank_mask_config_unsupported(config) + assert 0 <= dead_rank < ep_size + + worker_args = [(config, dead_rank)] * config.ep_size + results = list( + mpi_pool_executor.map( + _worker_rank_mask_one_rank_masked, + *zip(*worker_args), + ) + ) + + saw_dead = False + for rank, result in enumerate(results): + if result["status"] == "dead": + assert rank == dead_rank + saw_dead = True + continue + + assert result["status"] == "alive" + combined = result["combined"] + topk_target_ranks = result["topk_target_ranks"] + expected_target_ranks = result["expected_target_ranks"] + + assert combined.shape == (local_num_tokens, config.hidden_size) + + live_topk = topk_target_ranks[:local_num_tokens] + live_expected = expected_target_ranks[:local_num_tokens] + for token_idx in range(local_num_tokens): + seen_ranks: Set[int] = set() + for k in range(top_k): + expected = int(live_expected[token_idx, k].item()) + got = int(live_topk[token_idx, k].item()) + if expected == dead_rank: + assert got == -1, ( + f"rank {rank} token {token_idx} k={k}: token routed to dead " + f"rank {dead_rank} should have been dropped (got={got})" + ) + elif expected in seen_ranks: + assert got == -1 + else: + assert got == expected, ( + f"rank {rank} token {token_idx} k={k}: target rank mismatch " + f"(expected={expected}, got={got})" + ) + seen_ranks.add(expected) + + assert saw_dead, f"dead rank {dead_rank} did not appear in results" + + # ============================================================================ # Test Class # ============================================================================ @@ -1682,3 +2004,46 @@ def test_moe_comm_postquant(self, mpi_pool_executor, config: CommTestConfig): def test_moe_comm_non_divisible_ep(self, mpi_pool_executor, config: CommTestConfig): """Verify NVLinkOneSided with non-divisible EP (num_experts % ep_size != 0).""" _run_full_test(mpi_pool_executor, config) + + @pytest.mark.threadleak(enabled=False) + @pytest.mark.parametrize( + "mpi_pool_executor,local_num_tokens,top_k", + [ + (4, 16, 2), + (4, 32, 4), + ], + indirect=["mpi_pool_executor"], + ) + def test_moe_comm_rank_mask_all_active_matches_no_mask( + self, + mpi_pool_executor, + local_num_tokens: int, + top_k: int, + ): + """Verify all-active active_rank_mask matches omitted mask for NVLinkOneSided.""" + _run_rank_mask_all_active_test(mpi_pool_executor, local_num_tokens, top_k) + + @pytest.mark.threadleak(enabled=False) + @pytest.mark.parametrize( + "mpi_pool_executor,dead_rank,local_num_tokens,top_k", + [ + (4, 2, 16, 2), + (4, 0, 16, 4), + (4, 3, 32, 4), + ], + indirect=["mpi_pool_executor"], + ) + def test_moe_comm_rank_mask_one_rank_masked_completes( + self, + mpi_pool_executor, + dead_rank: int, + local_num_tokens: int, + top_k: int, + ): + """Verify masked-dead rank is skipped by raw NVLinkOneSided moe_a2a ops.""" + _run_rank_mask_one_rank_masked_test( + mpi_pool_executor, + dead_rank, + local_num_tokens, + top_k, + )