Prhub

#34546 [XPU] Fix/kimi linear xpu

原始 PR 作者 SKRohit 合并时间 2026-08-20 12:43 文件变更 3 提交数 14 评论 3 代码增减 +9 / -5

执行摘要

修复 KimiLinear 模型 XPU 崩溃,回退无 JIT kernel 路径

PR body 明确说明动机:"Enables KimiLinearForCausalLM (hybrid KDA linear-attention + MLA + MoE) to run on Intel XPU."。具体要解决三个崩溃点:

1) XPU 没有 KDA packed decode 所需的 tvm_ffi CUDA JIT kernel,batched decode 直接崩溃;
2) grouped-topk 的 sgl_kernel topk_sigmoid 在 XPU 上要求 fp32 的 correction bias;
3) get_stream("alt") 内部创建 torch.cuda.Stream(),在无 CUDA 的 XPU 上不可用。

值得精读的部分是平台 gating 的设计模式:用 is_xpu() 把不支持的 kernel 路径整体挡掉并复用既有 fallback(与 CPU/NPU 一致),以及 review 中『修改模型权重契约 vs 在 kernel 调用点适配』的决策——后者是更稳妥的做法。不建议作为架构参考,属于一次性适配;如后续保留该支持,建议补一个 XPU 单元测试或注释文档。

讨论亮点
  1. 模型参数 dtype 之争(已解决):mingfeima 在 kimi_linear.py 的 diff 上要求回退 e_score_correction_bias 改 fp32 的改动,原话 "we cannot do this type of change. need to keep what it was.";SKRohit 回复 "Removed changes from model implementation code." 并回退,最终把转换移到 topk.py 调用点。
  2. 测试与 KPI 质疑(已解决):mingfeima 在 CHANGES_REQUESTED 中问 "do we have test case to cover this change? model not in our KPI list (no current/upcoming device can run this model)"。作者未补测试,mingfeima 最终 APPROVED,支持责任主要由 XPU 侧承担。

实现拆解

  1. KDA 注意力解码路径(kda_triton.py):将 TritonKDAKernel.supports_packed_decodenot is_cpu() and not is_npu() 扩展为追加 not is_xpu()。原因是 XPU 上 fused_recurrent_kda_packed_decode 依赖的 tvm_ffi CUDA JIT kernel 不可用,关闭该标志后 XPU 自动复用非 packed 的 Triton decode() 路径(fused_sigmoid_gating_delta_rule_update),与 CPU/NPU 的 fallback 一致,batched decode 通过 query_start_loc 驱动。对 CUDA 等其它平台零影响。
  2. MoE 分组 topk(topk.py):在 biased_grouped_topk_gpu 的 XPU 分支中,调用 sgl_kernel.topk_sigmoid 前把 correction_bias 显式 .to(torch.float32)。原因是 XPU 的 topk_sigmoid AOT kernel 只接受 fp32 correction bias;CUDA 路径内部已有 cast,故不改动其它平台行为。
  3. 模型初始化(kimi_linear.py):新增 is_xpu 导入,并将 KimiLinearModel.__init__ 中的 self.alt_stream = get_stream("alt") 改为 self.alt_stream = None if is_xpu() else get_stream("alt")。原因是 get_stream("alt") 会构造 torch.cuda.Stream();置 NoneKimiMoE 内由 alt_stream is not None 守卫的双流重叠计算自动退化为串行。
  4. 测试与配套:无测试文件变更。review 期间曾尝试把 self.gate.e_score_correction_bias 参数创建为 fp32(匹配 DeepseekV2),被 reviewer 以 "need to keep what it was" 拒绝后回退,最终将 fp32 转换收敛到 topk.py 调用点,保持模型参数 dtype 契约不变。
文件 模块 状态 重要度
python/sglang/srt/layers/attention/linear/kernels/kda_triton.py 注意力内核 modified 5.8
python/sglang/srt/layers/moe/topk.py 专家路由 modified 4.82
python/sglang/srt/models/kimi_linear.py 模型实现 modified 5.21

关键符号

TritonKDAKernel.supports_packed_decode KimiLinearModel.__init__ biased_grouped_topk_gpu

关键源码片段

python/sglang/srt/layers/attention/linear/kernels/kda_triton.py platform-gating

核心改动:通过 supports_packed_decode 追加 not is_xpu(),决定 XPU 上 KDA decode 走哪条 kernel 路径,是修复 batched decode 崩溃的关键。

# python/sglang/srt/layers/attention/linear/kernels/kda_triton.py
from typing import Optionalimport torchfrom sglang.srt.layers.attention.linear.kernels.kernel_backend import (
    LinearAttnKernelBase,
)
from sglang.srt.utils import is_cpu, is_npu, is_xpu# 非 CPU 环境才导入依赖 tvm_ffi CUDA JIT 的 kernel;
# XPU 环境虽无 packed kernel,但仍可复用 Triton 版 decode 实现。
if not is_cpu():
    from sglang.kernels.ops.attention.fla.fused_recurrent import (
        fused_recurrent_kda_packed_decode,
    )
    from sglang.kernels.ops.attention.fla.fused_recurrent_linear_replayssm import (
        fused_recurrent_linear_replayssm_decode,
    )
    from sglang.kernels.ops.attention.fla.fused_sigmoid_gating_recurrent import (
        fused_sigmoid_gating_delta_rule_update,
    )
    from sglang.kernels.ops.attention.fla.kda import chunk_kda
​
​
class TritonKDAKernel(LinearAttnKernelBase):
    """基于 Triton 的 KDA(Kimi Delta Attention)线性注意力 kernel。"""
​
    # XPU 上没有 packed decode 所需的 tvm_ffi CUDA JIT kernel,
    # 因此把 XPU 也排除:batched decode 改走非 packed 的 Triton
    # decode() 路径(fused_sigmoid_gating_delta_rule_update),
    # 与 CPU/NPU 的 fallback 完全一致,由 query_start_loc 驱动。
    supports_packed_decode: bool = not is_cpu() and not is_npu() and not is_xpu()
​
    def packed_decode(self, mixed_qkv, a, b, *, A_log, dt_bias, scale, ssm_states,
                      cache_indices, num_v_heads, head_v_dim, lower_bound=None, **kwargs):
        """Packed decode 快速路径:仅在 supports_packed_decode 为 True 时被调用。"""
        B = mixed_qkv.shape[0]
        out = mixed_qkv.new_empty(B, 1, num_v_heads, head_v_dim)
​
        # 这里是 KDA ReplaySSM buffered decode 的入口,依赖 tvm_ffi 生成的
        # CUDA JIT kernel,故 XPU 上不会进入该方法。
        ...
        return out
python/sglang/srt/layers/moe/topk.py core-logic

修复 XPU 上 topk_sigmoid AOT kernel 对 fp32 correction_bias 的硬性要求,位置在 MoE 路由核心路径。

# python/sglang/srt/layers/moe/topk.py —— biased_grouped_topk_gpu 的 XPU 分支# 当满足 num_fused_shared_experts == 0、num_experts <= 256、topk <= 8 等条件时,
# 走 sgl_kernel 的 topk_sigmoid 快速路径。XPU 上该 AOT kernel 要求
# correction_bias 为 fp32,因此在调用点显式转换;CUDA 路径内部已自行
# cast 到 fp32,故此处转换不影响其它平台。
num_tokens = gating_output.shape[0]
topk_values = torch.empty(
    (num_tokens, topk), dtype=torch.float32, device=gating_output.device
)
topk_indices = torch.empty(
    (num_tokens, topk), dtype=torch.int32, device=gating_output.device
)if num_tokens == 0:
    return topk_values, topk_indicestopk_sigmoid(
    topk_values,
    topk_indices,
    gating_output,
    renormalize,
    # XPU 的 topk_sigmoid AOT kernel 只接受 fp32 的 correction bias
    correction_bias.to(torch.float32),
    scale,
)return topk_values, topk_indices
python/sglang/srt/models/kimi_linear.py data-contract

模型入口改动:XPU 上跳过 get_stream("alt") 创建,避免 torch.cuda.Stream() 崩溃,同时引入 is_xpu 依赖。

# python/sglang/srt/models/kimi_linear.py —— KimiLinearModel.__init__ 片段if self.pp_group.is_first_rank:
    self.embed_tokens = VocabParallelEmbedding(
        config.vocab_size,
        config.hidden_size,
        prefix=f"{prefix}.embed_tokens",
    )
else:
    self.embed_tokens = PPMissingLayer()# XPU 下 get_stream("alt") 内部会创建 torch.cuda.Stream(),直接调用会崩溃;
# 置为 None 后,KimiMoE 的双流重叠计算会被 alt_stream is not None 守卫跳过,
# 退化为串行执行(可用性优先,性能让位)。
self.alt_stream = None if is_xpu() else get_stream("alt")self.layers, self.start_layer, self.end_layer = make_layers(
    config.num_hidden_layers,
    lambda idx, prefix: KimiDecoderLayer(
        layer_idx=idx,
        config=config,
        quant_config=quant_config,
        prefix=prefix,
        alt_stream=self.alt_stream,
    ),
    pp_rank=self.pp_group.rank_in_group,
    pp_size=self.pp_group.world_size,
    prefix=f"{prefix}.layers",
)

评论区精华

模型参数 e_score_correction_bias 是否应改为 fp32 设计

mingfeima 在 kimi_linear.py 的 diff 上指出:"we cannot do this type of change. need to keep what it was." —— 不希望改动模型实现中路由 bias 参数的 dtype;SKRohit 回复 "Removed changes from model implementation code." 并回退了该改动。

结论:保持模型参数 dtype 契约不变,改为在 topk.py 调用 sgl_kernel topk_sigmoid 时显式 .to(torch.float32),把平台适配逻辑收敛到 kernel 调用点。 · 已解决

XPU 改动缺少测试覆盖且模型不在 KPI 列表 测试

mingfeima 在 CHANGES_REQUESTED 评论中质疑:"do we have test case to cover this change? model not in our KPI list (no current/upcoming device can run this model)"。

结论:作者未补充测试,mingfeima 最终 APPROVED;该支持主要面向 XPU 社区侧,回归风险由后续 XPU CI 承担。 · 已解决

风险与影响

  1. 正确性风险(kimi_linear.py)alt_stream=None 依赖 KimiMoE 内部所有使用点都有 alt_stream is not None 守卫;若存在未保护的 alt_stream.record_stream()torch.cuda 调用,会直接 NPE。当前仅凭 PR body 声明,无测试验证。
  2. 性能降级(kda_triton.py 与 kimi_linear.py):XPU 上 batched decode 从 packed 快速路径退回非 packed Triton 路径,MoE 双流重叠退化为串行,长上下文 decode 延迟可能明显升高;但 XPU 本无 packed kernel,属于可用性优先的正确取舍。
  3. 微小开销(topk.py)correction_bias.to(torch.float32) 每次调用产生一次潜在拷贝(若已是 fp32 则无拷贝),影响可忽略。
  4. 回归风险:三处改动均被 is_xpu() 完全门控,CUDA/CPU/NPU/ROCm 路径行为不变;但缺少测试覆盖,后续 XPU CI 若不纳入该模型,回归难以及时发现。

用户侧:Intel XPU 用户首次可以在 SGLang 上运行 KimiLinearForCausalLM(从崩溃变为可用),但 decode 与 MoE 为串行降级,不承诺性能。系统侧:改动是纯平台 gating,其它后端行为零变化,无 schema/配置/部署改动。团队侧:新增一处平台差异点,且该模型不在官方 KPI 列表,长期维护依赖 XPU 社区侧投入。

缺少测试覆盖 XPU 上 decode 与 MoE 退化为串行执行 alt_stream 守卫依赖未验证 平台 gating 逻辑分散在三处

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论