# PR #34546 完整报告

- 仓库：`sgl-project/sglang`
- 标题：[XPU] Fix/kimi linear xpu
- 合并时间：2026-08-20 12:43
- 原文链接：http://prhub.com.cn/sgl-project/sglang/pull/34546

---

# 执行摘要

- 一句话：修复 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 上不可用。

# 实现拆解

1. **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 等其它平台零影响。
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()`；置 `None` 后 `KimiMoE` 内由 `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`（模块 注意力内核；类别 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
# 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
# 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
# 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",
)

```

# 评论区精华

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 侧承担。

- 模型参数 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 承担。

# 风险与影响

- 风险：
 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 逻辑分散在三处

# 关联脉络

- 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。