From 1479e839aefd752b30d3f47c3442c22484c13d12 Mon Sep 17 00:00:00 2001 From: PanZezhong Date: Tue, 28 Jul 2026 15:39:59 +0800 Subject: [PATCH] fix: moe ep --- .../nvidia/moe_fused_dense_nvidia.cu | 178 ++++++++++++------ test/infinicore/ops/moe_fused_dense.py | 15 ++ 2 files changed, 137 insertions(+), 56 deletions(-) diff --git a/src/infiniop/ops/moe_fused_dense/nvidia/moe_fused_dense_nvidia.cu b/src/infiniop/ops/moe_fused_dense/nvidia/moe_fused_dense_nvidia.cu index b08ea3607..b1a64e2c7 100644 --- a/src/infiniop/ops/moe_fused_dense/nvidia/moe_fused_dense_nvidia.cu +++ b/src/infiniop/ops/moe_fused_dense/nvidia/moe_fused_dense_nvidia.cu @@ -71,7 +71,6 @@ size_t workspace_size(const MoeFusedDenseInfo &info) { bytes += align_up(info.num_experts * sizeof(cutlass::gemm::GemmCoord), 16); bytes += align_up(info.num_experts * sizeof(void *), 16) * 4; bytes += align_up(info.num_experts * sizeof(int64_t), 16) * 4; - bytes += align_up(sizeof(int), 16); return bytes + 256; #endif } @@ -291,29 +290,59 @@ __global__ void setup_prefill_gemm2_kernel(cutlass::gemm::GemmCoord *problems, } template -__global__ void setup_decode_gemm1_aligned_compact_kernel(cutlass::gemm::GemmCoord *problems, - void **ptr_a, - void **ptr_b, - void **ptr_c, - void **ptr_d, - int64_t *lda, - int64_t *ldb, - int64_t *ldc, - int64_t *ldd, - int *active_count, - int *output_permutation, - const T *hidden, - const T *w13, - T *gate_up, - const int *sorted_token_ids, - const int *expert_ids, - const int *num_tokens_post_padded, - int pairs, - int topk, - int num_experts, - int hidden_size, - int intermediate_size, - int block_size) { +__global__ void setup_decode_gemm1_defaults_kernel(cutlass::gemm::GemmCoord *problems, + void **ptr_a, + void **ptr_b, + void **ptr_c, + void **ptr_d, + int64_t *lda, + int64_t *ldb, + int64_t *ldc, + int64_t *ldd, + const T *hidden, + const T *w13, + T *gate_up, + int topk, + int hidden_size, + int intermediate_size) { + int route = blockIdx.x * blockDim.x + threadIdx.x; + if (route >= topk) { + return; + } + problems[route] = cutlass::gemm::GemmCoord(1, intermediate_size * 2, hidden_size); + ptr_a[route] = const_cast(hidden); + ptr_b[route] = const_cast(w13); + ptr_c[route] = gate_up + static_cast(route) * intermediate_size * 2; + ptr_d[route] = gate_up + static_cast(route) * intermediate_size * 2; + lda[route] = hidden_size; + ldb[route] = hidden_size; + ldc[route] = intermediate_size * 2; + ldd[route] = intermediate_size * 2; +} + +template +__global__ void setup_decode_gemm1_local_kernel(cutlass::gemm::GemmCoord *problems, + void **ptr_a, + void **ptr_b, + void **ptr_c, + void **ptr_d, + int64_t *lda, + int64_t *ldb, + int64_t *ldc, + int64_t *ldd, + int *output_permutation, + const T *hidden, + const T *w13, + T *gate_up, + const int *sorted_token_ids, + const int *expert_ids, + const int *num_tokens_post_padded, + int pairs, + int topk, + int num_experts, + int hidden_size, + int intermediate_size, + int block_size) { int row = blockIdx.x * blockDim.x + threadIdx.x; int total_rows = *num_tokens_post_padded; if (row >= total_rows) { @@ -328,7 +357,7 @@ __global__ void setup_decode_gemm1_aligned_compact_kernel(cutlass::gemm::GemmCoo return; } - int idx = atomicAdd(active_count, 1); + int idx = pair; int token = pair / topk; output_permutation[pair] = idx; problems[idx] = cutlass::gemm::GemmCoord(1, intermediate_size * 2, hidden_size); @@ -343,27 +372,58 @@ __global__ void setup_decode_gemm1_aligned_compact_kernel(cutlass::gemm::GemmCoo } template -__global__ void setup_decode_gemm2_aligned_compact_kernel(cutlass::gemm::GemmCoord *problems, - void **ptr_a, - void **ptr_b, - void **ptr_c, - void **ptr_d, - int64_t *lda, - int64_t *ldb, - int64_t *ldc, - int64_t *ldd, - const T *activated, - const T *w2, - T *expert_out, - const int *sorted_token_ids, - const int *expert_ids, - const int *num_tokens_post_padded, - const int *output_permutation, - int pairs, - int num_experts, - int hidden_size, - int intermediate_size, - int block_size) { +__global__ void setup_decode_gemm2_defaults_kernel(cutlass::gemm::GemmCoord *problems, + void **ptr_a, + void **ptr_b, + void **ptr_c, + void **ptr_d, + int64_t *lda, + int64_t *ldb, + int64_t *ldc, + int64_t *ldd, + const T *activated, + const T *w2, + T *expert_out, + int topk, + int hidden_size, + int intermediate_size) { + int route = blockIdx.x * blockDim.x + threadIdx.x; + if (route >= topk) { + return; + } + problems[route] = cutlass::gemm::GemmCoord(1, hidden_size, intermediate_size); + ptr_a[route] = const_cast(activated + static_cast(route) * intermediate_size); + ptr_b[route] = const_cast(w2); + ptr_c[route] = expert_out + static_cast(route) * hidden_size; + ptr_d[route] = expert_out + static_cast(route) * hidden_size; + lda[route] = intermediate_size; + ldb[route] = intermediate_size; + ldc[route] = hidden_size; + ldd[route] = hidden_size; +} + +template +__global__ void setup_decode_gemm2_local_kernel(cutlass::gemm::GemmCoord *problems, + void **ptr_a, + void **ptr_b, + void **ptr_c, + void **ptr_d, + int64_t *lda, + int64_t *ldb, + int64_t *ldc, + int64_t *ldd, + const T *activated, + const T *w2, + T *expert_out, + const int *sorted_token_ids, + const int *expert_ids, + const int *num_tokens_post_padded, + const int *output_permutation, + int pairs, + int num_experts, + int hidden_size, + int intermediate_size, + int block_size) { int row = blockIdx.x * blockDim.x + threadIdx.x; int total_rows = *num_tokens_post_padded; if (row >= total_rows) { @@ -521,9 +581,7 @@ infiniStatus_t calculate_typed(const MoeFusedDenseInfo &info, auto grouped_ldb = reinterpret_cast(advance_workspace(ptr, remaining, static_cast(num_experts) * sizeof(int64_t), alignof(int64_t))); auto grouped_ldc = reinterpret_cast(advance_workspace(ptr, remaining, static_cast(num_experts) * sizeof(int64_t), alignof(int64_t))); auto grouped_ldd = reinterpret_cast(advance_workspace(ptr, remaining, static_cast(num_experts) * sizeof(int64_t), alignof(int64_t))); - auto active_count = reinterpret_cast(advance_workspace(ptr, remaining, sizeof(int), alignof(int))); - - if (!counts || !offsets || !output_permutation || !packed_hidden || !gate_up || !activated || !expert_out || !grouped_problems || !grouped_ptr_a || !grouped_ptr_b || !grouped_ptr_c || !grouped_ptr_d || !grouped_lda || !grouped_ldb || !grouped_ldc || !grouped_ldd || !active_count) { + if (!counts || !offsets || !output_permutation || !packed_hidden || !gate_up || !activated || !expert_out || !grouped_problems || !grouped_ptr_a || !grouped_ptr_b || !grouped_ptr_c || !grouped_ptr_d || !grouped_lda || !grouped_ldb || !grouped_ldc || !grouped_ldd) { return INFINI_STATUS_INSUFFICIENT_WORKSPACE; } @@ -531,26 +589,34 @@ infiniStatus_t calculate_typed(const MoeFusedDenseInfo &info, auto w2_t = reinterpret_cast(w2); if (num_tokens == 1 && topk <= num_experts) { cudaMemsetAsync(output_permutation, 0xff, pairs * sizeof(int), stream); - cudaMemsetAsync(active_count, 0, sizeof(int), stream); - setup_decode_gemm1_aligned_compact_kernel<<<(max_num_tokens_padded + 255) / 256, 256, 0, stream>>>( + setup_decode_gemm1_defaults_kernel<<<(topk + 255) / 256, 256, 0, stream>>>( grouped_problems, grouped_ptr_a, grouped_ptr_b, grouped_ptr_c, grouped_ptr_d, - grouped_lda, grouped_ldb, grouped_ldc, grouped_ldd, active_count, output_permutation, + grouped_lda, grouped_ldb, grouped_ldc, grouped_ldd, + reinterpret_cast(hidden_states), w13_t, gate_up, + topk, hidden_size, intermediate_size); + setup_decode_gemm1_local_kernel<<<(max_num_tokens_padded + 255) / 256, 256, 0, stream>>>( + grouped_problems, grouped_ptr_a, grouped_ptr_b, grouped_ptr_c, grouped_ptr_d, + grouped_lda, grouped_ldb, grouped_ldc, grouped_ldd, output_permutation, reinterpret_cast(hidden_states), w13_t, gate_up, reinterpret_cast(sorted_token_ids), reinterpret_cast(expert_ids), reinterpret_cast(num_tokens_post_padded), pairs, topk, num_experts, hidden_size, intermediate_size, block_size); - // Decode has one token and exactly one valid aligned row for each of - // its top-k routes, so the compact problem count is always topk. - // A D2H count copy and stream sync cannot be captured for replay. + // Keep the problem count fixed at topk for graph replay. Routes owned + // by another EP rank retain valid dummy descriptors and are ignored + // by output_permutation when local expert outputs are combined. CHECK_STATUS(launch_cutlass_gemm_grouped_device_meta( topk, grouped_problems, grouped_ptr_a, grouped_ptr_b, grouped_ptr_c, grouped_ptr_d, grouped_lda, grouped_ldb, grouped_ldc, grouped_ldd, stream)); swiglu_kernel<<<(topk * intermediate_size + 255) / 256, 256, 0, stream>>>(gate_up, activated, topk, intermediate_size); - setup_decode_gemm2_aligned_compact_kernel<<<(max_num_tokens_padded + 255) / 256, 256, 0, stream>>>( + setup_decode_gemm2_defaults_kernel<<<(topk + 255) / 256, 256, 0, stream>>>( + grouped_problems, grouped_ptr_a, grouped_ptr_b, grouped_ptr_c, grouped_ptr_d, + grouped_lda, grouped_ldb, grouped_ldc, grouped_ldd, + activated, w2_t, expert_out, topk, hidden_size, intermediate_size); + setup_decode_gemm2_local_kernel<<<(max_num_tokens_padded + 255) / 256, 256, 0, stream>>>( grouped_problems, grouped_ptr_a, grouped_ptr_b, grouped_ptr_c, grouped_ptr_d, grouped_lda, grouped_ldb, grouped_ldc, grouped_ldd, activated, w2_t, expert_out, reinterpret_cast(sorted_token_ids), diff --git a/test/infinicore/ops/moe_fused_dense.py b/test/infinicore/ops/moe_fused_dense.py index b0ab5c44e..3cc6a6907 100644 --- a/test/infinicore/ops/moe_fused_dense.py +++ b/test/infinicore/ops/moe_fused_dense.py @@ -115,6 +115,8 @@ def torch_moe_fused_dense( for token in range(num_tokens): for route in range(topk): expert = int(ids[token, route].item()) + if expert < 0 or expert >= w13f.shape[0]: + continue gate_up = F.linear(hidden[token], w13f[expert]) gate, up = gate_up.chunk(2, dim=-1) activated = F.silu(gate) * up @@ -138,6 +140,19 @@ def parse_test_cases(): block_size=4, ), ), + ( + "decode EP-filtered routes", + _make_case( + 20260728, + num_tokens=1, + hidden_size=32, + intermediate_size=64, + num_experts=4, + topk=4, + topk_ids=torch.tensor([[6, 1, 3, 7]], dtype=torch.int32), + block_size=4, + ), + ), ( "prefill aligned missing experts", _make_case(