From 1a4160658c5e2ddfb6c643edb8686970221d76a7 Mon Sep 17 00:00:00 2001 From: samsja Date: Wed, 19 Aug 2026 22:33:35 +0000 Subject: [PATCH] feat(low_latency): support top-16 routing and hidden-dim 3584 MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The low_latency internode dispatch/combine kernels cap top-k at 11 and only template-specialize on a fixed set of a2a hidden dims. Two configs don't fit (both hit by Kimi-K3 — 896 experts, top-16, routed-expert hidden 3584): - top-16 routing trips EP_HOST_ASSERT(num_topk <= kNumMaxTopK) (dispatch) and the kNumMaxTopk equivalent (combine) -> raise both 11 -> 16. - routed-expert hidden 3584 falls through SWITCH_HIDDEN's default -> add case 3584 (satisfies the LL alignment invariants: %128, %256, and %512 for the combine kNumSendUnrolls=2 path). --- csrc/kernels/legacy/internode_ll.cu | 4 ++-- csrc/kernels/legacy/launch.cuh | 2 ++ 2 files changed, 4 insertions(+), 2 deletions(-) diff --git a/csrc/kernels/legacy/internode_ll.cu b/csrc/kernels/legacy/internode_ll.cu index b6be356ae..cc36718a4 100644 --- a/csrc/kernels/legacy/internode_ll.cu +++ b/csrc/kernels/legacy/internode_ll.cu @@ -702,7 +702,7 @@ void dispatch(void* packed_recv_x, int num_device_sms, cudaStream_t stream, int phases) { - constexpr int kNumMaxTopK = 11; + constexpr int kNumMaxTopK = 16; // Kimi-K3 routes top-16 const int num_warp_groups = ceil_div(num_experts, num_device_sms); const int num_warps_per_group = 32 / num_warp_groups; EP_HOST_ASSERT(num_warp_groups > 0 and num_warps_per_group > 0); @@ -1880,7 +1880,7 @@ void combine(void* combined_x, ); } - constexpr int kNumMaxTopk = 11; + constexpr int kNumMaxTopk = 16; // Kimi-K3 routes top-16 const int num_warp_groups = ceil_div(num_experts, num_device_sms); const int num_warps_per_group = 32 / num_warp_groups; const int num_recv_per_sm = ceil_div(num_combined_tokens, num_device_sms); diff --git a/csrc/kernels/legacy/launch.cuh b/csrc/kernels/legacy/launch.cuh index 60e308039..36351e52d 100644 --- a/csrc/kernels/legacy/launch.cuh +++ b/csrc/kernels/legacy/launch.cuh @@ -118,6 +118,8 @@ case_macro(2560); \ case 3072: \ case_macro(3072); /* for gpt-oss */ \ + case 3584: \ + case_macro(3584); /* for kimi-k3 */ \ case 4096: \ case_macro(4096); \ case 5120: \