diff --git a/python/sglang/kernels/jit/csrc/gemm/per_token_group_quant.cuh b/python/sglang/kernels/jit/csrc/gemm/per_token_group_quant.cuh index ff24e78f8..f4c751d6b 100644 --- a/python/sglang/kernels/jit/csrc/gemm/per_token_group_quant.cuh +++ b/python/sglang/kernels/jit/csrc/gemm/per_token_group_quant.cuh @@ -229,9 +229,10 @@ struct QuantTrait { static constexpr bool kAligned = kAligned_; static constexpr bool kFuseSiluAndMul = kFuseSiluAndMul_; static constexpr uint32_t kBlockSize = 256; - static constexpr uint32_t kVecSize = 32u / sizeof(InputType); + static constexpr uint32_t kVecSize = 32u / 2; static constexpr uint32_t kNumLanes = kGroupSize / kVecSize; - static_assert(sizeof(InputType) == 2, "only 16-bit inputs (bf16/fp16) are supported"); + static_assert(sizeof(InputType) == 2 || sizeof(InputType) == 4, "inputs must be 16-bit (bf16/fp16) or fp32"); + static_assert(sizeof(InputType) == 2 || !kFuseSiluAndMul, "fp32 inputs do not implement the fused silu"); static_assert(16 <= kGroupSize && kGroupSize <= 256, "supported group sizes are 16..256"); static_assert(kGroupSize % kVecSize == 0 && 1 <= kNumLanes && kNumLanes <= device::kWarpThreads); static_assert(!kUe8m0 || std::is_same_v, "ue8m0 scales imply fp8 output"); @@ -242,6 +243,19 @@ struct QuantTrait { const uint32_t token_idx, const uint32_t group_idx, const uint32_t lane_id) { + if constexpr (sizeof(InputType) == 4) { + run_fp32(params, expert_idx, token_idx, group_idx, lane_id); + } else { + run_packed16(params, expert_idx, token_idx, group_idx, lane_id); + } + } + + SGL_DEVICE static void run_packed16( + const QuantKernelParams& params, + const uint32_t expert_idx, + const uint32_t token_idx, + const uint32_t group_idx, + const uint32_t lane_id) { using deepseek_v4::fp8::cast_to_ue8m0; using deepseek_v4::fp8::inv_scale_ue8m0; using namespace device; @@ -316,6 +330,65 @@ struct QuantTrait { out.store(params.output.get(expert_idx, token_idx) + group_offset, lane_id); params.scale.store(expert_idx, token_idx, group_idx, scale_inv); } + + SGL_DEVICE static void run_fp32( + const QuantKernelParams& params, + const uint32_t expert_idx, + const uint32_t token_idx, + const uint32_t group_idx, + const uint32_t lane_id) { + using deepseek_v4::fp8::cast_to_ue8m0; + using deepseek_v4::fp8::inv_scale_ue8m0; + using namespace device; + using Q = QuantType; + using WTrait = detail::WeightTrait; + using Q2 = typename WTrait::packed2_t; + constexpr uint32_t kSubVec = kMaxVecBytes / sizeof(fp32_t); + constexpr uint32_t kNumSubVecs = kVecSize / kSubVec; + using in_vec_t = AlignedVector; + using out_vec_t = AlignedVector; + constexpr float kMaxValue = WTrait::kMaxValue; + constexpr float kMaxValueInv = 1.f / kMaxValue; + + const fp32_t* token_in = params.input.get(expert_idx, token_idx); + const uint32_t group_offset = group_idx * kGroupSize; + + in_vec_t in_vecs[kNumSubVecs]; +#pragma unroll + for (uint32_t v = 0; v < kNumSubVecs; ++v) { + in_vecs[v].load(token_in + group_offset, lane_id * kNumSubVecs + v); + } + const auto in = [&](const uint32_t i) { return in_vecs[i / kSubVec][i % kSubVec]; }; + + float local_amax = fabsf(in(0)); +#pragma unroll + for (uint32_t i = 1; i < kVecSize; ++i) { + local_amax = math::max(local_amax, fabsf(in(i))); + } + const auto amax = math::max(warp::reduce_max(local_amax), 1e-10f); + const float raw_scale = amax * kMaxValueInv; + + out_vec_t out; + detail::scale_t scale_inv; + float quant_scale; + if constexpr (kUe8m0) { + static_assert(std::is_same_v, "ue8m0 scales imply fp8 quantization"); + const auto exp = cast_to_ue8m0(raw_scale); + scale_inv = static_cast(exp); + quant_scale = inv_scale_ue8m0(exp); + } else { + scale_inv = raw_scale; + quant_scale = kMaxValue / amax; + } + const float2 quant_scale2 = {quant_scale, quant_scale}; +#pragma unroll + for (uint32_t i = 0; i < kVecSize / 2; ++i) { + out[i] = WTrait::quant(detail::mul2(float2{in(2 * i), in(2 * i + 1)}, quant_scale2)); + } + + out.store(params.output.get(expert_idx, token_idx) + group_offset, lane_id); + params.scale.store(expert_idx, token_idx, group_idx, scale_inv); + } }; // --------------------------------------------------------------------------- diff --git a/python/sglang/kernels/jit/csrc/kimi_k3/situ_and_mul.cuh b/python/sglang/kernels/jit/csrc/kimi_k3/situ_and_mul.cuh index dcb3ccf1f..5280de48c 100644 --- a/python/sglang/kernels/jit/csrc/kimi_k3/situ_and_mul.cuh +++ b/python/sglang/kernels/jit/csrc/kimi_k3/situ_and_mul.cuh @@ -64,11 +64,13 @@ struct SituAndMulParams { uint32_t stride_in_vecs; // input row stride in vector units (2*D/vec if dense) }; -template +template __global__ void situ_and_mul_kernel(const __grid_constant__ SituAndMulParams params) { using namespace device; - constexpr auto kVecSize = kMaxVecBytes / sizeof(T); - using vec_t = AlignedVector; + constexpr auto kWidest = sizeof(TIn) > sizeof(TOut) ? sizeof(TIn) : sizeof(TOut); + constexpr auto kVecSize = kMaxVecBytes / kWidest; + using vec_t = AlignedVector; + using out_vec_t = AlignedVector; const auto num_vecs = params.hidden_dim / kVecSize; // per token const auto tid = blockIdx.x * blockDim.x + threadIdx.x; @@ -94,23 +96,24 @@ __global__ void situ_and_mul_kernel(const __grid_constant__ SituAndMulParams par const float linear_beta = params.linear_beta; const float inv_linear_beta = params.inv_linear_beta; - vec_t out; + out_vec_t out; #pragma unroll for (int i = 0; i < kVecSize; ++i) { const float g = cast(gate[i]); const float u = cast(up[i]); - out[i] = cast(kimi_k3::situ_activate(g, u, beta, inv_beta, linear_beta, inv_linear_beta)); + out[i] = cast(kimi_k3::situ_activate(g, u, beta, inv_beta, linear_beta, inv_linear_beta)); } - store_as(params.out, out, output_offset); + store_as(params.out, out, output_offset); } // Host launcher -template +template struct SituAndMulKernel { - static constexpr auto kVecSize = device::kMaxVecBytes / sizeof(T); + static constexpr auto kWidest = sizeof(TIn) > sizeof(TOut) ? sizeof(TIn) : sizeof(TOut); + static constexpr auto kVecSize = device::kMaxVecBytes / kWidest; static constexpr auto kBlockSize = 256u; static void @@ -128,11 +131,11 @@ struct SituAndMulKernel { device_.set_options(); TensorMatcher({N, D_out}) // - .with_dtype() + .with_dtype() .with_device(device_) .verify(out); TensorMatcher({N, D_in}) // - .with_dtype() + .with_dtype() .with_device(device_) .with_strides({-1, 1}) .verify(input); @@ -166,9 +169,11 @@ struct SituAndMulKernel { }; if (has_linear_beta) { - LaunchKernel(num_blocks, kBlockSize, device).enable_pdl(kUsePDL)(situ_and_mul_kernel, params); + LaunchKernel(num_blocks, kBlockSize, device) + .enable_pdl(kUsePDL)(situ_and_mul_kernel, params); } else { - LaunchKernel(num_blocks, kBlockSize, device).enable_pdl(kUsePDL)(situ_and_mul_kernel, params); + LaunchKernel(num_blocks, kBlockSize, device) + .enable_pdl(kUsePDL)(situ_and_mul_kernel, params); } } }; diff --git a/python/sglang/kernels/jit/csrc/moe/route_quant_fused.cuh b/python/sglang/kernels/jit/csrc/moe/route_quant_fused.cuh index 19445c6e7..4a37e7506 100644 --- a/python/sglang/kernels/jit/csrc/moe/route_quant_fused.cuh +++ b/python/sglang/kernels/jit/csrc/moe/route_quant_fused.cuh @@ -26,8 +26,9 @@ struct RouteQuantFusedParams { // One quant CTA covers one token row: thread pairs (2g, 2g+1) hold group g // with lanes (0, 1) — the same subwarp layout the flat quant kernel derives // from global_tid, so the group reduction and stores are bit-identical. -using RouteQuantTrait = QuantTrait< - bf16_t, +template +using RouteQuantTraitT = QuantTrait< + TX, fp8_e4m3_t, /*kGroupSize=*/32, /*kUe8m0=*/true, @@ -35,10 +36,12 @@ using RouteQuantTrait = QuantTrait< /*kAligned=*/true, /*kFuseSiluAndMul=*/false>; +using RouteQuantTrait = RouteQuantTraitT; + inline constexpr uint32_t kQuantGroupsPerRow_ = LargeRouterRadixTrait::kBlockSize / RouteQuantTrait::kNumLanes; inline constexpr uint32_t kQuantHidden_ = kQuantGroupsPerRow_ * RouteQuantTrait::kGroupSize; // 3584 -template +template __global__ __launch_bounds__(LargeRouterRadixTrait::kBlockSize) // void route_quant_fused_kernel(const __grid_constant__ RouteQuantFusedParams params) { const auto M = static_cast(params.route.M); @@ -50,9 +53,9 @@ __global__ __launch_bounds__(LargeRouterRadixTrait::kBlockSize) // // as the routing CTAs, so they carry their own PDL wait/trigger. device::PDLWaitPrimary(); const uint32_t token_idx = blockIdx.x - M; - const uint32_t group_idx = threadIdx.x / RouteQuantTrait::kNumLanes; - const uint32_t lane_id = threadIdx.x % RouteQuantTrait::kNumLanes; - RouteQuantTrait::run(params.quant, /*expert_idx=*/0, token_idx, group_idx, lane_id); + const uint32_t group_idx = threadIdx.x / RouteQuantTraitT::kNumLanes; + const uint32_t lane_id = threadIdx.x % RouteQuantTraitT::kNumLanes; + RouteQuantTraitT::run(params.quant, /*expert_idx=*/0, token_idx, group_idx, lane_id); device::PDLTriggerSecondary(); } } @@ -99,11 +102,16 @@ struct RouteQuantFusedKernel { // Quant half: shape/stride/alignment checks + byte-stride munging shared // with the standalone flat kernel. - const auto ctx = build_quant_context(x, out_q, out_s); + auto x_dtype = SymbolicDType{}; + TensorMatcher({M_, -1}).with_dtype(x_dtype).with_device(device).with_strides({-1, 1}).verify(x); + const auto quant_params = + x_dtype.is_type() + ? build_quant_context, /*kMasked=*/false>(x, out_q, out_s).params + : build_quant_context, /*kMasked=*/false>(x, out_q, out_s).params; RuntimeCheck( - ctx.params.hidden_size == kQuantHidden_, "route_quant_fused is specialized for a 3584-wide activation row"); + quant_params.hidden_size == kQuantHidden_, "route_quant_fused is specialized for a 3584-wide activation row"); RuntimeCheck( - ctx.params.num_tokens == static_cast(M_.unwrap()), + quant_params.num_tokens == static_cast(M_.unwrap()), "route_quant_fused: scores and activations must have the same token count"); const auto M = static_cast(M_.unwrap()); @@ -125,16 +133,27 @@ struct RouteQuantFusedKernel { renormalize ? 1 : 0, apply_scale ? 1 : 0, /*sorted=*/0}, - .quant = ctx.params, + .quant = quant_params, }; +#define SGL_ROUTE_QUANT_LAUNCH(TS, TX) \ + LaunchKernel(2 * M, LargeRouterRadixTrait::kBlockSize, device.unwrap()) \ + .enable_pdl(kUsePDL)(route_quant_fused_kernel, params) + if (score_dtype.is_type()) { - LaunchKernel(2 * M, LargeRouterRadixTrait::kBlockSize, device.unwrap()) - .enable_pdl(kUsePDL)(route_quant_fused_kernel, params); + if (x_dtype.is_type()) { + SGL_ROUTE_QUANT_LAUNCH(fp32_t, fp32_t); + } else { + SGL_ROUTE_QUANT_LAUNCH(fp32_t, bf16_t); + } } else { - LaunchKernel(2 * M, LargeRouterRadixTrait::kBlockSize, device.unwrap()) - .enable_pdl(kUsePDL)(route_quant_fused_kernel, params); + if (x_dtype.is_type()) { + SGL_ROUTE_QUANT_LAUNCH(bf16_t, fp32_t); + } else { + SGL_ROUTE_QUANT_LAUNCH(bf16_t, bf16_t); + } } +#undef SGL_ROUTE_QUANT_LAUNCH } }; diff --git a/python/sglang/kernels/ops/gemm/cutedsl_bf16_gemm.py b/python/sglang/kernels/ops/gemm/cutedsl_bf16_gemm.py index 84ecb4c59..ac8ddff9f 100644 --- a/python/sglang/kernels/ops/gemm/cutedsl_bf16_gemm.py +++ b/python/sglang/kernels/ops/gemm/cutedsl_bf16_gemm.py @@ -119,8 +119,10 @@ class TgvGemmCuteExtKernel: pdl_launch: Optional[bool] = None, pdl_count: int = -1, has_bias: bool = False, + out_dtype: Type[cutlass.Numeric] = cutlass.BFloat16, ): self.acc_dtype = acc_dtype + self.out_dtype = out_dtype self.cta_m = cta_m self.cta_n = cta_n self.cta_k = cta_k @@ -164,6 +166,7 @@ class TgvGemmCuteExtKernel: f"TgvGemmCuteExtKernel_cta{self.cta_m}x{self.cta_n}x{self.cta_k}" f"_2cta{int(self.use_2cta)}_pdl{int(self.use_pdl)}" f"_bias{int(self.has_bias)}" + f"_out{self.out_dtype.__name__.lower()}" ) @cute.experimental.jit @@ -995,6 +998,11 @@ _TORCH_TO_CUTLASS_DTYPE = { torch.bfloat16: cutlass.BFloat16, } +_TORCH_TO_CUTLASS_OUT_DTYPE = { + torch.bfloat16: cutlass.BFloat16, + torch.float32: cutlass.Float32, +} + # Per-process cache mapping (dtype, config…) → compiled cute_ext callable. # We need to construct concrete cute.Tensors once and reuse the resulting # compiled function across all live calls; a fresh build per call would @@ -1029,7 +1037,12 @@ def _make_layout_tensor( def _make_compile_repr_tensors( - dtype: torch.dtype, has_bias: bool, a_leading: int, b_leading: int, c_leading: int + dtype: torch.dtype, + c_dtype: torch.dtype, + has_bias: bool, + a_leading: int, + b_leading: int, + c_leading: int, ): """Build representative tensors with strides matching the requested leading-dim pattern. After the A↔B swap, the cute_ext kernel sees: @@ -1051,7 +1064,7 @@ def _make_compile_repr_tensors( (L, K, M), dtype, b_leading ) # kernel B shape (L, K, M_pt) C_t = _make_layout_tensor( - (L, N, M), dtype, c_leading + (L, N, M), c_dtype, c_leading ) # kernel C shape (L, N_pt, M_pt) a_ = from_dlpack(A_t, assumed_align=32).mark_layout_dynamic(leading_dim=a_leading) @@ -1071,6 +1084,7 @@ def _make_compile_repr_tensors( def _get_compiled_cute_ext_kernel( dtype: torch.dtype, + c_dtype: torch.dtype, cta_m: int, cta_n: int, cta_k: int, @@ -1091,6 +1105,7 @@ def _get_compiled_cute_ext_kernel( """ key = ( dtype, + c_dtype, cta_m, cta_n, cta_k, @@ -1110,6 +1125,11 @@ def _get_compiled_cute_ext_kernel( raise ValueError( f"TGV cute_ext backend supports {list(_TORCH_TO_CUTLASS_DTYPE)}; got {dtype}." ) + if c_dtype not in _TORCH_TO_CUTLASS_OUT_DTYPE: + raise ValueError( + f"TGV cute_ext output supports {list(_TORCH_TO_CUTLASS_OUT_DTYPE)}; " + f"got {c_dtype}." + ) gemm = TgvGemmCuteExtKernel( acc_dtype=cutlass.Float32, @@ -1120,10 +1140,12 @@ def _get_compiled_cute_ext_kernel( use_2cta=use_2cta, use_pdl=use_pdl, has_bias=has_bias, + out_dtype=_TORCH_TO_CUTLASS_OUT_DTYPE[c_dtype], ) a_, b_, c_, bias_ = _make_compile_repr_tensors( dtype, + c_dtype, has_bias, a_leading, b_leading, @@ -1243,6 +1265,7 @@ def _run_tgv( compiled = _get_compiled_cute_ext_kernel( dtype=a.dtype, + c_dtype=out.dtype, cta_m=cta_m, cta_n=cta_n, cta_k=_TGV_CUTE_EXT_CTA_K, @@ -1385,7 +1408,7 @@ def _tgv_bf16_gemm_out_run( if not is_sm100_supported(): raise RuntimeError("cutedsl_bf16_gemm requires an SM10x GPU") assert x.dtype == torch.bfloat16 and weight.dtype == torch.bfloat16 - assert out.dtype == torch.bfloat16 and out.device == x.device + assert out.dtype in (torch.bfloat16, torch.float32) and out.device == x.device assert x.ndim == 2 and weight.ndim == 2 and out.ndim == 2 assert x.stride(-1) == 1, "x must be K-major [M, K]" assert weight.stride(-1) == 1, "weight must be K-major [N, K]" diff --git a/python/sglang/kernels/ops/kimi_k3/activation.py b/python/sglang/kernels/ops/kimi_k3/activation.py index 20527e00e..8a942f04b 100644 --- a/python/sglang/kernels/ops/kimi_k3/activation.py +++ b/python/sglang/kernels/ops/kimi_k3/activation.py @@ -32,9 +32,9 @@ def _fast_math_flags() -> list[str]: @cache_once -def _jit_situ_and_mul_module(dtype: torch.dtype) -> Module: - """Compile and cache the JIT SiTU-and-mul module for a given dtype.""" - args = make_cpp_args(dtype, is_arch_support_pdl()) +def _jit_situ_and_mul_module(in_dtype: torch.dtype, out_dtype: torch.dtype) -> Module: + """Compile and cache the JIT SiTU-and-mul module for an (in, out) dtype pair.""" + args = make_cpp_args(in_dtype, out_dtype, is_arch_support_pdl()) return load_jit( _make_name("situ_and_mul"), *args, @@ -50,7 +50,7 @@ def situ_and_mul( beta: float, linear_beta: Optional[float], ) -> torch.Tensor: - """Fused SiTU (SoftCap-GLU) activation: bf16 -> bf16. + """Fused SiTU (SoftCap-GLU) activation. gate_out = beta * tanh(gate / beta) * sigmoid(gate) up_out = linear_beta * tanh(up / linear_beta) [if linear_beta is not None] @@ -58,14 +58,16 @@ def situ_and_mul( Parameters ---------- - input : bf16 CUDA tensor [*, 2*D] - out : optional pre-allocated bf16 CUDA tensor [*, D] + input : bf16 or fp32 CUDA tensor [*, 2*D] + out : optional pre-allocated CUDA tensor [*, D]; its dtype selects + the output dtype (bf16 for an fp32 input) beta : gate softcap scalar (e.g. 4.0) linear_beta : up softcap scalar (e.g. 25.0), or None to skip """ hidden_size = input.shape[-1] // 2 if out is None: - out = input.new_empty(*input.shape[:-1], hidden_size) + out_dtype = torch.bfloat16 if input.dtype == torch.float32 else input.dtype + out = input.new_empty(*input.shape[:-1], hidden_size, dtype=out_dtype) # 2D inputs may be row-strided (e.g. a slice of a fused-GEMM output); # higher-rank inputs keep the dense-view path. @@ -76,7 +78,7 @@ def situ_and_mul( out_2d = out.view(-1, hidden_size) has_linear_beta = linear_beta is not None - module = _jit_situ_and_mul_module(input.dtype) + module = _jit_situ_and_mul_module(input_2d.dtype, out_2d.dtype) module.run( input_2d, out_2d, diff --git a/python/sglang/kernels/ops/moe/moe_route_quant_fused.py b/python/sglang/kernels/ops/moe/moe_route_quant_fused.py index 4320db5fd..58857df2f 100644 --- a/python/sglang/kernels/ops/moe/moe_route_quant_fused.py +++ b/python/sglang/kernels/ops/moe/moe_route_quant_fused.py @@ -81,7 +81,7 @@ def covered( and x.shape[0] == scores.shape[0] and 0 < x.shape[0] <= _MAX_TOKENS and x.shape[1] == _HIDDEN - and x.dtype == torch.bfloat16 + and x.dtype in (torch.bfloat16, torch.float32) and x.stride(1) == 1 and x.data_ptr() % 32 == 0 and (x.stride(0) * x.element_size()) % 32 == 0 diff --git a/python/sglang/kernels/ops/quantization/per_token_group_quant.py b/python/sglang/kernels/ops/quantization/per_token_group_quant.py index d485a5816..8d499a347 100644 --- a/python/sglang/kernels/ops/quantization/per_token_group_quant.py +++ b/python/sglang/kernels/ops/quantization/per_token_group_quant.py @@ -16,7 +16,7 @@ from sglang.srt.utils.custom_op import register_custom_op if TYPE_CHECKING: from tvm_ffi.module import Module -_SUPPORTED_INPUT_DTYPES = (torch.bfloat16, torch.float16) +_SUPPORTED_INPUT_DTYPES = (torch.bfloat16, torch.float16, torch.float32) _SUPPORTED_OUTPUT_DTYPES = (torch.float8_e4m3fn, torch.int8) _SUPPORTED_GROUP_SIZES = (16, 32, 64, 128, 256) diff --git a/python/sglang/srt/models/kimi_k3.py b/python/sglang/srt/models/kimi_k3.py index 98a213b9a..4c21d790e 100644 --- a/python/sglang/srt/models/kimi_k3.py +++ b/python/sglang/srt/models/kimi_k3.py @@ -137,11 +137,16 @@ def _k3_bf16_gemm( x: torch.Tensor, weight: torch.Tensor, out: Optional[torch.Tensor] = None, + out_dtype: Optional[torch.dtype] = None, ) -> torch.Tensor: """F.linear / torch.mm with the same TGV dispatch module-level GEMMs get through UnquantizedLinearMethod. The fused MoE front and the deferred shared down GEMM call torch directly on raw merged weights, so the --bf16-gemm-backend cutedsl selection would silently skip them.""" + if out is None and out_dtype is not None and out_dtype != x.dtype: + out = torch.empty( + (x.shape[0], weight.shape[0]), dtype=out_dtype, device=x.device + ) if x.dtype == torch.bfloat16 and weight.dtype == torch.bfloat16: from sglang.srt.layers.quantization.unquant import get_bf16_gemm_backend @@ -155,15 +160,11 @@ def _k3_bf16_gemm( if use_cutedsl_bf16_gemm(x.shape[0], weight.shape[0], weight.shape[1]): if out is None: return cutedsl_bf16_gemm(x, weight) - if out.is_contiguous(): - # TGV stores straight into caller memory (same entry the - # UnquantizedLinearMethod out-buffer path uses); no - # staging tensor + copy. - return cutedsl_bf16_gemm_out(x, weight, out) - out.copy_(cutedsl_bf16_gemm(x, weight)) - return out + return cutedsl_bf16_gemm_out(x, weight, out) if out is None: return torch.nn.functional.linear(x, weight) + if out.dtype != x.dtype: + return torch.mm(x, weight.t(), out=out, out_dtype=out.dtype) return torch.mm(x, weight.t(), out=out) @@ -478,14 +479,6 @@ class KimiK3MoE(nn.Module): _a2a_backend = get_moe_a2a_backend() self._ep_a2a = _a2a_backend.is_megamoe() or _a2a_backend.is_deepep() - # The flashinfer_mxfp4 (trtllm-gen) runner quantizes routed_input with - # the strided-input JIT group quant (_use_jit_mxfp8_quant in mxfp4.py), - # so the fused-front split view can be consumed as is; other runners - # (e.g. marlin) require a dense buffer. - self._moe_front_needs_contiguous = ( - not get_moe_runner_backend().is_flashinfer_mxfp4() - ) - # Defer the trtllm-gen finalize (top-k weighted unpermute) out of the # MoE op and fuse it into the push all-reduce's staging pass # (k3_ar_fusion.finalize_all_reduce_push_norm): the rank-local latent @@ -628,6 +621,7 @@ class KimiK3MoE(nn.Module): # Invalidate the cached properties. for prop in ( "_eligible_for_fused_front", + "_front_fp32", "_routing_contract_ok", "_ep_front_eligible", ): @@ -654,6 +648,19 @@ class KimiK3MoE(nn.Module): in (torch.bfloat16, torch.float16) ) + @cached_property + def _front_fp32(self) -> bool: + """Emit the merged front in fp32 so the router reads exact logits. + + The situ activation and the flashinfer_mxfp4 quantizer read the fp32 + slices directly. Every other runner takes routed_input rounded back to + bf16 in _forward_fused, which is bit-identical to the bf16 front.""" + return ( + not _is_hip + and self._eligible_for_fused_front + and self._front_w.dtype == torch.bfloat16 + ) + def _forward_mega_experts( self, routed_input: torch.Tensor, topk_output ) -> torch.Tensor: @@ -971,6 +978,28 @@ class KimiK3MoE(nn.Module): return _add3(out, shared_output, prefix_sum) return out if prefix_sum is None else out + prefix_sum + @cached_property + def _moe_front_needs_dense_bf16(self) -> bool: + """Whether routed_input must be repaired into a dense bf16 buffer. + + Only the SM100 trtllm-gen mxfp4 runner reads the front slice as it + comes: its group quant (route_quant_fused / per_token_group_quant) + takes both a strided row and an fp32 row. The SM90/SM120 cutlass mxfp4 + kernels return from apply() before that quant, and precision="bf16" + skips it as well, so those keep the bf16 contract even though the + runner backend is the same.""" + from sglang.srt.layers.quantization.mxfp4 import Mxfp4MoEMethod + + method = self.experts.quant_method + return not ( + isinstance(method, Mxfp4MoEMethod) + and method.use_flashinfer + and not method.use_marlin + and method._fi_kernel == "trtllm_sm100" + and method.flashinfer_mxfp4_moe_precision == "default" + and method.hidden_size == self.moe_hidden_size + ) + @cached_property def _route_quant_fuse_eligible(self) -> bool: """Whether to stage routed_input for the fused route+pack+quant launch @@ -1051,14 +1080,20 @@ class KimiK3MoE(nn.Module): ) num_tokens, hidden_size = hidden_states.shape - fused = _k3_bf16_gemm(hidden_states, self._front_w) + fused = _k3_bf16_gemm( + hidden_states, + self._front_w, + out_dtype=torch.float32 if self._front_fp32 else None, + ) gate_up, router_logits, routed_input = torch.split( fused, self._front_sizes, dim=-1 ) if num_tokens > 1 and _is_hip and not _aiter_k3_opt: router_logits = router_logits.contiguous() - if num_tokens > 1 and self._moe_front_needs_contiguous: - routed_input = routed_input.contiguous() + if self._moe_front_needs_dense_bf16: + # off an fp32 front the cast allocates the dense buffer, so the + # contiguous() behind it is free; off a bf16 front it is the copy + routed_input = routed_input.to(hidden_states.dtype).contiguous() latent_numel = num_tokens * self.moe_hidden_size if k3_ar_fusion.enabled(): # the shared-expert AR is pull-only, so its input must be a