Fix the router GEMM inaccuracy when using _front_w in Kimi-K3 (#33764)
This commit is contained in:
@@ -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<QuantType, fp8_e4m3_t>, "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<Q>(expert_idx, token_idx) + group_offset, lane_id);
|
||||
params.scale.store<kUe8m0, kRowMajor, kAligned>(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<Q>;
|
||||
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<fp32_t, kSubVec>;
|
||||
using out_vec_t = AlignedVector<Q2, kVecSize / 2>;
|
||||
constexpr float kMaxValue = WTrait::kMaxValue;
|
||||
constexpr float kMaxValueInv = 1.f / kMaxValue;
|
||||
|
||||
const fp32_t* token_in = params.input.get<const fp32_t>(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<kNumLanes>(local_amax), 1e-10f);
|
||||
const float raw_scale = amax * kMaxValueInv;
|
||||
|
||||
out_vec_t out;
|
||||
detail::scale_t<kUe8m0> scale_inv;
|
||||
float quant_scale;
|
||||
if constexpr (kUe8m0) {
|
||||
static_assert(std::is_same_v<Q, fp8_e4m3_t>, "ue8m0 scales imply fp8 quantization");
|
||||
const auto exp = cast_to_ue8m0(raw_scale);
|
||||
scale_inv = static_cast<uint8_t>(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<Q>(expert_idx, token_idx) + group_offset, lane_id);
|
||||
params.scale.store<kUe8m0, kRowMajor, kAligned>(expert_idx, token_idx, group_idx, scale_inv);
|
||||
}
|
||||
};
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
@@ -64,11 +64,13 @@ struct SituAndMulParams {
|
||||
uint32_t stride_in_vecs; // input row stride in vector units (2*D/vec if dense)
|
||||
};
|
||||
|
||||
template <typename T, bool kHasLinearBeta, bool kUsePDL>
|
||||
template <typename TIn, typename TOut, bool kHasLinearBeta, bool kUsePDL>
|
||||
__global__ void situ_and_mul_kernel(const __grid_constant__ SituAndMulParams params) {
|
||||
using namespace device;
|
||||
constexpr auto kVecSize = kMaxVecBytes / sizeof(T);
|
||||
using vec_t = AlignedVector<T, kMaxVecBytes / sizeof(T)>;
|
||||
constexpr auto kWidest = sizeof(TIn) > sizeof(TOut) ? sizeof(TIn) : sizeof(TOut);
|
||||
constexpr auto kVecSize = kMaxVecBytes / kWidest;
|
||||
using vec_t = AlignedVector<TIn, kVecSize>;
|
||||
using out_vec_t = AlignedVector<TOut, kVecSize>;
|
||||
|
||||
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<fp32_t>(gate[i]);
|
||||
const float u = cast<fp32_t>(up[i]);
|
||||
|
||||
out[i] = cast<T>(kimi_k3::situ_activate<kHasLinearBeta>(g, u, beta, inv_beta, linear_beta, inv_linear_beta));
|
||||
out[i] = cast<TOut>(kimi_k3::situ_activate<kHasLinearBeta>(g, u, beta, inv_beta, linear_beta, inv_linear_beta));
|
||||
}
|
||||
|
||||
store_as<vec_t>(params.out, out, output_offset);
|
||||
store_as<out_vec_t>(params.out, out, output_offset);
|
||||
}
|
||||
|
||||
// Host launcher
|
||||
|
||||
template <typename T, bool kUsePDL>
|
||||
template <typename TIn, typename TOut, bool kUsePDL>
|
||||
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<kDLCUDA>();
|
||||
|
||||
TensorMatcher({N, D_out}) //
|
||||
.with_dtype<T>()
|
||||
.with_dtype<TOut>()
|
||||
.with_device(device_)
|
||||
.verify(out);
|
||||
TensorMatcher({N, D_in}) //
|
||||
.with_dtype<T>()
|
||||
.with_dtype<TIn>()
|
||||
.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<T, true, kUsePDL>, params);
|
||||
LaunchKernel(num_blocks, kBlockSize, device)
|
||||
.enable_pdl(kUsePDL)(situ_and_mul_kernel<TIn, TOut, true, kUsePDL>, params);
|
||||
} else {
|
||||
LaunchKernel(num_blocks, kBlockSize, device).enable_pdl(kUsePDL)(situ_and_mul_kernel<T, false, kUsePDL>, params);
|
||||
LaunchKernel(num_blocks, kBlockSize, device)
|
||||
.enable_pdl(kUsePDL)(situ_and_mul_kernel<TIn, TOut, false, kUsePDL>, params);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
@@ -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 <typename TX>
|
||||
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<bf16_t>;
|
||||
|
||||
inline constexpr uint32_t kQuantGroupsPerRow_ = LargeRouterRadixTrait::kBlockSize / RouteQuantTrait::kNumLanes;
|
||||
inline constexpr uint32_t kQuantHidden_ = kQuantGroupsPerRow_ * RouteQuantTrait::kGroupSize; // 3584
|
||||
|
||||
template <bool kUsePDL, typename TScore>
|
||||
template <bool kUsePDL, typename TScore, typename TX>
|
||||
__global__ __launch_bounds__(LargeRouterRadixTrait::kBlockSize) //
|
||||
void route_quant_fused_kernel(const __grid_constant__ RouteQuantFusedParams params) {
|
||||
const auto M = static_cast<uint32_t>(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<kUsePDL>();
|
||||
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<TX>::kNumLanes;
|
||||
const uint32_t lane_id = threadIdx.x % RouteQuantTraitT<TX>::kNumLanes;
|
||||
RouteQuantTraitT<TX>::run(params.quant, /*expert_idx=*/0, token_idx, group_idx, lane_id);
|
||||
device::PDLTriggerSecondary<kUsePDL>();
|
||||
}
|
||||
}
|
||||
@@ -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<Trait, /*kMasked=*/false>(x, out_q, out_s);
|
||||
auto x_dtype = SymbolicDType{};
|
||||
TensorMatcher({M_, -1}).with_dtype<bf16_t, fp32_t>(x_dtype).with_device(device).with_strides({-1, 1}).verify(x);
|
||||
const auto quant_params =
|
||||
x_dtype.is_type<fp32_t>()
|
||||
? build_quant_context<RouteQuantTraitT<fp32_t>, /*kMasked=*/false>(x, out_q, out_s).params
|
||||
: build_quant_context<RouteQuantTraitT<bf16_t>, /*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<uint32_t>(M_.unwrap()),
|
||||
quant_params.num_tokens == static_cast<uint32_t>(M_.unwrap()),
|
||||
"route_quant_fused: scores and activations must have the same token count");
|
||||
|
||||
const auto M = static_cast<uint32_t>(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<kUsePDL, TS, TX>, params)
|
||||
|
||||
if (score_dtype.is_type<fp32_t>()) {
|
||||
LaunchKernel(2 * M, LargeRouterRadixTrait::kBlockSize, device.unwrap())
|
||||
.enable_pdl(kUsePDL)(route_quant_fused_kernel<kUsePDL, fp32_t>, params);
|
||||
if (x_dtype.is_type<fp32_t>()) {
|
||||
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<kUsePDL, bf16_t>, params);
|
||||
if (x_dtype.is_type<fp32_t>()) {
|
||||
SGL_ROUTE_QUANT_LAUNCH(bf16_t, fp32_t);
|
||||
} else {
|
||||
SGL_ROUTE_QUANT_LAUNCH(bf16_t, bf16_t);
|
||||
}
|
||||
}
|
||||
#undef SGL_ROUTE_QUANT_LAUNCH
|
||||
}
|
||||
};
|
||||
|
||||
|
||||
@@ -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]"
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user