执行摘要
- 一句话:修复 KimiLinear 模型 XPU 崩溃,回退无 JIT kernel 路径
- 推荐动作:值得精读的部分是平台 gating 的设计模式:用
is_xpu() 把不支持的 kernel 路径整体挡掉并复用既有 fallback(与 CPU/NPU 一致),以及 review 中『修改模型权重契约 vs 在 kernel 调用点适配』的决策——后者是更稳妥的做法。不建议作为架构参考,属于一次性适配;如后续保留该支持,建议补一个 XPU 单元测试或注释文档。
功能与动机
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 上不可用。
实现拆解
- KDA 注意力解码路径(kda_triton.py):将
TritonKDAKernel.supports_packed_decode 从 not 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 等其它平台零影响。
- 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,故不改动其它平台行为。
- 模型初始化(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();置 None 后 KimiMoE 内由 alt_stream is not None 守卫的双流重叠计算自动退化为串行。
- 测试与配套:无测试文件变更。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(模块 注意力内核;类别 source;类型 platform-gating;符号 TritonKDAKernel, supports_packed_decode): 核心改动:通过 supports_packed_decode 追加 not is_xpu(),决定 XPU 上 KDA decode 走哪条 kernel 路径,是修复 batched decode 崩溃的关键。
python/sglang/srt/layers/moe/topk.py(模块 专家路由;类别 source;类型 core-logic;符号 biased_grouped_topk_gpu, topk_sigmoid): 修复 XPU 上 topk_sigmoid AOT kernel 对 fp32 correction_bias 的硬性要求,位置在 MoE 路由核心路径。
python/sglang/srt/models/kimi_linear.py(模块 模型实现;类别 source;类型 data-contract;符号 KimiLinearModel, alt_stream): 模型入口改动:XPU 上跳过 get_stream("alt") 创建,避免 torch.cuda.Stream() 崩溃,同时引入 is_xpu 依赖。
关键符号:TritonKDAKernel.supports_packed_decode, KimiLinearModel.init, biased_grouped_topk_gpu
关键源码片段
python/sglang/srt/layers/attention/linear/kernels/kda_triton.py
核心改动:通过 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 Optional
import torch
from 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
修复 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_indices
topk_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
模型入口改动: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",
)
评论区精华
- 模型参数 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 调用点。
- 测试与 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 侧承担。
- 模型参数 e_score_correction_bias 是否应改为 fp32 (design): 保持模型参数 dtype 契约不变,改为在 topk.py 调用 sgl_kernel topk_sigmoid 时显式 .to(torch.float32),把平台适配逻辑收敛到 kernel 调用点。
- XPU 改动缺少测试覆盖且模型不在 KPI 列表 (testing): 作者未补充测试,mingfeima 最终 APPROVED;该支持主要面向 XPU 社区侧,回归风险由后续 XPU CI 承担。
风险与影响
- 风险:
- 正确性风险(kimi_linear.py):
alt_stream=None 依赖 KimiMoE 内部所有使用点都有 alt_stream is not None 守卫;若存在未保护的 alt_stream.record_stream() 或 torch.cuda 调用,会直接 NPE。当前仅凭 PR body 声明,无测试验证。
- 性能降级(kda_triton.py 与 kimi_linear.py):XPU 上 batched decode 从 packed 快速路径退回非 packed Triton 路径,MoE 双流重叠退化为串行,长上下文 decode 延迟可能明显升高;但 XPU 本无 packed kernel,属于可用性优先的正确取舍。
- 微小开销(topk.py):
correction_bias.to(torch.float32) 每次调用产生一次潜在拷贝(若已是 fp32 则无拷贝),影响可忽略。
- 回归风险:三处改动均被
is_xpu() 完全门控,CUDA/CPU/NPU/ROCm 路径行为不变;但缺少测试覆盖,后续 XPU CI 若不纳入该模型,回归难以及时发现。
- 影响:用户侧:Intel XPU 用户首次可以在 SGLang 上运行 KimiLinearForCausalLM(从崩溃变为可用),但 decode 与 MoE 为串行降级,不承诺性能。系统侧:改动是纯平台 gating,其它后端行为零变化,无 schema/配置/部署改动。团队侧:新增一处平台差异点,且该模型不在官方 KPI 列表,长期维护依赖 XPU 社区侧投入。
- 风险标记:缺少测试覆盖, XPU 上 decode 与 MoE 退化为串行执行, alt_stream 守卫依赖未验证, 平台 gating 逻辑分散在三处
关联脉络
- PR #35337 [XPU][CI] key persistent JIT kernel cache by image content ID: 同为 XPU 平台适配:本 PR 中 KDA packed decode 恰因 XPU 缺少 tvm_ffi CUDA JIT kernel 而回退,35337 则修复 XPU 上 JIT kernel 缓存与 dlopen 过期 .so 的问题,二者共同反映 XPU 上 kernel 生态的约束。
- PR #34481 [AMD] Keep the PTX-inline-asm diffusion norm fusions off on ROCm (fix FLUX warmup crash): 与本 PR 同属『按硬件平台关闭不支持的 kernel 快速路径、回退到通用 Triton 实现』的典型模式:一个是 ROCm 禁用 NV PTX norm fusion,一个是 XPU 禁用 KDA packed decode。
参与讨论