2026/10/11 4:04:00

TileRT稀疏MLA实现原理:FlashSparseMLA与TopK近似索引让长上下文不再拖慢解码

TileRT稀疏MLA实现原理:FlashSparseMLA与TopK近似索引让长上下文不再拖慢解码 【免费下载链接】TileRTTile-Based Runtime for Ultra-Low-Latency LLM Inference项目地址https://gitcode.com/gh_mirrors/ti/TileRT点击查看免费下载TileRT 是一款面向超低延迟 LLM 推理的 Tile 级运行时。本文拆解它最有代表性的稀疏 MLAMulti-head Latent Attention实现用 TopK 近似索引先筛出关键位置再用 FlashSparseMLA 只对这些位置做注意力让百万 token 的长上下文不再拖慢逐 token 解码。为什么长上下文解码会拖慢LLM 解码是逐 token 串行的每生成一个 token都要对所有层的 KV Cache 做一遍注意力。全量注意力的计算量和显存读取量都随上下文长度线性增长——上下文从 1K 涨到 1M注意力部分要读的数据量就是原来的 1000 倍。TileRT 的解法来自 DeepSeek-V3.2 / GLM-5 的稀疏注意力思想DSA不要每个历史位置都看用一个轻量索引器给所有历史 token 打分只取最相关的 2048 个位置index_topk 2048见 tilert/models/deepseek_v3_2/model_args.py。KV Cache 本身也小MLA 把 KV 压成低秩表示kv_lora_rank 512配合独立的低维索引缓存index_head_dim 128单 token 的缓存体积远低于传统 GQA。于是扫 100 万个位置变成打分 100 万个位置很便宜 只读 2048 个位置很快。三步流水线从打分到稀疏注意力整条链路由三个自定义算子完成核心源码在 tilert/models/deepseek_v3_2/ops/ 目录下。第 1 步Indexer 计算相关性得分解码时当前 token 会生成一组索引 query与每个历史位置的索引缓存点积得到一条覆盖全部上下文的得分向量IDX_LOGITS形状约为[1, seq, max_seq_lenpad]。这一步由 sparse_index.py 中的sparse_index完成输入是 bf16 的 q/kv/weights得分以 fp32 累加索引头配置为 64 头 × 128 维——维度很小所以即使上下文达到百万级打分本身也很廉价。更进一步sparse_index_topk 把打分 取 TopK融合进一个 kernel省掉了中间得分向量的往返读写。第 2 步TopK 近似索引筛出 2048 个位置在几十万上百万个得分中挑出前 2048 名如果对每个 token、每层都做一次全量排序排序本身就会成为新瓶颈。TileRT 在 topk.py 中提供了两条路径topk_approximate近似 TopK kernel固定topk2048用分桶近似的思路快速定位候选位置延迟稳定且极低适合解码热路径topk_accurate精确 TopK支持topk ∈ {512, 1024, 2048}并带ratio压缩因子用于 token 维度的下采样。两者由 TopK 模块统一封装use_approximate开关一键切换同时保留了 PyTorch 原生topk的参考实现用于数值校验。第 3 步FlashSparseMLA 只算被选中的位置拿到 2048 个位置索引后真正的 MLA 注意力就只对它们计算。flash_sparse_mla.py 中的flash_sparse_mla有几个关键设计分片计算split-KKV 方向按split_size64切块2048 个位置正好分成最多 32 个 split各 split 独立算出部分输出output_acc和 log-sum-exp 归一化项lse_accLSE 归并融合由 FlashSparseMLACombine 模块把 32 个部分结果按对数归一化精确归并得到与全量注意力数值一致的结果bf16mma计算内核原生支持 MTP算子直接接受seqlen4的批量输入解码 1 token 3 个草稿 token稀疏 MLA 与多 token 预测零成本叠加模块同时提供golden_forward参考实现方便在 Tile 级 kernel 与标准 einsum 实现之间对拍验证。多 GPU 如何保持索引步调一致TileRT 把 128 个注意力头切到 8 张 GPU 上并行但 TopK 选出的位置索引必须在所有卡上一致。协作方式很直接GPU 0 使用 SparseSelectMlaV2额外挂有索引器投影ProjxWis负责算分、选 TopK 并持有三个缓存——ki_cache索引 KV、kv_cache低秩 KV、pe_cache位置编码见 get_cache_vars其余 GPU 使用 PureMlaV2通过 broadcast_selected_token_ids / receive_selected_token_ids 两个 P2P 算子把[1, S, 2048]的 int32 索引以写缓冲区 同步 flag的方式从 GPU 0 直达对端再对各自负责的头部分片做稀疏 MLA最后经 AllReduce 合并。整层逻辑61 层的 DSA 堆叠、临时变量布局IDX_SCORES → IDX_LOGITS → IDX_SELECTS定义在 modules/dsa.pyGLM-5 复用了同一套算子见 tilert/models/glm_5/_dsa_v32/ops/。实测效果1M 上下文的 TPS 还剩多少稀疏 MLA 的收益直接体现在输入越长、衰减越慢。下图是 GLM-5.2/5.3-FP8 在 8× MI350X 上的基准输出固定 1K输入从 1K 拉到 1M——由于注意力只读 2048 个位置生成速度仅从 314 tok/s 缓降到 207 tok/s而叠加 MTP平均接受长度 3.2后 1K 输入可达 648 tok/s。早期版本 GLM-5.1 在 8× B200 上的数据同样验证了这一趋势192K 输入下无 MTP 仍有 169 tok/s开启 MTP 后升至 352 tok/s长尾衰减曲线非常平缓。源码导航功能文件稀疏索引打分 / 打分TopK 融合ops/sparse_index.py近似/精确 TopKops/topk.pyFlashSparseMLA 与 LSE 归并ops/flash_sparse_mla.py索引的跨卡广播/接收ops/broadcast_selected_token_ids.py稀疏/纯 MLA 层组装modules/mla_v2.pyDSA 全模型堆叠modules/dsa.py模型超参topk、序列长度等model_args.py小结TileRT 的稀疏 MLA 实现可以浓缩成一句话用廉价的低维 Indexer 给全部历史打分用 TopK 近似索引以固定 2048 的代价锁定关键位置再用分片 LSE 归并的 FlashSparseMLA 精确算出稀疏注意力。配合 MLA 的低秩 KV 压缩与跨卡索引广播注意力成本与上下文长度基本解耦——这正是 TileRT 能让百万 token 会话保持每秒数百 token 解码速度的关键所在。赞分享【免费下载链接】TileRTTile-Based Runtime for Ultra-Low-Latency LLM Inference项目地址https://gitcode.com/gh_mirrors/ti/TileRT点击查看免费下载相关推荐MiniCPM 2.0 系列技术详解128k 长上下文、MoE 与稀疏化推理实践MiniCPM 2.0 系列技术详解128k 长上下文、MoE 与稀疏化推理实践 MiniCPM 2.0 是 OpenBMB 开源社区在 MiniCPM 1.大模型本地部署模型量化微调LoRA工具调用openBMBAscendKafka稀疏索引实战algorithm-pattern数据索引原理完全讲解Kafka稀疏索引实战algorithm pattern数据索引原理完全讲解 algorithm pattern 是一套算法模板刷题仓库提供最科学的刷题方式教程TileLang 实现 DeepSeek V3.2 稀疏 MLA 全流程Lightning Indexer、Top-k Selector 与稀疏注意力内核实战TileLang 实现 DeepSeek V3.2 稀疏 MLA 全流程Lightning Indexer、Top k Selector 与稀疏注意力内核实战编译器编程语言高性能计算人工智能深度学习创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考