From 666fc3a35459369f9bf23463d99a3f4fa4ecd4d6 Mon Sep 17 00:00:00 2001 From: Xin Guan <294044352+xguannv@users.noreply.github.com> Date: Thu, 17 Sep 2026 16:55:35 +0800 Subject: [PATCH 1/5] [None][feat] Add SiTU support to Rubin FC12 MoE backend Signed-off-by: Xin Guan <294044352+xguannv@users.noreply.github.com> --- .../_torch/custom_ops/cute_dsl_custom_ops.py | 57 ++- ...ous_grouped_blockscaled_gemm_fused_fc12.py | 340 ++++++++++++------ .../moe/fused_moe/fused_moe_cute_dsl_fc12.py | 33 +- .../test_lists/test-db/l0_b200.yml | 2 +- .../_torch/moe/test_kimi_k3_situ_moe.py | 76 +++- tests/unittest/_torch/moe/test_moe_backend.py | 88 ++++- tests/unittest/_torch/moe/test_moe_impl.py | 18 + 7 files changed, 472 insertions(+), 142 deletions(-) diff --git a/tensorrt_llm/_torch/custom_ops/cute_dsl_custom_ops.py b/tensorrt_llm/_torch/custom_ops/cute_dsl_custom_ops.py index 5459137e9b76..5716b07700a1 100644 --- a/tensorrt_llm/_torch/custom_ops/cute_dsl_custom_ops.py +++ b/tensorrt_llm/_torch/custom_ops/cute_dsl_custom_ops.py @@ -17773,7 +17773,7 @@ def _( # builds that do not ship cutlass.utils.rubin_helpers (SM107 only). if IS_CUTLASS_DSL_RUBIN_AVAILABLE: # ---------------------------------------------------------------- - # Rubin NVFP4 Fused FC12 (FC1 gather+SwiGLU + FC2 finalize in ONE kernel) + # Rubin NVFP4 Fused FC12 (FC1 gather+gated-act + FC2 finalize in ONE kernel) # ---------------------------------------------------------------- # Compat shim: the delivered fused kernel accesses ``cutlass.memory.*`` # (SmemAllocator / TmemAllocator / get_smem_capacity_in_bytes) and @@ -17866,16 +17866,20 @@ class Sm107BlockScaledContiguousGroupedGemmFusedFc12Runner( kernel_cache = dict() tuning_config_cache = dict() - def __init__(self, - num_experts: int, - top_k: int, - num_local_experts: int, - local_expert_offset: int, - tile_size: int, - scaling_vector_size: int = 16, - swiglu_limit: float = float("inf"), - ep_size: int = 1, - enable_alltoall: bool = False): + def __init__( + self, + num_experts: int, + top_k: int, + num_local_experts: int, + local_expert_offset: int, + tile_size: int, + scaling_vector_size: int = 16, + swiglu_limit: float = float("inf"), + ep_size: int = 1, + enable_alltoall: bool = False, + activation_type: ActivationType = ActivationType.Swiglu, + situ_beta: Optional[float] = None, + situ_linear_beta: Optional[float] = None): super().__init__() self.num_experts = num_experts self.top_k = top_k @@ -17884,6 +17888,9 @@ def __init__(self, self.tile_size = tile_size self.scaling_vector_size = scaling_vector_size self.swiglu_limit = swiglu_limit + self.activation_type = ActivationType(activation_type) + self.situ_beta = situ_beta + self.situ_linear_beta = situ_linear_beta # Used only by the in-op output memset (moved here so the memset # is the fused kernel's immediate stream predecessor). self.ep_size = ep_size @@ -17908,6 +17915,9 @@ def unique_id(self): self.tile_size, self.scaling_vector_size, self.swiglu_limit, + int(self.activation_type), + self.situ_beta, + self.situ_linear_beta, ) def get_valid_tactics( @@ -18205,7 +18215,8 @@ def forward(self, inputs: List[torch.Tensor], cache_key = (self.scaling_vector_size, self.tile_size, self.top_k, mma_tiler, mma_inst_shape, cluster_shape_mn, max_active_clusters, - self.swiglu_limit) + self.swiglu_limit, int(self.activation_type), + self.situ_beta, self.situ_linear_beta) if cache_key not in self.__class__.kernel_cache: gemm = self.__class__.kernel_class( self.scaling_vector_size, @@ -18217,6 +18228,9 @@ def forward(self, inputs: List[torch.Tensor], use_pdl=True, swiglu_limit=self.swiglu_limit, scheduler="l2_atomic", + activation_type=self.activation_type, + situ_beta=self.situ_beta, + situ_linear_beta=self.situ_linear_beta, ) @cute.jit @@ -18452,6 +18466,9 @@ def _run_nvfp4_fc12_fused_rubin( ep_size: int, enable_alltoall: bool, tuner_key: str, + activation_type: ActivationType = ActivationType.Swiglu, + situ_beta: Optional[float] = None, + situ_linear_beta: Optional[float] = None, ) -> torch.Tensor: tuner = AutoTuner.get() runner = Sm107BlockScaledContiguousGroupedGemmFusedFc12Runner( @@ -18464,6 +18481,9 @@ def _run_nvfp4_fc12_fused_rubin( swiglu_limit=swiglu_limit, ep_size=ep_size, enable_alltoall=enable_alltoall, + activation_type=activation_type, + situ_beta=situ_beta, + situ_linear_beta=situ_linear_beta, ) # Input order matches Fc12FusedInputsHelper (FC1 prefix 0..9 mirrors # GatherGroupedGemmInputsHelper; FC2 tensors 10..14; expanded_idx 15 @@ -18500,7 +18520,9 @@ def _run_nvfp4_fc12_fused_rubin( "SymInt local_expert_offset, SymInt tile_size, float swiglu_limit, " "SymInt ep_size, bool enable_alltoall, " "SymInt scaling_vector_size=16, " - "str? precomputed_tactic=None) -> ()", + "str? precomputed_tactic=None, " + f"SymInt activation_type={int(ActivationType.Swiglu)}, " + "float? situ_beta=None, float? situ_linear_beta=None) -> ()", device_types="cuda") def cute_dsl_nvfp4_fc12_fused_rubin( input: torch.Tensor, @@ -18529,6 +18551,9 @@ def cute_dsl_nvfp4_fc12_fused_rubin( enable_alltoall: bool, scaling_vector_size: int = 16, precomputed_tactic: Optional[str] = None, + activation_type: int = int(ActivationType.Swiglu), + situ_beta: Optional[float] = None, + situ_linear_beta: Optional[float] = None, ) -> None: # In-place: finalize scatter-adds into ``output`` (mutates_args); # the op returns nothing so it does not alias its own input. @@ -18541,7 +18566,8 @@ def cute_dsl_nvfp4_fc12_fused_rubin( local_expert_offset, tile_size, scaling_vector_size, swiglu_limit, precomputed_tactic, expanded_idx_to_permuted_idx, ep_size, enable_alltoall, - "trtllm::cute_dsl_nvfp4_fc12_fused_rubin") + "trtllm::cute_dsl_nvfp4_fc12_fused_rubin", + ActivationType(activation_type), situ_beta, situ_linear_beta) @torch.library.register_fake("trtllm::cute_dsl_nvfp4_fc12_fused_rubin") def _( @@ -18571,5 +18597,8 @@ def _( enable_alltoall: bool, scaling_vector_size: int = 16, precomputed_tactic: Optional[str] = None, + activation_type: int = int(ActivationType.Swiglu), + situ_beta: Optional[float] = None, + situ_linear_beta: Optional[float] = None, ) -> None: return None diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/rubin/moe/rubin_contiguous_grouped_blockscaled_gemm_fused_fc12.py b/tensorrt_llm/_torch/cute_dsl_kernels/rubin/moe/rubin_contiguous_grouped_blockscaled_gemm_fused_fc12.py index 75e5bd11b463..37e7f3de4537 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/rubin/moe/rubin_contiguous_grouped_blockscaled_gemm_fused_fc12.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/rubin/moe/rubin_contiguous_grouped_blockscaled_gemm_fused_fc12.py @@ -29,6 +29,8 @@ from dataclasses import dataclass from typing import NamedTuple, Optional, Tuple, Type, Union +from ....utils import ActivationType + import cuda.bindings.driver as cuda import cutlass import cutlass.cute as cute @@ -790,6 +792,22 @@ def silu_f32( return a * sigmoid_f32(a, fastmath=fastmath) +_FC12_ACTIVATION_TYPES = ( + ActivationType.Swiglu, + ActivationType.SiTu, +) + + +def validate_fc12_activation_type(activation_type) -> ActivationType: + """Fused FC12 is gated-only: SwiGLU or SiTU. Relu2 would keep full N.""" + activation_type = ActivationType(int(activation_type)) + if activation_type not in _FC12_ACTIVATION_TYPES: + raise ValueError( + f"Fused FC12 supports SwiGLU and SiTU only, got {activation_type}" + ) + return activation_type + + class S2TCopyBundle(NamedTuple): """Bundle of tiled copy and partitioned tensors for smem-to-tmem copies.""" @@ -841,12 +859,12 @@ def sm100_tcgen05_st_32x32b_x4( Compute: fc1_acc = fc1_alpha * (SFA * A[token_ids]) * (SFB * B) - C = up * silu(gate) # SwiGLU on interleaved acc + C = gated act on interleaved acc (SwiGLU or SiTU; both halve N) fc2_acc = fc2_alpha * (SFC * C) * (FC2_SFB * FC2_B) + optional NVFP4 quantization (generates SFC) when c_dtype == Float4E2M1FN. Shapes: A is M×K×1; B is N×K×L (L = num experts), interleaved [up, gate] at -granularity=64; C is M×(N/2)×1 (N halved by SwiGLU). SFA/SFB layouts follow +granularity=64; C is M×(N/2)×1 (N halved by the gated epilogue). SFA/SFB layouts follow BlockScaledBasicChunk. ``permuted_idx_to_expanded_idx`` is shared by both phases: FC1 divides by ``topk`` to recover the receive row, while FC2 also uses the remainder to select the route scale. @@ -882,8 +900,9 @@ class Sm107BlockScaledContiguousGroupedGemmFusedFc12Kernel: - A/SFA load: four gather warps issue CpAsync128.CG into a shared fc1_a_pipeline. SFA is consumed directly by a Cp128x128b UTCCP in the MMA warp using the 128dp_Unique layout. - - SwiGLU epilogue: C = up * silu(gate), where up/gate come from - interleaved accumulator at granularity=64 → output N is halved. + - Gated epilogue: SwiGLU (``up * silu(gate)``) or SiTU (Kimi K3); + up/gate come from interleaved accumulator at granularity=64 → + output N is halved. Selected by ``activation_type`` at trace time. - Optional NVFP4 quant: when c_dtype == Float4E2M1FN, the epilogue also generates SFC and quantizes the output. @@ -902,6 +921,11 @@ class Sm107BlockScaledContiguousGroupedGemmFusedFc12Kernel: :param vectorized_f32: Use vectorized f32x2 ops in epilogue. :param topk: Experts selected per token. :param swiglu_limit: DS-v4 SwiGLU clamp limit; ``+inf`` disables clamp. + Mutually exclusive with ``ActivationType.SiTu``. + :param activation_type: ``Swiglu`` (default) or ``SiTu``. Folded at trace + time; a different value compiles a different kernel. + :param situ_beta: Gate-side SiTU constant; required for ``SiTu`` only. + :param situ_linear_beta: Linear-side SiTU constant; required for ``SiTu`` only. :param scheduler: Persistent work-ID scheduler: ``static`` or ``l2_atomic``. """ @@ -916,6 +940,9 @@ def __init__( use_pdl: bool = True, swiglu_limit: cutlass.Float32 = float("inf"), scheduler: str = "static", + activation_type: ActivationType = ActivationType.Swiglu, + situ_beta: Optional[float] = None, + situ_linear_beta: Optional[float] = None, ): self.sf_vec_size = sf_vec_size self.topk = topk @@ -937,6 +964,31 @@ def __init__( swiglu_limit = float("inf") self.swiglu_limit = swiglu_limit self.has_swiglu_limit = swiglu_limit != float("inf") + self.activation_type = validate_fc12_activation_type(activation_type) + if self.activation_type == ActivationType.SiTu: + if situ_beta is None or situ_linear_beta is None: + raise ValueError( + "ActivationType.SiTu requires both situ_beta and " + f"situ_linear_beta, got {situ_beta} and {situ_linear_beta}." + ) + if situ_beta <= 0 or situ_linear_beta <= 0: + raise ValueError( + f"SiTU betas must be positive, got {situ_beta} and {situ_linear_beta}." + ) + if self.has_swiglu_limit: + raise ValueError( + "Fused FC12 SiTU does not support a SwiGLU clamp; " + "drop swiglu_limit for SiTU checkpoints." + ) + elif situ_beta is not None or situ_linear_beta is not None: + raise ValueError( + "situ_beta / situ_linear_beta require " + f"ActivationType.SiTu, got {self.activation_type.name}." + ) + self.situ_beta = None if situ_beta is None else float(situ_beta) + self.situ_linear_beta = ( + None if situ_linear_beta is None else float(situ_linear_beta) + ) self.use_2cta_instrs = mma_inst_shape[0] == 256 self.cta_group = ( @@ -1542,7 +1594,7 @@ def __call__( This method performs FC1 layer computation: 1. GEMM: acc = fc1_alpha * (SFA * A[token_ids]) * (SFB * B) - 2. SwiGLU: C = up * silu(gate), where up/gate are extracted from interleaved acc (granularity=64) + 2. Gated act: SwiGLU or SiTU on interleaved acc (granularity=64; N/2) 3. Optional Quant: When c_dtype is Float4E2M1FN, generates SFC and quantizes output Data loading: @@ -1562,7 +1614,7 @@ def __call__( shared memory when use_2cta_instrs is True - TMA warp: Load B and SFB with multicast - MMA warp: Perform matrix multiply-accumulate - - Epilogue warps: Apply SwiGLU activation, optional quantization, and store results + - Epilogue warps: Apply gated activation, optional quantization, and store results :param fc1_a: FC1 input A (MxKx1), gathered with the shared route mapping :type fc1_a: cute.Tensor @@ -2199,6 +2251,173 @@ class SharedStorageCpasync2cta: ) return + @cute.jit + def _apply_swiglu_epilogue( + self, + acc_vec_up: cute.Tensor, + acc_vec_gate: cute.Tensor, + alpha_val, + tCompute: cute.Tensor, + ): + """SwiGLU: ``tCompute[i] = (alpha * up[i]) * silu(alpha * gate[i])``. + + Optional DS-v4 clamp via ``has_swiglu_limit`` is folded at trace time. + """ + if cutlass.const_expr(self.vectorized_f32): + LOG2_E = cutlass.Float32(1.4426950408889634) + for i in cutlass.range_constexpr(0, cute.size(acc_vec_up.shape), 2): + acc_vec_up_alpha = cute.arch.mul_packed_f32x2( + (acc_vec_up[i], acc_vec_up[i + 1]), + (cutlass.Float32(alpha_val), cutlass.Float32(alpha_val)), + ) + acc_vec_gate_alpha = cute.arch.mul_packed_f32x2( + (acc_vec_gate[i], acc_vec_gate[i + 1]), + (cutlass.Float32(alpha_val), cutlass.Float32(alpha_val)), + ) + if cutlass.const_expr(self.has_swiglu_limit): + acc_vec_gate_alpha = ( + fmin(acc_vec_gate_alpha[0], self.swiglu_limit), + fmin(acc_vec_gate_alpha[1], self.swiglu_limit), + ) + acc_vec_up_alpha = ( + fclip_xorsign(acc_vec_up_alpha[0], self.swiglu_limit), + fclip_xorsign(acc_vec_up_alpha[1], self.swiglu_limit), + ) + tCompute_log2e = cute.arch.mul_packed_f32x2( + (acc_vec_gate_alpha[0], acc_vec_gate_alpha[1]), + (-LOG2_E, -LOG2_E), + ) + ( + tCompute[i], + tCompute[i + 1], + ) = cute.arch.add_packed_f32x2( + ( + cute.math.exp2(tCompute_log2e[0], fastmath=True), + cute.math.exp2(tCompute_log2e[1], fastmath=True), + ), + (1.0, 1.0), + ) + tCompute[i] = cute.arch.rcp_approx(tCompute[i]) + tCompute[i + 1] = cute.arch.rcp_approx(tCompute[i + 1]) + ( + tCompute[i], + tCompute[i + 1], + ) = cute.arch.mul_packed_f32x2( + (tCompute[i], tCompute[i + 1]), + (acc_vec_gate_alpha[0], acc_vec_gate_alpha[1]), + ) + ( + tCompute[i], + tCompute[i + 1], + ) = cute.arch.mul_packed_f32x2( + (tCompute[i], tCompute[i + 1]), + (acc_vec_up_alpha[0], acc_vec_up_alpha[1]), + ) + else: + for i in cutlass.range_constexpr(cute.size(acc_vec_up.shape)): + acc_vec_up_alpha = acc_vec_up[i] * cutlass.Float32(alpha_val) + acc_vec_gate_alpha = acc_vec_gate[i] * cutlass.Float32(alpha_val) + if cutlass.const_expr(self.has_swiglu_limit): + acc_vec_gate_alpha = fmin(acc_vec_gate_alpha, self.swiglu_limit) + acc_vec_up_alpha = fclip_xorsign( + acc_vec_up_alpha, self.swiglu_limit + ) + tCompute[i] = acc_vec_up_alpha * silu_f32( + acc_vec_gate_alpha, fastmath=True + ) + + @cute.jit + def _apply_situ_epilogue( + self, + acc_vec_up: cute.Tensor, + acc_vec_gate: cute.Tensor, + alpha_val, + tCompute: cute.Tensor, + ): + """SiTU (Kimi K3), matching ``SituAndMul`` / the two-op Rubin kernel:: + + g = alpha * gate, u = alpha * up + situ_gate = beta * tanh(g / beta) * sigmoid(g) + situ_up = linear_beta * tanh(u / linear_beta) + tCompute = situ_gate * situ_up + + Packed path uses ``tanh(z) = 2 * sigmoid(2z) - 1``. Betas fold at + trace time; a different pair compiles a different kernel. + """ + beta = self.situ_beta + linear_beta = self.situ_linear_beta + if cutlass.const_expr(self.vectorized_f32): + LOG2_E = cutlass.Float32(1.4426950408889634) + neg_log2e_pair = (-LOG2_E, -LOG2_E) + one_pair = (cutlass.Float32(1.0), cutlass.Float32(1.0)) + + inv_2beta = cutlass.Float32(2.0 / beta) + two_beta = cutlass.Float32(2.0 * beta) + neg_beta = cutlass.Float32(-beta) + inv_2lbeta = cutlass.Float32(2.0 / linear_beta) + two_lbeta = cutlass.Float32(2.0 * linear_beta) + neg_lbeta = cutlass.Float32(-linear_beta) + + def _sigmoid(p0, p1): + neg = cute.arch.mul_packed_f32x2((p0, p1), neg_log2e_pair) + e = ( + cute.math.exp2(neg[0], fastmath=True), + cute.math.exp2(neg[1], fastmath=True), + ) + d = cute.arch.add_packed_f32x2(e, one_pair) + return (cute.arch.rcp_approx(d[0]), cute.arch.rcp_approx(d[1])) + + alpha_pair = (cutlass.Float32(alpha_val), cutlass.Float32(alpha_val)) + for i in cutlass.range_constexpr(0, cute.size(acc_vec_up.shape), 2): + g = cute.arch.mul_packed_f32x2( + (acc_vec_gate[i], acc_vec_gate[i + 1]), alpha_pair + ) + u = cute.arch.mul_packed_f32x2( + (acc_vec_up[i], acc_vec_up[i + 1]), alpha_pair + ) + + sigmoid_g = _sigmoid(g[0], g[1]) + + gs = _sigmoid( + *cute.arch.mul_packed_f32x2(g, (inv_2beta, inv_2beta)) + ) + tanh_g = cute.arch.add_packed_f32x2( + cute.arch.mul_packed_f32x2(gs, (two_beta, two_beta)), + (neg_beta, neg_beta), + ) + + us = _sigmoid( + *cute.arch.mul_packed_f32x2(u, (inv_2lbeta, inv_2lbeta)) + ) + tanh_u = cute.arch.add_packed_f32x2( + cute.arch.mul_packed_f32x2(us, (two_lbeta, two_lbeta)), + (neg_lbeta, neg_lbeta), + ) + + situ_gate = cute.arch.mul_packed_f32x2(tanh_g, sigmoid_g) + out_pair = cute.arch.mul_packed_f32x2(situ_gate, tanh_u) + tCompute[i] = out_pair[0] + tCompute[i + 1] = out_pair[1] + else: + inv_2beta = cutlass.Float32(2.0 / beta) + two_beta = cutlass.Float32(2.0 * beta) + beta_f32 = cutlass.Float32(beta) + inv_2lbeta = cutlass.Float32(2.0 / linear_beta) + two_lbeta = cutlass.Float32(2.0 * linear_beta) + lbeta_f32 = cutlass.Float32(linear_beta) + for i in cutlass.range_constexpr(cute.size(acc_vec_up.shape)): + g = acc_vec_gate[i] * cutlass.Float32(alpha_val) + u = acc_vec_up[i] * cutlass.Float32(alpha_val) + tanh_g = ( + two_beta * sigmoid_f32(g * inv_2beta, fastmath=True) + - beta_f32 + ) + tanh_u = ( + two_lbeta * sigmoid_f32(u * inv_2lbeta, fastmath=True) + - lbeta_f32 + ) + tCompute[i] = (tanh_g * sigmoid_f32(g, fastmath=True)) * tanh_u + # GPU device kernel @cute.kernel def kernel( @@ -4313,9 +4532,9 @@ def kernel( # acc_pipeline.consumer_wait(acc_consumer_state) - # SwiGLU epilogue. Acc has full N cols with interleaved + # Gated epilogue. Acc has full N cols with interleaved # [up, gate] at granularity=64; C has N/2 cols. Iterate M and - # N output subtiles → up * silu(gate). + # N output subtiles → SwiGLU or SiTU. # tFC1TR_tAcc: (T2R, T2R_M, T2R_N, EPI_M, EPI_N, # STAGE), sliced on STAGE. bSG_gC: ((ATOM_V, REST_V), # EPI_M, EPI_N, loopM, loopN, loopL). @@ -4355,103 +4574,22 @@ def kernel( acc_vec_up = tFC1TR_rAcc_up.load() acc_vec_gate = tFC1TR_rAcc_gate.load() - # SwiGLU: output = up * silu(gate), silu(x) = x * sigmoid(x). + # Gated epilogue. Acc has full N cols with interleaved + # [up, gate] at granularity=64; C has N/2 cols. SwiGLU + # vs SiTU is a const_expr specialization (trace-time). tCompute = cute.make_rmem_tensor( acc_vec_gate.shape, self.acc_dtype ) - if cutlass.const_expr(self.vectorized_f32): - # SwiGLU Packed Version: uses f32x2 packed operations for better performance - # Computes: output = (alpha * up) * silu(alpha * gate) - # where silu(x) = x * sigmoid(x) = x / (1 + exp(-x)) - LOG2_E = cutlass.Float32(1.4426950408889634) - for i in cutlass.range_constexpr( - 0, cute.size(tFC1TR_rAcc_up), 2 - ): - acc_vec_up_alpha = cute.arch.mul_packed_f32x2( - (acc_vec_up[i], acc_vec_up[i + 1]), - ( - cutlass.Float32(alpha_val), - cutlass.Float32(alpha_val), - ), - ) - acc_vec_gate_alpha = cute.arch.mul_packed_f32x2( - (acc_vec_gate[i], acc_vec_gate[i + 1]), - ( - cutlass.Float32(alpha_val), - cutlass.Float32(alpha_val), - ), - ) - if cutlass.const_expr(self.has_swiglu_limit): - acc_vec_gate_alpha = ( - fmin( - acc_vec_gate_alpha[0], self.swiglu_limit - ), - fmin( - acc_vec_gate_alpha[1], self.swiglu_limit - ), - ) - acc_vec_up_alpha = ( - fclip_xorsign( - acc_vec_up_alpha[0], self.swiglu_limit - ), - fclip_xorsign( - acc_vec_up_alpha[1], self.swiglu_limit - ), - ) - tCompute_log2e = cute.arch.mul_packed_f32x2( - (acc_vec_gate_alpha[0], acc_vec_gate_alpha[1]), - (-LOG2_E, -LOG2_E), - ) - ( - tCompute[i], - tCompute[i + 1], - ) = cute.arch.add_packed_f32x2( - ( - cute.math.exp2( - tCompute_log2e[0], fastmath=True - ), - cute.math.exp2( - tCompute_log2e[1], fastmath=True - ), - ), - (1.0, 1.0), - ) - tCompute[i] = cute.arch.rcp_approx(tCompute[i]) - tCompute[i + 1] = cute.arch.rcp_approx(tCompute[i + 1]) - ( - tCompute[i], - tCompute[i + 1], - ) = cute.arch.mul_packed_f32x2( - (tCompute[i], tCompute[i + 1]), - (acc_vec_gate_alpha[0], acc_vec_gate_alpha[1]), - ) - ( - tCompute[i], - tCompute[i + 1], - ) = cute.arch.mul_packed_f32x2( - (tCompute[i], tCompute[i + 1]), - (acc_vec_up_alpha[0], acc_vec_up_alpha[1]), - ) + if cutlass.const_expr( + self.activation_type == ActivationType.SiTu + ): + self._apply_situ_epilogue( + acc_vec_up, acc_vec_gate, alpha_val, tCompute + ) else: - # SwiGLU Unpacked Version: scalar operations - # Computes: output = (alpha * up) * silu(alpha * gate) - for i in cutlass.range_constexpr(cute.size(tFC1TR_rAcc_up)): - acc_vec_up_alpha = acc_vec_up[i] * cutlass.Float32( - alpha_val - ) - acc_vec_gate_alpha = acc_vec_gate[i] * cutlass.Float32( - alpha_val - ) - if cutlass.const_expr(self.has_swiglu_limit): - acc_vec_gate_alpha = fmin( - acc_vec_gate_alpha, self.swiglu_limit - ) - acc_vec_up_alpha = fclip_xorsign( - acc_vec_up_alpha, self.swiglu_limit - ) - tCompute[i] = acc_vec_up_alpha * silu_f32( - acc_vec_gate_alpha, fastmath=True - ) + self._apply_swiglu_epilogue( + acc_vec_up, acc_vec_gate, alpha_val, tCompute + ) if cutlass.const_expr(self.fc1_generate_sfc): # Float4E2M1FN quantization: per-vector absmax → diff --git a/tensorrt_llm/_torch/moe/fused_moe/fused_moe_cute_dsl_fc12.py b/tensorrt_llm/_torch/moe/fused_moe/fused_moe_cute_dsl_fc12.py index 5969ad27d6c3..88d68beeaa8f 100644 --- a/tensorrt_llm/_torch/moe/fused_moe/fused_moe_cute_dsl_fc12.py +++ b/tensorrt_llm/_torch/moe/fused_moe/fused_moe_cute_dsl_fc12.py @@ -15,7 +15,7 @@ """TrtllmCutedslFusedFc12Nvfp4Impl: ``trtllm.cutedsl.fused_fc12.nvfp4``. FC1+FC2-fused CuteDSL NVFP4 MoE backend for Rubin (SM107). One persistent kernel -(``trtllm::cute_dsl_nvfp4_fc12_fused_rubin``) does gather + FC1 GEMM + SwiGLU + +(``trtllm::cute_dsl_nvfp4_fc12_fused_rubin``) does gather + FC1 GEMM + SwiGLU/SiTU + requant + FC2 GEMM + finalize (scatter-add into ``moe_output``), so the FC1->FC2 intermediate never round-trips through global memory. @@ -78,6 +78,21 @@ class CuteDslFc12FusedMoENvfp4Runner(CuteDslFusedMoENvfp4Runner): def _tile_sizes(): return [128, 256] + def forward( + self, + inputs: List[torch.Tensor], + tactic: Optional[int], + do_preparation: bool = False, + ) -> torch.Tensor: + if do_preparation: + # Prime the fused inner op before outer CUDA-graph profiling. + # The parent's preparation passes a two-op memset knob that + # FC12 does not accept: its memset is owned by the fused op. + for tile_size in self._tile_sizes(): + super().forward(inputs, tactic=tile_size) + return inputs[4] + return super().forward(inputs, tactic=tactic) + @register_moe_impl class TrtllmCutedslFusedFc12Nvfp4Impl(MoEImplBase): @@ -106,7 +121,7 @@ class TrtllmCutedslFusedFc12Nvfp4Impl(MoEImplBase): supports_apply_router_weight_on_input=True, ), input_requirement=MoEInputRequirement(routing_scales_dtype=torch.float32), - doc="Rubin (SM107) NVFP4 fused FC1 + SwiGLU + FC2 + finalize persistent " + doc="Rubin (SM107) NVFP4 fused FC1 + gated activation + FC2 + finalize persistent " "CuTe DSL grouped GEMM.", ) # Taken off the descriptor rather than restated: the scheduler reads these @@ -116,11 +131,11 @@ class TrtllmCutedslFusedFc12Nvfp4Impl(MoEImplBase): capabilities = descriptor.capabilities input_requirement = descriptor.input_requirement - # The fused FC12 kernel has no activation selector: FC1 weights are the - # gate/up pair and the epilogue is SwiGLU (+ clamp). Relu2 is not gated and - # is not supported, so the resolver turns such layers down here. + # Both activations share the gate/up geometry. SiTU softcaps are baked + # into the epilogue, so they must be uniform across the local experts. activation_support = MoEActivationSupport( - kinds=frozenset({ActivationType.Swiglu}), + kinds=frozenset({ActivationType.Swiglu, ActivationType.SiTu}), + alpha_beta=ActivationParamShape.UNIFORM_SCALAR, limit=ActivationParamShape.UNIFORM_SCALAR, limit_when_absent=float("inf"), ) @@ -380,6 +395,7 @@ def run_moe_nvfp4( local_expert_offset=weight_view.slot_start, enable_finalize_fusion=self.use_fused_finalize, enable_alltoall=enable_alltoall, + workload_identity=(int(self.activation_type), self.act_alpha, self.act_beta), ) inputs = [x, token_selected_experts, token_final_scales, x_sf, moe_output, weight_view] _, best_tactic = tuner.choose_one( @@ -436,7 +452,7 @@ def run_moe_nvfp4_impl( tile_tokens_dim=tile_size, ) - # One fused op: gather + FC1 GEMM + SwiGLU + requant + FC2 GEMM + + # One fused op: gather + FC1 GEMM + gated activation + requant + FC2 GEMM + # finalize (scatter-add into moe_output = a2a combine workspace). # fc1_alpha/fc2_alpha map 1:1 to the two-op path's per-expert global # scales (the kernel takes split alphas). The three atomic counters are @@ -471,6 +487,9 @@ def run_moe_nvfp4_impl( ep_size=self.mapping.moe_ep_size, enable_alltoall=enable_alltoall, scaling_vector_size=16, + activation_type=int(self.activation_type), + situ_beta=self.act_alpha, + situ_linear_beta=self.act_beta, ) return moe_output diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index 6d5edcb10f32..f4ced2cc0f6f 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -183,7 +183,7 @@ l0_b200: - unittest/_torch/moe/test_moe_backend.py::test_megamoe_streaming_reload_resets_slot_claims - unittest/_torch/moe/test_moe_backend.py::test_trtllm_gen_nvfp4_situ_selects_padded_quant_method - unittest/_torch/moe/test_moe_backend.py::test_trtllm_gen_nvfp4_situ_fc31_scale_c_drops_dequant_scale - - unittest/_torch/moe/test_moe_backend.py::test_megamoe_bakes_situ_softcaps_as_uniform_scalars + - unittest/_torch/moe/test_moe_backend.py::test_codegen_baked_situ_softcaps_are_uniform_scalars - unittest/_torch/moe/test_moe_backend.py::test_megamoe_plain_swiglu_carries_no_constants - unittest/_torch/moe/test_moe_backend.py::test_create_moe_forwards_situ_activation_as_one_carrier - unittest/_torch/moe/test_moe_backend.py::test_create_moe_backend_rejects_apply_router_weight_on_input_by_declaration diff --git a/tests/unittest/_torch/moe/test_kimi_k3_situ_moe.py b/tests/unittest/_torch/moe/test_kimi_k3_situ_moe.py index b53d8731792f..bdfe256eb797 100644 --- a/tests/unittest/_torch/moe/test_kimi_k3_situ_moe.py +++ b/tests/unittest/_torch/moe/test_kimi_k3_situ_moe.py @@ -1249,13 +1249,12 @@ def block_scales(*shape): def _make_nvfp4_moe(gate, num_experts=_TP_EXPERTS, moe_backend="CUTLASS"): - """NVFP4 + SiTU routed MoE on either FP4 backend. + """NVFP4 + SiTU routed MoE on a supported FP4 backend. - Both take the same ``SiTuActivation`` carrier; CUTLASS serves it as an - ``ActivationType`` its kernels branch on, TRTLLM-Gen with the fused - ``Bmm_E2m1_E2m1E2m1_..._siTuGlu_*`` FC1 cubins (group-16 block scales). - ``_make_routed_moe`` already mirrors both of KimiK3MoERuntime's branches, - so the backend is the only variable. + All take the same ``SiTuActivation`` carrier. CUTLASS branches on the + ``ActivationType`` in its kernels, TRTLLM-Gen selects the fused + ``Bmm_E2m1_E2m1E2m1_..._siTuGlu_*`` FC1 cubins, and FC12 specializes its + Rubin fused epilogue. The backend is the only variable. """ from tensorrt_llm.models.modeling_utils import QuantAlgo, QuantConfig @@ -1714,7 +1713,20 @@ def _swiglu_reference_moe(x, router_logits, routing_method, w1, w2, w3, alpha, b @nvfp4_moe_supported -@pytest.mark.parametrize("moe_backend", ["CUTLASS", "TRTLLM"]) +@pytest.mark.parametrize( + "moe_backend", + [ + "CUTLASS", + "TRTLLM", + pytest.param( + "CUTEDSL_FC12", + marks=pytest.mark.skipif( + not torch.cuda.is_available() or get_sm_version() != 107, + reason="FC12 requires Rubin (SM107)", + ), + ), + ], +) def test_nvfp4_kernel_actually_applies_situ(moe_backend): """Which activation does the QUANTIZED kernel actually run? @@ -1726,12 +1738,12 @@ def test_nvfp4_kernel_actually_applies_situ(moe_backend): wrong in exactly the way the GSM8K collapse showed, while every shape-and-buffer check stayed green. - Run for both FP4 backends: CUTLASS resolves SiTU through the activation - enum, TRTLLM-Gen through a distinct fused-cubin family - (``Bmm_E2m1_E2m1E2m1_..._siTuGlu_*``). A silent fallback is a different - failure on each, and this comparison is tolerance-free, so it catches - both -- including the degenerate all-zero FC1 output, which scores 0 - against both references and so fails the assertion below. + Run for all supported FP4 backends: CUTLASS resolves SiTU through the + activation enum, TRTLLM-Gen through a distinct fused-cubin family + (``Bmm_E2m1_E2m1E2m1_..._siTuGlu_*``), and FC12 through its Rubin fused + epilogue. A silent fallback is a different failure on each, and this + comparison is tolerance-free, so it also catches the degenerate all-zero + FC1 output, which scores 0 against both references and fails below. Reported rather than merely asserted: which reference the kernel is closer to is the diagnosis. @@ -1741,7 +1753,7 @@ def test_nvfp4_kernel_actually_applies_situ(moe_backend): torch.manual_seed(91) x = torch.randn(8, hidden, dtype=torch.bfloat16, device="cuda") * 0.5 - act_scale = float(x.abs().max().float() / (448 * 6)) + act_scale = 1.0 if moe_backend == "CUTEDSL_FC12" else float(x.abs().max().float() / (448 * 6)) w1 = [ torch.randn(inter, hidden, dtype=torch.bfloat16, device="cuda") * 0.05 for _ in range(num_experts) @@ -1757,10 +1769,33 @@ def test_nvfp4_kernel_actually_applies_situ(moe_backend): bank = [_quantize_expert_to_nvfp4(w1[e], w2[e], w3[e], act_scale) for e in range(num_experts)] moe = _make_nvfp4_moe(gate, num_experts=num_experts, moe_backend=moe_backend) + if moe_backend == "CUTEDSL_FC12": + from tensorrt_llm._torch.moe.fused_moe.fused_moe_cute_dsl_fc12 import ( + TrtllmCutedslFusedFc12Nvfp4Impl, + ) + + assert type(moe.backend) is TrtllmCutedslFusedFc12Nvfp4Impl _load_nvfp4_bank_for(moe, bank, moe_backend) router_logits = gate.compute_logits(x) actual = moe.forward(x, router_logits, all_rank_num_tokens=None).float() + assert torch.isfinite(actual).all() + + if moe_backend == "CUTEDSL_FC12": + # Warm up on a side stream before capture; compilation and tuning + # must finish before CUDA Graph records the fused op. + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(stream): + for _ in range(3): + moe.forward(x, router_logits, all_rank_num_tokens=None) + torch.cuda.current_stream().wait_stream(stream) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + captured = moe.forward(x, router_logits, all_rank_num_tokens=None) + for _ in range(2): + graph.replay() + torch.testing.assert_close(captured.float(), actual, rtol=0.02, atol=0.02) situ = _situ_reference_moe( x, router_logits, gate.routing_method, w1, w2, w3, beta=4.0, linear_beta=25.0 @@ -1788,6 +1823,19 @@ def score(ref): f"quantized path is not applying SiTU" ) + if moe_backend == "CUTEDSL_FC12": + # The unfixed FC12 kernel runs ordinary SwiGLU, not SwiGLU with the + # SiTU constants misinterpreted as sigmoid scale / linear bias. + plain_swiglu = _swiglu_reference_moe( + x, router_logits, gate.routing_method, w1, w2, w3, alpha=1.0, beta=0.0 + ) + plain_cos, plain_l2 = score(plain_swiglu) + distance_situ = torch.linalg.vector_norm(actual - situ) + distance_swiglu = torch.linalg.vector_norm(actual - plain_swiglu) + print(f"FC12 plain SwiGLU: cosine={plain_cos:.6f}, rel_l2={plain_l2:.6f}") + assert situ_l2 < 0.30 + assert distance_situ < distance_swiglu + def test_fp8_block_scaled_dequantization(): """FP8_PB_WO attention weights must be dequantized, never reinterpreted. diff --git a/tests/unittest/_torch/moe/test_moe_backend.py b/tests/unittest/_torch/moe/test_moe_backend.py index 45e49f9d137b..ac11de86c97d 100644 --- a/tests/unittest/_torch/moe/test_moe_backend.py +++ b/tests/unittest/_torch/moe/test_moe_backend.py @@ -69,6 +69,10 @@ CuteDslFusedMoENvfp4Runner, ) from tensorrt_llm._torch.moe.fused_moe.fused_moe_cute_dsl_b12x import CuteDslB12xFusedMoE +from tensorrt_llm._torch.moe.fused_moe.fused_moe_cute_dsl_fc12 import ( + CuteDslFc12FusedMoENvfp4Runner, + TrtllmCutedslFusedFc12Nvfp4Impl, +) from tensorrt_llm._torch.moe.fused_moe.fused_moe_cutlass import CutlassFusedMoE from tensorrt_llm._torch.moe.fused_moe.fused_moe_marlin import MarlinFusedMoE from tensorrt_llm._torch.moe.fused_moe.fused_moe_trtllm_gen import ( @@ -1279,14 +1283,15 @@ def test_megamoe_cutedsl_cache_derived_state_survives_the_read_only_reader_walk( wrapper.cache_derived_state.assert_not_called() -def test_megamoe_bakes_situ_softcaps_as_uniform_scalars(): - # MegaMoE declares UNIFORM_SCALAR for alpha/beta because the kernels bake - # them at codegen time, so a per-expert tensor is reduced here. +@pytest.mark.parametrize("impl", [MegaMoEDeepGemm, TrtllmCutedslFusedFc12Nvfp4Impl]) +def test_codegen_baked_situ_softcaps_are_uniform_scalars(impl): + # These kernels bake alpha/beta at codegen time, so the activation adapter + # accepts only one value shared by every local expert. params = materialize_activation_params( SiTuActivation(gate_softcap=torch.full((8,), 4.0), linear_softcap=25.0), - MegaMoEDeepGemm.activation_support, + impl.activation_support, num_local_experts=8, - owner="MegaMoEDeepGemm", + owner=impl.__name__, ) assert params.activation_type is ActivationType.SiTu @@ -1294,6 +1299,79 @@ def test_megamoe_bakes_situ_softcaps_as_uniform_scalars(): assert params.beta == 25.0 +def test_fc12_rejects_nonuniform_situ_softcaps() -> None: + with pytest.raises(ValueError, match="uniform"): + materialize_activation_params( + SiTuActivation(gate_softcap=torch.tensor([4.0, 5.0]), linear_softcap=25.0), + TrtllmCutedslFusedFc12Nvfp4Impl.activation_support, + num_local_experts=2, + owner="FC12", + ) + + +@pytest.mark.parametrize("clamp", [None, 7.0]) +def test_fc12_swiglu_keeps_clamp_without_situ_constants(clamp: Optional[float]) -> None: + params = materialize_activation_params( + SwigluActivation(clamp=clamp), + TrtllmCutedslFusedFc12Nvfp4Impl.activation_support, + num_local_experts=2, + owner="FC12", + ) + assert params.activation_type is ActivationType.Swiglu + assert (params.alpha, params.beta) == (None, None) + assert params.clamp == (float("inf") if clamp is None else clamp) + + +def test_fc12_outer_tuning_separates_activation_and_softcaps() -> None: + identities = [ + (int(ActivationType.Swiglu), None, None), + (int(ActivationType.SiTu), 4.0, 25.0), + (int(ActivationType.SiTu), 5.0, 25.0), + (int(ActivationType.SiTu), 4.0, 30.0), + ] + keys = [ + CuteDslFc12FusedMoENvfp4Runner( + forward_impl=MagicMock(), + num_experts=8, + top_k=2, + num_local_experts=8, + local_expert_offset=0, + workload_identity=identity, + ).unique_id() + for identity in identities + ] + assert len(set(keys)) == len(identities) + + +def test_fc12_preparation_primes_fused_tiles_without_two_op_memset_knob() -> None: + def fused_forward( + *inputs, + enable_alltoall, + tile_size, + recv_expert_count, + deep_ep_expert_capacity, + use_count_native_expert_metadata, + ): + assert not enable_alltoall + assert recv_expert_count is None and deep_ep_expert_capacity is None + assert not use_count_native_expert_metadata + tiles.append(tile_size) + return inputs[4] + + tiles = [] + runner = CuteDslFc12FusedMoENvfp4Runner( + forward_impl=fused_forward, + num_experts=8, + top_k=2, + num_local_experts=8, + local_expert_offset=0, + workload_identity=(int(ActivationType.SiTu), 4.0, 25.0), + ) + inputs = [object() for _ in range(6)] + assert runner.forward(inputs, tactic=None, do_preparation=True) is inputs[4] + assert tiles == [128, 256] + + def test_megamoe_plain_swiglu_carries_no_constants(): params = materialize_activation_params( SwigluActivation(), diff --git a/tests/unittest/_torch/moe/test_moe_impl.py b/tests/unittest/_torch/moe/test_moe_impl.py index 1f03273219c9..1f11a0fd413d 100644 --- a/tests/unittest/_torch/moe/test_moe_impl.py +++ b/tests/unittest/_torch/moe/test_moe_impl.py @@ -23,6 +23,7 @@ or numerical parity across backends, belong in ``test_moe_backend.py``. """ +from dataclasses import replace from unittest.mock import MagicMock import pytest @@ -934,6 +935,23 @@ def test_fc12_identity_round_trips_through_registry(): assert impl.input_requirement is impl.descriptor.input_requirement +@pytest.mark.parametrize("activation,accepted", [("SiTu", True), ("Relu2", False)]) +def test_fc12_pinned_activation_support(activation: str, accepted: bool) -> None: + problem = replace( + _fc12_problem(), + activation=activation, + activation_constants=frozenset({"alpha", "beta"}) if activation == "SiTu" else frozenset(), + ) + report = resolve_moe_impl( + ModelConfig(), problem=problem, deployment=_fc12_deployment(), impl_id=_FC12_IMPL_ID + ) + if accepted: + assert impl_class_for(report) is TrtllmCutedslFusedFc12Nvfp4Impl + else: + assert report.winner is None + assert report.rejected[0].reason is MoERejectReason.ACTIVATION_UNSUPPORTED + + def test_pinned_fc12_identity_fails_hard_where_the_backend_literal_degrades(): """On Rubin both tracks land on FC12; off Rubin the literal degrades and the pin does not.""" config = ModelConfig() From 65deaa423cddb45a8cc7cfde170c375e242214ea Mon Sep 17 00:00:00 2001 From: Xin Guan <294044352+xguannv@users.noreply.github.com> Date: Mon, 21 Sep 2026 16:30:12 +0800 Subject: [PATCH 2/5] [None][test] Gate FC12 SiTU test on Rubin DSL support Signed-off-by: Xin Guan <294044352+xguannv@users.noreply.github.com> --- tests/unittest/_torch/moe/test_kimi_k3_situ_moe.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/tests/unittest/_torch/moe/test_kimi_k3_situ_moe.py b/tests/unittest/_torch/moe/test_kimi_k3_situ_moe.py index bdfe256eb797..83106ff91d04 100644 --- a/tests/unittest/_torch/moe/test_kimi_k3_situ_moe.py +++ b/tests/unittest/_torch/moe/test_kimi_k3_situ_moe.py @@ -45,6 +45,7 @@ from utils.util import check_accuracy import tensorrt_llm._torch.models.modeling_kimi_linear as modeling_kimi_linear +from tensorrt_llm._torch.cute_dsl_utils import IS_CUTLASS_DSL_RUBIN_AVAILABLE from tensorrt_llm._torch.model_config import ModelConfig from tensorrt_llm._torch.models.modeling_kimi_linear import KimiK3MoEGate, KimiK3MoERuntime from tensorrt_llm._torch.moe.fused_moe.communication import CommunicationFactory @@ -1721,8 +1722,10 @@ def _swiglu_reference_moe(x, router_logits, routing_method, w1, w2, w3, alpha, b pytest.param( "CUTEDSL_FC12", marks=pytest.mark.skipif( - not torch.cuda.is_available() or get_sm_version() != 107, - reason="FC12 requires Rubin (SM107)", + not torch.cuda.is_available() + or get_sm_version() != 107 + or not IS_CUTLASS_DSL_RUBIN_AVAILABLE, + reason="FC12 requires Rubin (SM107) with CuTe DSL Rubin support", ), ), ], From 44c32c98e0981d98011939ae6a55b95e62507339 Mon Sep 17 00:00:00 2001 From: Xin Guan <294044352+xguannv@users.noreply.github.com> Date: Tue, 22 Sep 2026 13:47:14 +0800 Subject: [PATCH 3/5] [None][fix] Allow FC12 CuTe DSL for Kimi K3 SiTU Signed-off-by: Xin Guan <294044352+xguannv@users.noreply.github.com> --- tensorrt_llm/_torch/models/modeling_kimi_linear.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/tensorrt_llm/_torch/models/modeling_kimi_linear.py b/tensorrt_llm/_torch/models/modeling_kimi_linear.py index 93285b67c835..7ac94890ecde 100644 --- a/tensorrt_llm/_torch/models/modeling_kimi_linear.py +++ b/tensorrt_llm/_torch/models/modeling_kimi_linear.py @@ -1403,13 +1403,15 @@ def _routed_moe_model_config(model_config: ModelConfig) -> ModelConfig: "CUTLASS", "TRTLLM", "CUTEDSL", + "CUTEDSL_FC12", "MEGAMOE_DEEPGEMM", "MEGAMOE_CUTEDSL", } if model_config.moe_backend not in supported_backends: raise ValueError( "Kimi K3 SiTU routed experts only support the CUTLASS, TRTLLM, " - "CUTEDSL, MEGAMOE_DEEPGEMM, and MEGAMOE_CUTEDSL backends; " + "CUTEDSL, CUTEDSL_FC12, MEGAMOE_DEEPGEMM, and " + "MEGAMOE_CUTEDSL backends; " f"got {model_config.moe_backend!r}." ) if model_config.moe_load_balancer is not None: From 0f893d05f07ce5d8a8b1b5e4d328f42ac2776974 Mon Sep 17 00:00:00 2001 From: Xin Guan <294044352+xguannv@users.noreply.github.com> Date: Tue, 22 Sep 2026 14:49:51 +0800 Subject: [PATCH 4/5] [None][fix] Make Kimi K3 FC12 backend selection strict Signed-off-by: Xin Guan <294044352+xguannv@users.noreply.github.com> --- .../_torch/models/modeling_kimi_linear.py | 7 +++++- .../_torch/moe/test_kimi_k3_situ_moe.py | 23 ++++++++++++++----- 2 files changed, 23 insertions(+), 7 deletions(-) diff --git a/tensorrt_llm/_torch/models/modeling_kimi_linear.py b/tensorrt_llm/_torch/models/modeling_kimi_linear.py index 7ac94890ecde..1d4563c0518a 100644 --- a/tensorrt_llm/_torch/models/modeling_kimi_linear.py +++ b/tensorrt_llm/_torch/models/modeling_kimi_linear.py @@ -1190,7 +1190,12 @@ def __init__( # CUTLASS is absent on purpose: it is the fallback target, so # "degraded to CUTLASS" is not a thing that can happen to it. allow_backend_degradation=routed_moe_model_config.moe_backend - not in ("MEGAMOE_DEEPGEMM", "MEGAMOE_CUTEDSL", "CUTEDSL"), + not in ( + "MEGAMOE_DEEPGEMM", + "MEGAMOE_CUTEDSL", + "CUTEDSL", + "CUTEDSL_FC12", + ), ) self._check_trtllm_situ_quant( routed_moe_model_config.moe_backend, routed_quant_config.quant_algo diff --git a/tests/unittest/_torch/moe/test_kimi_k3_situ_moe.py b/tests/unittest/_torch/moe/test_kimi_k3_situ_moe.py index 48cefd68acee..33cfdf6014d4 100644 --- a/tests/unittest/_torch/moe/test_kimi_k3_situ_moe.py +++ b/tests/unittest/_torch/moe/test_kimi_k3_situ_moe.py @@ -555,7 +555,17 @@ def test_kimi_k3_moe_split_selection() -> None: assert KimiK3MoERuntime._select_moe_tp_ep(tep) == (4, 2) -@pytest.mark.parametrize("backend", ["CUTLASS", "TRTLLM", "MEGAMOE_DEEPGEMM", "MEGAMOE_CUTEDSL"]) +@pytest.mark.parametrize( + "backend", + [ + "CUTLASS", + "TRTLLM", + "CUTEDSL", + "CUTEDSL_FC12", + "MEGAMOE_DEEPGEMM", + "MEGAMOE_CUTEDSL", + ], +) def test_kimi_k3_routed_config_preserves_explicit_backend(backend): model_config = ModelConfig( mapping=Mapping(world_size=1, rank=0, tp_size=1), @@ -673,7 +683,8 @@ def admits_situ(name): assert "CUTEDSL" in declares_situ -def test_explicit_cutedsl_fails_instead_of_degrading_to_cutlass(monkeypatch): +@pytest.mark.parametrize("backend", ["CUTEDSL", "CUTEDSL_FC12"]) +def test_explicit_cutedsl_fails_instead_of_degrading_to_cutlass(monkeypatch, backend): """K3 must propagate strict backend selection through create_moe.""" from transformers.configuration_utils import PretrainedConfig @@ -688,7 +699,7 @@ def test_explicit_cutedsl_fails_instead_of_degrading_to_cutlass(monkeypatch): from tensorrt_llm.models.modeling_utils import QuantConfig # Keep the real resolver and K3 caller; only make eligibility deterministic. - for backend_cls in BACKEND_FAMILY["CUTEDSL"]: + for backend_cls in BACKEND_FAMILY[backend]: monkeypatch.setattr( backend_cls, "can_implement", @@ -707,7 +718,7 @@ def test_explicit_cutedsl_fails_instead_of_degrading_to_cutlass(monkeypatch): model_config = ModelConfig( pretrained_config=pretrained_config, mapping=Mapping(world_size=1, rank=0, tp_size=1), - moe_backend="CUTEDSL", + moe_backend=backend, quant_config_dict={"layers.0.mlp.experts": quant_config}, ) cfg = _K3Config(routed_expert_hidden_size=512, latent_moe_use_norm=True) @@ -726,8 +737,8 @@ def test_explicit_cutedsl_fails_instead_of_degrading_to_cutlass(monkeypatch): assert report.degraded assert impl_class_for(report) is CutlassFusedMoE - # Removing CUTEDSL from K3's no-degradation list must fail this assertion. - with pytest.raises(ValueError, match="CUTEDSL.*degradation disallowed") as excinfo: + # Removing either backend from K3's no-degradation list must fail this assertion. + with pytest.raises(ValueError, match=rf"{backend}.*degradation disallowed") as excinfo: KimiK3MoERuntime(model_config, cfg, layer_idx=0, aux_stream_dict={}) assert "dep_missing" in str(excinfo.value) From 6de6759c94c1d127e5237dad9516688fc624db3f Mon Sep 17 00:00:00 2001 From: Xin Guan <294044352+xguannv@users.noreply.github.com> Date: Sun, 27 Sep 2026 22:42:03 +0800 Subject: [PATCH 5/5] [None][fix] Validate FC12 SiTU parameters and test coverage Signed-off-by: Xin Guan <294044352+xguannv@users.noreply.github.com> --- .../deployment-guide-for-kimi-k3-on-trtllm.md | 5 ++- ...ntiguous_gather_grouped_gemm_act_fusion.py | 3 ++ ...ous_grouped_blockscaled_gemm_fused_fc12.py | 7 ++- .../_torch/models/modeling_kimi_linear.py | 4 +- .../moe/fused_moe/fused_moe_cute_dsl_fc12.py | 10 +++-- .../test_lists/test-db/l0_b200.yml | 5 +++ .../_torch/moe/test_kimi_k3_situ_moe.py | 45 ++++++++++++++++++- tests/unittest/_torch/moe/test_moe_backend.py | 7 ++- 8 files changed, 73 insertions(+), 13 deletions(-) diff --git a/docs/source/deployment-guide/deployment-guide-for-kimi-k3-on-trtllm.md b/docs/source/deployment-guide/deployment-guide-for-kimi-k3-on-trtllm.md index be0c68577dda..286746bc92f5 100644 --- a/docs/source/deployment-guide/deployment-guide-for-kimi-k3-on-trtllm.md +++ b/docs/source/deployment-guide/deployment-guide-for-kimi-k3-on-trtllm.md @@ -18,7 +18,7 @@ This guide uses Slurm and the `trtllm-llmapi-launch` multi-node launcher. The co | TEP8 | attention-TP, EP8 | 213 GB | GB300-class per-GPU memory | | TEP16 | attention-TP, EP16 | 115 GB | validated on GB200 (`SM100`) | - DEP16 replicates the BF16 non-expert weights on every rank (114 GB per rank) on top of the MXFP4 routed experts at 16-way expert parallelism (90 GB per rank). TEP16 shards those non-expert weights instead, which is what brings it within `SM100` per-GPU memory; TEP8 does not fit because its 8-way expert share alone is 181 GB per rank. On B200 (`SM100`), Kimi K3 is functionally supported at the kernel and module level and covered by unit tests in CI. Other GPU architectures are not supported. + DEP16 replicates the BF16 non-expert weights on every rank (114 GB per rank) on top of the MXFP4 routed experts at 16-way expert parallelism (90 GB per rank). TEP16 shards those non-expert weights instead, which is what brings it within `SM100` per-GPU memory; TEP8 does not fit because its 8-way expert share alone is 181 GB per rank. On B200 (`SM100`), Kimi K3 is functionally supported at the kernel and module level and covered by unit tests in CI. Full-model deployments on other GPU architectures are not covered by this guide. * Multi-node launcher: Slurm with the pyxis/enroot container plugin (or an equivalent MPI launcher) to start one rank per GPU across the nodes. * High-speed inter-node interconnect (e.g., NVLink/InfiniBand) for the expert-parallel traffic. * Shared filesystem visible to all nodes for the repository, the model weights, and the configuration file. @@ -34,7 +34,8 @@ The checkpoint and the configuration file must live on a shared filesystem visib ## Feature Support Notes -* **Blackwell only.** NVIDIA Blackwell GPUs are supported. The performance results in this guide were validated on NVIDIA GB300 NVL GPUs. Kimi K3 kernels and modules are also functional on B200 (`SM100`) and covered by unit tests in CI, and the TEP16 deployment is validated end-to-end on GB200 (`SM100`); DEP16 and TEP8 require GB300-class per-GPU memory (see Prerequisites). Support for other GPU architectures may be added in a future release. +* **Full-model deployments in this guide use Blackwell.** The performance results were validated on NVIDIA GB300 NVL GPUs. Kimi K3 kernels and modules are also functional on B200 (`SM100`) and covered by unit tests in CI, and the TEP16 deployment is validated end-to-end on GB200 (`SM100`); DEP16 and TEP8 require GB300-class per-GPU memory (see Prerequisites). +* **Rubin FC12 MoE kernel support.** On SM107, `CUTEDSL_FC12` supports SwiGLU and Kimi K3 SiTU for NVFP4 routed experts. This kernel/backend support does not validate the full Kimi K3 deployment on Rubin; the configurations and results below are for Blackwell. * **High-throughput and low-latency deployments are provided.** DEP16 (`enable_attention_dp: true`, `moe_expert_parallel_size: 16`) is the high-throughput deployment. TEP16 (`enable_attention_dp: false`, `moe_expert_parallel_size: 16`) is the low-latency deployment. An 8-GPU deployment, TEP8 (`enable_attention_dp: false`, `moe_expert_parallel_size: 8`), is also provided. Select the deployment and concurrency appropriate for your workload. * **CUDA graphs and the overlap scheduler are enabled.** The performance-sweep recipes set `disable_overlap_scheduler: false` and enable CUDA graphs. DEP16 additionally sets `cuda_graph_config.enable_padding: true`. * **Chunked prefill is supported and enabled** (`enable_chunked_prefill: true`), so prompts longer than `max_num_tokens` are scheduled across multiple steps. diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/blackwell/blockscaled_contiguous_gather_grouped_gemm_act_fusion.py b/tensorrt_llm/_torch/cute_dsl_kernels/blackwell/blockscaled_contiguous_gather_grouped_gemm_act_fusion.py index 8c9d85e68f46..f06f092507c4 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/blackwell/blockscaled_contiguous_gather_grouped_gemm_act_fusion.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/blackwell/blockscaled_contiguous_gather_grouped_gemm_act_fusion.py @@ -2857,6 +2857,9 @@ def _apply_situ_epilogue( beta * tanh(x/beta) = beta * (2*sigmoid(2x/beta) - 1) = 2*beta*sigmoid((2/beta)*x) - beta + + Keep the packed SiTU algebra in sync with the Rubin fused FC12 + ``rubin_contiguous_grouped_blockscaled_gemm_fused_fc12.py`` epilogue. """ beta = self.situ_beta linear_beta = self.situ_linear_beta diff --git a/tensorrt_llm/_torch/cute_dsl_kernels/rubin/moe/rubin_contiguous_grouped_blockscaled_gemm_fused_fc12.py b/tensorrt_llm/_torch/cute_dsl_kernels/rubin/moe/rubin_contiguous_grouped_blockscaled_gemm_fused_fc12.py index 37e7f3de4537..95e99c0b1032 100644 --- a/tensorrt_llm/_torch/cute_dsl_kernels/rubin/moe/rubin_contiguous_grouped_blockscaled_gemm_fused_fc12.py +++ b/tensorrt_llm/_torch/cute_dsl_kernels/rubin/moe/rubin_contiguous_grouped_blockscaled_gemm_fused_fc12.py @@ -971,9 +971,9 @@ def __init__( "ActivationType.SiTu requires both situ_beta and " f"situ_linear_beta, got {situ_beta} and {situ_linear_beta}." ) - if situ_beta <= 0 or situ_linear_beta <= 0: + if not (0 < situ_beta < float("inf") and 0 < situ_linear_beta < float("inf")): raise ValueError( - f"SiTU betas must be positive, got {situ_beta} and {situ_linear_beta}." + f"SiTU betas must be finite and positive, got {situ_beta} and {situ_linear_beta}." ) if self.has_swiglu_limit: raise ValueError( @@ -2343,6 +2343,9 @@ def _apply_situ_epilogue( Packed path uses ``tanh(z) = 2 * sigmoid(2z) - 1``. Betas fold at trace time; a different pair compiles a different kernel. + + Keep the packed SiTU algebra in sync with the Blackwell two-op + ``blockscaled_contiguous_gather_grouped_gemm_act_fusion.py`` epilogue. """ beta = self.situ_beta linear_beta = self.situ_linear_beta diff --git a/tensorrt_llm/_torch/models/modeling_kimi_linear.py b/tensorrt_llm/_torch/models/modeling_kimi_linear.py index 1d4563c0518a..a524efd42f64 100644 --- a/tensorrt_llm/_torch/models/modeling_kimi_linear.py +++ b/tensorrt_llm/_torch/models/modeling_kimi_linear.py @@ -260,9 +260,9 @@ def _resolve_kimi_situ_betas(cfg: Any) -> tuple[float, float]: "Kimi K3 routed SiTu experts require activation_situ_linear_beta; " "None means an identity linear branch that the fused kernels cannot represent." ) - if situ_beta <= 0 or situ_linear_beta <= 0: + if not (0 < situ_beta < float("inf") and 0 < situ_linear_beta < float("inf")): raise ValueError( - f"Kimi K3 SiTu betas must be positive; got {situ_beta} and {situ_linear_beta}." + f"Kimi K3 SiTu betas must be finite and positive; got {situ_beta} and {situ_linear_beta}." ) return float(situ_beta), float(situ_linear_beta) diff --git a/tensorrt_llm/_torch/moe/fused_moe/fused_moe_cute_dsl_fc12.py b/tensorrt_llm/_torch/moe/fused_moe/fused_moe_cute_dsl_fc12.py index 88d68beeaa8f..54a010b6654e 100644 --- a/tensorrt_llm/_torch/moe/fused_moe/fused_moe_cute_dsl_fc12.py +++ b/tensorrt_llm/_torch/moe/fused_moe/fused_moe_cute_dsl_fc12.py @@ -84,10 +84,14 @@ def forward( tactic: Optional[int], do_preparation: bool = False, ) -> torch.Tensor: + """Prime both fused tiles before outer CUDA-graph profiling. + + The inner fused op may need to compile or tune, which cannot happen + inside the outer profiler's CUDA graph. The parent's preparation also + passes a two-op memset argument that FC12 does not accept. This path + runs for SwiGLU as well as SiTU when the outer tuner prepares a shape. + """ if do_preparation: - # Prime the fused inner op before outer CUDA-graph profiling. - # The parent's preparation passes a two-op memset knob that - # FC12 does not accept: its memset is owned by the fused op. for tile_size in self._tile_sizes(): super().forward(inputs, tactic=tile_size) return inputs[4] diff --git a/tests/integration/test_lists/test-db/l0_b200.yml b/tests/integration/test_lists/test-db/l0_b200.yml index f4ced2cc0f6f..a5d29dc9bc6f 100644 --- a/tests/integration/test_lists/test-db/l0_b200.yml +++ b/tests/integration/test_lists/test-db/l0_b200.yml @@ -184,6 +184,11 @@ l0_b200: - unittest/_torch/moe/test_moe_backend.py::test_trtllm_gen_nvfp4_situ_selects_padded_quant_method - unittest/_torch/moe/test_moe_backend.py::test_trtllm_gen_nvfp4_situ_fc31_scale_c_drops_dequant_scale - unittest/_torch/moe/test_moe_backend.py::test_codegen_baked_situ_softcaps_are_uniform_scalars + - unittest/_torch/moe/test_moe_backend.py::test_fc12_rejects_nonuniform_situ_softcaps + - unittest/_torch/moe/test_moe_backend.py::test_fc12_swiglu_keeps_clamp_without_situ_constants + - unittest/_torch/moe/test_moe_backend.py::test_fc12_outer_tuning_separates_activation_and_softcaps + - unittest/_torch/moe/test_moe_backend.py::test_fc12_preparation_primes_fused_tiles_without_two_op_memset_knob + - unittest/_torch/moe/test_moe_impl.py::test_fc12_pinned_activation_support - unittest/_torch/moe/test_moe_backend.py::test_megamoe_plain_swiglu_carries_no_constants - unittest/_torch/moe/test_moe_backend.py::test_create_moe_forwards_situ_activation_as_one_carrier - unittest/_torch/moe/test_moe_backend.py::test_create_moe_backend_rejects_apply_router_weight_on_input_by_declaration diff --git a/tests/unittest/_torch/moe/test_kimi_k3_situ_moe.py b/tests/unittest/_torch/moe/test_kimi_k3_situ_moe.py index 33cfdf6014d4..76168cc45aaa 100644 --- a/tests/unittest/_torch/moe/test_kimi_k3_situ_moe.py +++ b/tests/unittest/_torch/moe/test_kimi_k3_situ_moe.py @@ -110,7 +110,16 @@ def test_kimi_situ_betas_require_linear_beta(): @pytest.mark.parametrize( "situ_beta,situ_linear_beta", - [(0.0, 25.0), (4.0, 0.0), (-1.0, 25.0), (4.0, -1.0)], + [ + (0.0, 25.0), + (4.0, 0.0), + (-1.0, 25.0), + (4.0, -1.0), + (float("nan"), 25.0), + (4.0, float("nan")), + (float("inf"), 25.0), + (4.0, float("inf")), + ], ) def test_kimi_situ_betas_must_be_positive(situ_beta, situ_linear_beta): cfg = SimpleNamespace( @@ -118,10 +127,42 @@ def test_kimi_situ_betas_must_be_positive(situ_beta, situ_linear_beta): activation_situ_linear_beta=situ_linear_beta, ) - with pytest.raises(ValueError, match="must be positive"): + with pytest.raises(ValueError, match="must be finite and positive"): modeling_kimi_linear._resolve_kimi_situ_betas(cfg) +@pytest.mark.parametrize( + "situ_beta,situ_linear_beta", + [ + (float("nan"), 25.0), + (4.0, float("nan")), + (float("inf"), 25.0), + (4.0, float("inf")), + ], + ids=["nan_gate", "nan_linear", "inf_gate", "inf_linear"], +) +def test_fc12_kernel_rejects_nonfinite_situ_betas(situ_beta, situ_linear_beta): + pytest.importorskip("cutlass") + pytest.importorskip("cutlass.utils.rubin_helpers") + from tensorrt_llm._torch.cute_dsl_kernels.rubin.moe.rubin_contiguous_grouped_blockscaled_gemm_fused_fc12 import ( # noqa: E501 + Sm107BlockScaledContiguousGroupedGemmFusedFc12Kernel, + ) + from tensorrt_llm._torch.utils import ActivationType + + with pytest.raises(ValueError, match="finite and positive"): + Sm107BlockScaledContiguousGroupedGemmFusedFc12Kernel( + sf_vec_size=16, + mma_inst_shape=(128, 128, 256), + mma_tiler=(128, 128, 256), + cluster_shape_mn=(1, 1), + vectorized_f32=True, + topk=8, + activation_type=ActivationType.SiTu, + situ_beta=situ_beta, + situ_linear_beta=situ_linear_beta, + ) + + @pytest.mark.parametrize( "situ_beta,situ_linear_beta", [ diff --git a/tests/unittest/_torch/moe/test_moe_backend.py b/tests/unittest/_torch/moe/test_moe_backend.py index a79aedaa30f3..2bab60fccc49 100644 --- a/tests/unittest/_torch/moe/test_moe_backend.py +++ b/tests/unittest/_torch/moe/test_moe_backend.py @@ -1301,10 +1301,13 @@ def test_codegen_baked_situ_softcaps_are_uniform_scalars(impl): assert params.beta == 25.0 -def test_fc12_rejects_nonuniform_situ_softcaps() -> None: +@pytest.mark.parametrize("nonuniform_softcap", ["gate", "linear"]) +def test_fc12_rejects_nonuniform_situ_softcaps(nonuniform_softcap: str) -> None: + gate_softcap = torch.tensor([4.0, 5.0]) if nonuniform_softcap == "gate" else 4.0 + linear_softcap = torch.tensor([25.0, 26.0]) if nonuniform_softcap == "linear" else 25.0 with pytest.raises(ValueError, match="uniform"): materialize_activation_params( - SiTuActivation(gate_softcap=torch.tensor([4.0, 5.0]), linear_softcap=25.0), + SiTuActivation(gate_softcap=gate_softcap, linear_softcap=linear_softcap), TrtllmCutedslFusedFc12Nvfp4Impl.activation_support, num_local_experts=2, owner="FC12",