From 30e7a3072d3f1e9bd70cd5e44146ca27c80522c4 Mon Sep 17 00:00:00 2001 From: Brayden Zhong Date: Tue, 8 Sep 2026 11:56:40 -0400 Subject: [PATCH] Keep fp32 routing weights in the fp8 block-scale and bf16 trtllm MoE (#33631) Co-authored-by: Brayden Zhong --- python/sglang/kernels/ops/moe/__init__.py | 10 -- .../sglang/kernels/ops/moe/pack_topk_ids.py | 101 ------------------ .../srt/layers/moe/flashinfer_trtllm_moe.py | 4 +- .../moe/moe_runner/flashinfer_trtllm.py | 59 +++++----- .../ops/test_kimi_k3_prerequisite_ops.py | 15 ++- 5 files changed, 51 insertions(+), 138 deletions(-) delete mode 100644 python/sglang/kernels/ops/moe/pack_topk_ids.py diff --git a/python/sglang/kernels/ops/moe/__init__.py b/python/sglang/kernels/ops/moe/__init__.py index 51e40eea3..7f06a870b 100644 --- a/python/sglang/kernels/ops/moe/__init__.py +++ b/python/sglang/kernels/ops/moe/__init__.py @@ -175,16 +175,6 @@ for _mod, _fn in _PHASE25_TRITON_KERNELS: ) del _mod, _fn -# Packed (topk_id << 16 | bf16-weight) kernel migrated from -# srt/layers/quantization/mxfp4_flashinfer_trtllm_moe (RFC #29630, Phase 2.5). -register_kernel( - KernelSpec( - op="moe.pack_topk_ids", - backend=KernelBackend.TRITON, - target="sglang.kernels.ops.moe.pack_topk_ids:PackTopkIds.triton", - ) -) - # Single-CTA align for tiny batches: covers the corner the AOT/JIT # moe_align_block_size small-batch path leaves out (num_experts > 64), and is # selected by the moe_runner call site on numel <= SMALL_NUMEL_LIMIT. diff --git a/python/sglang/kernels/ops/moe/pack_topk_ids.py b/python/sglang/kernels/ops/moe/pack_topk_ids.py deleted file mode 100644 index 437aef580..000000000 --- a/python/sglang/kernels/ops/moe/pack_topk_ids.py +++ /dev/null @@ -1,101 +0,0 @@ -"""Pack ``(topk_id, topk_weight)`` pairs into one int32 per entry. - -Migrated from ``sglang.srt.layers.quantization.mxfp4_flashinfer_trtllm_moe`` -(RFC #29630, Phase 2.5). Used by the FlashInfer TRT-LLM routed-MoE path, which -consumes routing ids and bf16 weights packed as ``(id << 16) | weight_bits``. -Routing ids may be int32 or the int64 dtype produced by ``torch.topk``; they are -converted to int32 by the Triton kernel before packing. -""" - -import torch -import triton -import triton.language as tl - -from sglang.kernels.jit.utils import is_arch_support_pdl - - -class PackTopkIds: - @classmethod - def execute( - cls, topk_ids: torch.Tensor, topk_weights: torch.Tensor - ) -> torch.Tensor: - return cls.triton(topk_ids, topk_weights) - - @classmethod - def vanilla( - cls, topk_ids: torch.Tensor, topk_weights: torch.Tensor - ) -> torch.Tensor: - weight_bits = ( - topk_weights.to(torch.bfloat16).view(torch.int16).to(torch.int32) & 0xFFFF - ) - return (topk_ids.to(torch.int32) << 16) | weight_bits - - @classmethod - def triton(cls, topk_ids: torch.Tensor, topk_weights: torch.Tensor) -> torch.Tensor: - assert topk_ids.shape == topk_weights.shape, ( - f"shape mismatch: {topk_ids.shape=} vs {topk_weights.shape=}" - ) - assert topk_ids.ndim >= 1, f"expected >=1D, got {topk_ids.shape=}" - - assert topk_ids.dtype in ( - torch.int32, - torch.int64, - ), f"topk_ids must be int32 or int64, got {topk_ids.dtype}" - assert topk_weights.dtype == torch.float32, ( - f"topk_weights must be float32, got {topk_weights.dtype}" - ) - - assert topk_ids.is_contiguous(), "topk_ids must be contiguous" - assert topk_weights.is_contiguous(), "topk_weights must be contiguous" - - out = torch.empty_like(topk_ids, dtype=torch.int32) - numel = out.numel() - if numel == 0: - return out - - BLOCK_SIZE = 1024 - grid = (triton.cdiv(numel, BLOCK_SIZE),) - pdl_kwargs = ( - {"USE_PDL": True, "launch_pdl": True} if is_arch_support_pdl() else {} - ) - _pack_topk_ids_triton_kernel[grid]( - topk_ids, - topk_weights, - out, - numel, - BLOCK_SIZE=BLOCK_SIZE, - **pdl_kwargs, - ) - return out - - -@triton.jit -def _pack_topk_ids_triton_kernel( - topk_ids_ptr, - topk_weights_ptr, - out_ptr, - numel, - BLOCK_SIZE: tl.constexpr, - USE_PDL: tl.constexpr = False, -): - pid = tl.program_id(0) - offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) - mask = offsets < numel - - if USE_PDL: - tl.extra.cuda.gdc_wait() - - ids = tl.load(topk_ids_ptr + offsets, mask=mask, other=0) - w = tl.load(topk_weights_ptr + offsets, mask=mask, other=0.0) - - if USE_PDL: - tl.extra.cuda.gdc_launch_dependents() - - w_bf16 = w.to(tl.bfloat16) - w_i16 = w_bf16.to(tl.int16, bitcast=True) - w_i32 = w_i16.to(tl.int32) & 0xFFFF - - ids_i32 = ids.to(tl.int32) - packed = (ids_i32 << 16) | w_i32 - - tl.store(out_ptr + offsets, packed, mask=mask) diff --git a/python/sglang/srt/layers/moe/flashinfer_trtllm_moe.py b/python/sglang/srt/layers/moe/flashinfer_trtllm_moe.py index c786aeeb3..d7f50b173 100644 --- a/python/sglang/srt/layers/moe/flashinfer_trtllm_moe.py +++ b/python/sglang/srt/layers/moe/flashinfer_trtllm_moe.py @@ -120,6 +120,7 @@ def trtllm_fp8_block_scale_moe_out_wrapper( def _fake_fp8_block_scale_routed_moe_out( topk_ids: torch.Tensor, + topk_weights: Optional[torch.Tensor], routing_bias: Optional[torch.Tensor], hidden_states: torch.Tensor, hidden_states_scale: torch.Tensor, @@ -156,6 +157,7 @@ def _fake_fp8_block_scale_routed_moe_out( ) def trtllm_fp8_block_scale_routed_moe_out_wrapper( topk_ids: torch.Tensor, + topk_weights: Optional[torch.Tensor], routing_bias: Optional[torch.Tensor], hidden_states: torch.Tensor, hidden_states_scale: torch.Tensor, @@ -192,7 +194,7 @@ def trtllm_fp8_block_scale_routed_moe_out_wrapper( ) from e kwargs = { - "topk_ids": topk_ids, + "topk_ids": topk_ids if topk_weights is None else (topk_ids, topk_weights), "routing_bias": routing_bias, "hidden_states": hidden_states, "hidden_states_scale": hidden_states_scale, diff --git a/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py b/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py index d79a1e105..5b3311445 100644 --- a/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py +++ b/python/sglang/srt/layers/moe/moe_runner/flashinfer_trtllm.py @@ -9,7 +9,6 @@ import torch from torch.nn import Module from torch.nn.parameter import Parameter -from sglang.kernels.ops.moe.pack_topk_ids import PackTopkIds from sglang.kernels.ops.quantization.fp8_kernel import ( per_token_group_quant_fp8, scaled_fp8_quant, @@ -168,20 +167,29 @@ def _is_gated(layer: Module) -> bool: return True if is_gated is None else is_gated -def _get_packed_topk_ids_for_flashinfer_routed(topk_output) -> torch.Tensor: - """Return FlashInfer routed packed top-k ids, using prepacked output if present.""" - packed_topk_ids = getattr(topk_output, "packed_topk_ids", None) - if packed_topk_ids is not None: - return packed_topk_ids +FlashInferRouting = torch.Tensor | tuple[torch.Tensor, torch.Tensor] + +def _get_routing_for_flashinfer_routed(topk_output) -> FlashInferRouting: + """Return the `topk_ids` kernel argument for the trtllm routed MoEs.""" from sglang.srt.layers.moe.topk import TopKOutputChecker + if TopKOutputChecker.format_is_packed(topk_output): + return topk_output.packed_topk_ids + assert TopKOutputChecker.format_is_standard(topk_output) - return PackTopkIds.execute( - topk_output.topk_ids.contiguous(), topk_output.topk_weights.contiguous() + return ( + topk_output.topk_ids.contiguous(), + topk_output.topk_weights.contiguous(), ) +def _routing_top_k(routing: FlashInferRouting) -> int: + """Both routing forms carry top_k as the last dim of their ids tensor.""" + ids = routing[0] if isinstance(routing, tuple) else routing + return ids.shape[1] + + def _align_fp8_moe_weights( w13: torch.Tensor, w2: torch.Tensor, @@ -802,10 +810,14 @@ def fused_experts_none_to_flashinfer_trtllm_fp8( assert runner_config.top_k is not None, ( "runner_config.top_k is required for flashinfer_trtllm_routed." ) - packed_topk_ids = _get_packed_topk_ids_for_flashinfer_routed(topk_output) + routing = _get_routing_for_flashinfer_routed(topk_output) + topk_ids, topk_weights = ( + routing if isinstance(routing, tuple) else (routing, None) + ) trtllm_fp8_block_scale_routed_moe_out_wrapper( - topk_ids=packed_topk_ids, + topk_ids=topk_ids, + topk_weights=topk_weights, routing_bias=None, hidden_states=a_q, hidden_states_scale=a_sf_t, @@ -1063,16 +1075,16 @@ def _fused_experts_flashinfer_mxfp4_sm100_trtllm_gen( from flashinfer.tllm_enums import ActivationType, RoutingMethodType if is_standard: - if prepared_packed_topk is not None: - packed_topk = prepared_packed_topk - else: - packed_topk = PackTopkIds.execute( - topk_output.topk_ids, topk_output.topk_weights - ) + routing = ( + prepared_packed_topk + if prepared_packed_topk is not None + else _get_routing_for_flashinfer_routed(topk_output) + ) + routed_top_k = _routing_top_k(routing) defer_finalize = _deferred_finalize_enabled.get() result = trtllm_fp4_block_scale_routed_moe( - topk_ids=packed_topk, + topk_ids=routing, routing_bias=None, hidden_states=x_quant, hidden_states_scale=x_scale, @@ -1089,7 +1101,7 @@ def _fused_experts_flashinfer_mxfp4_sm100_trtllm_gen( output1_scale_gate_scalar=None, output2_scale_scalar=None, num_experts=quant_info.global_num_experts, - top_k=packed_topk.shape[1], + top_k=routed_top_k, n_group=None, topk_group=None, intermediate_size=quant_info.intermediate_size_per_partition, @@ -1109,7 +1121,7 @@ def _fused_experts_flashinfer_mxfp4_sm100_trtllm_gen( gemm2_out=gemm2_out, expert_weights=topk_weights, expanded_idx_to_permuted_idx=expanded_idx, - top_k=packed_topk.shape[1], + top_k=routed_top_k, ) return StandardCombineInput(hidden_states=result) # The finalized kernel writes to its explicit output argument. Do @@ -1375,9 +1387,9 @@ def fused_experts_none_to_flashinfer_trtllm_fp4( ) if use_routed_topk: - packed_topk_ids = _get_packed_topk_ids_for_flashinfer_routed(topk_output) + routing = _get_routing_for_flashinfer_routed(topk_output) result = trtllm_fp4_block_scale_routed_moe( - topk_ids=packed_topk_ids, + topk_ids=routing, routing_bias=None, hidden_states=hs_fp4, hidden_states_scale=hs_scale, @@ -1395,7 +1407,7 @@ def fused_experts_none_to_flashinfer_trtllm_fp4( output2_scale_scalar=quant_info.g2_alphas, per_token_scale=per_token_scale, num_experts=quant_info.global_num_experts, - top_k=packed_topk_ids.shape[1], + top_k=_routing_top_k(routing), n_group=0, topk_group=0, intermediate_size=quant_info.intermediate_size_per_partition, @@ -1543,9 +1555,8 @@ def fused_experts_none_to_flashinfer_trtllm_bf16( elif routing_method_type == RoutingMethodType.DeepSeekV3: routing_method_type = RoutingMethodType.TopK - packed_topk_ids = _get_packed_topk_ids_for_flashinfer_routed(topk_output) final_hidden_states = trtllm_bf16_routed_moe( - topk_ids=packed_topk_ids, + topk_ids=_get_routing_for_flashinfer_routed(topk_output), hidden_states=hidden_states, gemm1_weights=quant_info.gemm1_weights, gemm2_weights=quant_info.gemm2_weights, diff --git a/test/registered/kernels/ops/test_kimi_k3_prerequisite_ops.py b/test/registered/kernels/ops/test_kimi_k3_prerequisite_ops.py index e4b6088f3..6f47dc3d6 100644 --- a/test/registered/kernels/ops/test_kimi_k3_prerequisite_ops.py +++ b/test/registered/kernels/ops/test_kimi_k3_prerequisite_ops.py @@ -33,7 +33,6 @@ from sglang.kernels.ops.mm.process.image import ( from sglang.kernels.ops.moe import moe_route_quant_fused from sglang.kernels.ops.moe.moe_route_radix import route_radix from sglang.kernels.ops.moe.moe_topk_sum import moe_topk_sum -from sglang.kernels.ops.moe.pack_topk_ids import PackTopkIds from sglang.kernels.ops.quantization.per_token_group_quant import ( per_token_group_quant, ) @@ -53,6 +52,18 @@ MLA_DIM = NOPE_DIM + ROPE_DIM MLA_PAGES = 256 +def _pack_topk_oracle(topk_ids, topk_weights): + """Pure-torch reference for the packed ids the fused route+quant kernel emits. + + FlashInfer's routed MoE reads one int32 per entry: the expert id in the high + half and the bf16 weight bits in the low half. + """ + weight_bits = ( + topk_weights.to(torch.bfloat16).view(torch.int16).to(torch.int32) & 0xFFFF + ) + return (topk_ids.to(torch.int32) << 16) | weight_bits + + def _route_oracle( scores, bias, topk, renormalize, routed_scaling_factor, apply_scale, sorted ): @@ -311,7 +322,7 @@ class TestKimiK3PrerequisiteOps(CustomTestCase): self.skipTest("fused route+quant kernel unavailable") hidden = torch.randn(8, 3584, device="cuda", dtype=torch.bfloat16) ref_weights, ref_ids = route_radix(*args, sorted=False) - ref_packed = PackTopkIds.execute(ref_ids, ref_weights) + ref_packed = _pack_topk_oracle(ref_ids, ref_weights) ref_q, ref_scale = per_token_group_quant( hidden, group_size=32, scale_ue8m0=True )