04 问题深拆①:sm80 没有稀疏 MLA 注意力后端怎么办
系列:GLM-5.3-NVFP4 部署实录(8×A800 / sm80)。
上一篇:03 踩坑总览。
本文覆盖坑 1、2、3、4、5、10、11------整个部署中工作量最大的一组问题。
- GLM-5.3 的注意力机制是 DSA(DeepSeek 式稀疏注意力)+ MLA :每个 query 只从 2048 个(index_topk)候选 token 里算注意力。这个设计在 Hopper 之后的卡上有专门的高性能后端,但在 A800(sm80)上,vLLM 给出的答案是没有后端,直接抛异常。本文记录我们如何把这条路径从无到有修通。
1. 第一现场
- 服务启动、模型加载完,进入注意力后端选择阶段即崩:
bash
No valid attention backend found
- 查看 0.28.0 源码,sm80 分支的候选后端列表里,与稀疏 MLA 相关的只有三个:
flashmla_sparse、flashattn_mla_sparse、flashinfer_mla_sparse------全部要求 SM90+/SM100+。Ampere 上一个都没有。
2. 解法:移植社区 PR #47629(TRITON_MLA_SPARSE)
- 社区已经验证过一条 sm80 路径:PR #47629(接管自 #38476)的 TRITON_MLA_SPARSE 后端------纯 Triton 实现的稀疏 MLA,带 split-KV decode 加速。它在 0.26.0 上服务过 GLM-5.2,我们的任务是把 vendor 进 0.28.0:
- 新增 3 个文件 :后端类
triton_mla_sparse.py、稀疏 MLA Triton 内核triton_mla_sparse_kernel.py(_DIM_QK=576,即 512 kv_lora + 64 rope)、FP8 MQA logits Triton 内核mqa_logits_triton.py; - 注册两处 :
registry.py加枚举;platforms/cuda.py的 sm80 候选表(_get_backend_priorities的 else 分支)末尾追加TRITON_MLA_SPARSE。改动本身小到可以全贴出来:
python
# vllm/v1/attention/backends/registry.py ------ 加一个枚举成员
class AttentionBackendEnum(Enum):
...
TRITON_MLA_SPARSE = (auto())
# vllm/platforms/cuda.py ------ sm80 走的 else 分支候选表末尾追加
else: # sm80 (Ampere)
priorities = [..., AttentionBackendEnum.TRITON_MLA_SPARSE]
验证通过的标志日志:
bash
Using TRITON_MLA_SPARSE attention backend out of potential backends: ['TRITON_MLA_SPARSE']
一个容易误判的细节 :0.28.0 对 VLLM_ATTENTION_BACKEND 环境变量会打 "Unknown environment variable" 警告并忽略它------后端命中靠的是 cuda.py 候选表补丁,不是这个 env。如果你盯着 env 没生效就以为补丁没打上,会白绕一圈。
3. 修通后端之后的连环小坑
- 后端选上只是第一步,稀疏路径上还有一串 0.28.0 的兼容性断点:
坑 3/4:metadata 计数字段缺失
0.28.0 的 mla_attention.py 有硬断言:
python
assert (
attn_metadata.num_decodes is not None
and attn_metadata.num_prefills is not None
and attn_metadata.num_decode_tokens is not None
)
- 而 0.28.0 原生
XPUMLASparseMetadata没有这些字段(它们是 0.26.0 时代加的)。解法是在 vendor 的xpu_mla_sparse.py里做计数 graft 。graft 分两层:metadata dataclass 补字段(全部带默认值、置于非默认字段之后,这就是坑 4),构造函数里用 0.28.0 现成的split_decodes_and_prefills()真正算出这些计数:
python
# 补字段(坑 4:带默认值的字段必须放在无默认值字段之后)
@dataclass
class XPUMLASparseMetadata(...):
... # 既有非默认字段
num_decode_tokens: int = 0
num_prefill_tokens: int = 0
num_decodes: int = 0
num_prefills: int = 0
prefill_max_seq_len: int = 0
seq_lens: Optional[torch.Tensor] = None
prefill: Optional[object] = None
# 构造时填充(0.26.0 与 0.28.0 的该工具函数签名一致)
(num_decodes, num_prefills, num_decode_tokens, num_prefill_tokens) = \
split_decodes_and_prefills(common_attn_metadata, ...)
坑 2:prefill 崩在 forward_mha
- 后端通了、能 forward 了,但真实 prefill(935 token prompt)崩在
forward_mha的 NotImplementedError。0.28.0 里sparse_mla_force_mqa配置项仍然存在(vllm/config/attention.py),它能让注意力强制走 MQA 路径绕开未实现的分支。启动参数加:
bash
--attention-config '{"sparse_mla_force_mqa": true}'
此项必填,去掉必崩。
坑 10:masked_mha_available 属性探测
- 0.28.0 的
mla_attention.py对所有 sparse 后端实现探测impl.masked_mha_available(Blackwell masked-MHA 快路径的开关),vendor 的TritonMLASparseImpl没有这个属性,直接 AttributeError。解法一行:impl 的__init__补self.masked_mha_available = False。
4. 最大的坑(11):C++ indexer op 的 deep_gemm 依赖
- 这是静态阶段就标记的唯一风险,实机坐实。0.28.0 把 DSA indexer(负责从 2048 个候选里挑 token 的那部分)重写成了单个 C++ op
torch.ops.vllm.sparse_attn_indexer。实机在profile_run阶段崩:
bash
attention.hpp:213 Unsupported architecture
- 原因:这个 C++ op 内部直调 deep_gemm 的
fp8_fp4_mqa_logits(prefill)/fp8_fp4_paged_mqa_logits(decode),deep_gemm 显式断言 SM90+。
更隐蔽的是坑 5 的铺垫:indexer.py 里的探测函数用的是 has_deep_gemm()------纯 import 探测 。sm80 上 vendor 的 vllm.third_party.deep_gemm 是可以 import 成功的,于是探测放行,然后才在内核里崩。0.28.0 其实已经提供了架构感知的 is_deep_gemm_supported()(= VLLM_USE_DEEP_GEMM and has_deep_gemm() and platform.support_deep_gemm(),sm80 返回 False),只是调用方没换,我们先把这个换掉即可。
4.1 方案决策:A / B / C
-
方案 A:整体回退 0.26.0 的 indexer 实现。改动大,等于放弃 0.28.0 的重写。
-
方案 B :直接用现成的
v0.26.0-glm52-sm80镜像软链到 GLM-5.3 权重。保底逃生门,不解决版本迁移本身。 -
方案 C(最终采用) :最小化干预 ------不动 C++ op 本身,只在 0.28.0
sparse_attn_indexer.py的 prefill/decode 两个 deep_gemm 调用点加is_deep_gemm_supported()门控:deep_gemm 平台路径原样保留(SM90+ 用户零影响),sm80 回退到已 vendor 的 Triton logits 内核(uint8 + LUT 解码 fp8,因为 sm80 连 Triton 的 fp8e4nv 都没有)。 -
方案 C 的改动只有两个调用点的 if/else,sm90+ 行为完全不变,风险最小,验证也最快。实际代码形态(decode 侧同理):
python
# sparse_attn_indexer.py prefill 调用点
if is_deep_gemm_supported():
logits = fp8_fp4_mqa_logits(...) # SM90+ 原路径,原样保留
elif not is_deep_gemm_supported():
# sm80 fallback: deep_gemm requires SM90+
logits = fp8_mqa_logits_triton( # vendor 的 Triton 回退
q_slice_cast, (k_quant_cast, k_scale_cast), weights[...],
cu_seqlen_ks, cu_seqlen_ke, clean_logits=False)
4.2 回退内核的两个工程细节
- vendor 的 Triton logits 内核(
mqa_logits_triton.py)本身有两个值得展开的点,它们决定了回退路径"不仅正确,还尽量快":
① 没有 fp8 硬件类型,就查表解码。 sm80 连 Triton 的 fp8e4nv 类型都没有(坑 8 会再次撞上这一点),内核无法直接做 fp8 算术。解法是把 fp8 当普通 uint8 加载,再用一张 256 项的 **e4m3fn→bf16 查找表(LUT)**还原数值:
python
lut = torch.arange(256, dtype=torch.uint8, device=device) \
.view(torch.float8_e4m3fn).to(torch.bfloat16)
lut[0x7F] = 480.0 # NaN 编码位按饱和值处理
lut[0xFF] = -480.0
查表在 Triton 里就是一次 tl.load(lut_ptr + u),比任何软件位运算解码都便宜,且 LUT 按设备缓存一份。
② autotune 搜索空间按实测裁剪。 内核带 Triton autotune,但搜索空间不是拍脑袋的全量组合,而是按 A100/sm80 实测收敛后收窄的------decode 内核在 (num_heads=32, head_dim=128, block_size=64) 下 num_warps=4 全面占优(其余组合慢 1.5--1.7×),只留 num_warps=4 × num_stages∈{2,4};prefill 内核保留 BLOCK_N∈{32,64,128} 自由度。另外 autotune 的 warmup 形状刻意模拟真实服务形态(chunked prefill 的"小 M 长 N":M=8、N=8192),避免 autotune 在 launch 开销主导的 dummy 网格上选中错误的 tile 配置------这是 Triton autotune 在推理服务里的一个经典暗坑。
5. 验证清单
这一组问题修完后的完整验证点:
- ✅ 后端命中:
Using TRITON_MLA_SPARSE attention backend - ✅ 真实 prefill(935 token)无
forward_mha崩溃 - ✅ prefill 内分页 logits 走 Triton 回退,无 deep_gemm 崩溃
- ✅ 4 并发 × ~30k 上下文全部 OK
小结
- sm80 缺稀疏后端不是"慢一点"的问题,而是"无路可走";社区 PR #47629 + 三处注册点修改即可移植;
- 0.28.0 的兼容断点(计数断言、属性探测、force_mqa)都是小修,但要一个个撞到才知道;
- C++ op 内部的 deep_gemm 依赖用调用点门控 + Triton 回退解决,最小化改动、保住 sm90+ 路径;
- 任何"import 探测"式的平台能力判断都不可信,必须用架构感知的判断。
下一篇:05 问题深拆②:fp8e4nv 编译失败与 inductor 崩溃------torch.compile 恢复之路。