CANN 推理优化实践:LightningIndexer 算子的 TileLang 实现与使用指南
【免费下载链接】cann-recipes-infer本项目针对LLM与多模态模型推理业务中的典型模型、加速算法,提供基于CANN平台的优化样例项目地址: https://gitcode.com/cann/cann-recipes-infer
导读
LightningIndexer 是 CANN 推理优化样例仓 cann-recipes-infer 中面向稀疏注意力(Sparse Attention)场景的关键前置算子,用于从 Query 与 Key 中高效提取每个查询位置最相关的 Top-K 键索引,从而将注意力计算的长度从完整的序列长度压缩到 Top-K。本文以 LightningIndexer 算子说明文档 为主体,结合 TileLang 算子实现、单元测试 与 DeepSeek-V3.2-Exp 算子开发指南,完整介绍算子的功能语义、参数约束、调用方式,并从源码级剖析其 Cube/Vector 双核协作、内存层次设计与增量式 Top-K 排序实现原理。
LightningIndexer 在稀疏注意力中的地位
在大语言模型推理中,FlashAttention 等标准注意力机制的计算复杂度为 O(N²),随序列长度增长迅速失控。Sparse Flash Attention 通过显式的索引张量 index为每个查询 token 指定其需要交互的键/值子集,将注意力计算复杂度降低到 O(N·S)(S 为稀疏关联大小),特别适用于超长序列或结构化稀疏场景。
LightningIndexer 正是该方案中的索引生成算子:它作为 SparseFlashAttention 的前置算子,输入 Query 与 Key,针对每个 Query 输出 Top-K 的 Key/Value 索引,从而稀疏化 Key/Value,将后续注意力机制的计算长度压缩到 Top-K。两者共同构成“先选索引、再算注意力”的两段式稀疏推理管线,其完整配套实现位于 ops/tilelang/ds_v32/ 目录(包含 LightningIndexer.md、SparseFlashAttention.md 两个算子说明及对应实现与测试)。
产品支持情况
| 产品 | 是否支持 |
|---|---|
| Atlas A3 推理系列产品 | √ |
LightningIndexer 算子面向 Atlas A3 推理系列产品,属于推理场景专用算子,其 TileLang 实现利用了 A2 架构 AI Core 中 Cube 核与 Vector 核的硬件特性(A2 的 CV 核默认配比为 1:2),这一点在后续源码解析中会进一步体现。
功能说明与计算公式
算子的核心功能是高效处理索引数据:计算 Query 与 Key 之间的相似度得分,经过 ReLU 激活、分组加权与 Top-K 选取后,输出每个查询位置对应的键索引。
计算公式如下:
$$ Indices(query,key,weights)=Topk(broadcast_vmul(relu(query \cdot key)), weights) $$
公式语义分步拆解为:
- 相似度计算:
query · key计算查询与键之间的点积相似度矩阵; - ReLU 激活:
relu(·)将负相似度置零,从源码实现看(T.copy(..., enable_relu=True))该操作在数据搬运回写过程中利用 Fixpipe 算子原生能力完成,不额外占用计算指令; - 分组加权:
broadcast_vmul(·, weights)将分组权重以广播方式逐元素乘到相似度得分上,等价于对不同分组(Group)的相似度做加权; - Top-K 选取:
Topk(·)对每个查询位置取加权后得分最高的 K 个键,输出其索引。
函数原型
custom.lightning_indexer(query, key, weights) -> Tensor参数说明
说明:query、key、weights 参数维度含义:B(Batch Size)表示输入样本批量大小、S(Sequence Length)表示输入样本序列长度、H(Head Size)表示 hidden 层的大小、N(Head Num)表示多头数、D(Head Dim)表示 hidden 层最小的单元尺寸,且满足 D=H/N。
| 参数 | 类型 | 必选 | 数据格式 | 数据类型 | 说明 |
|---|---|---|---|---|---|
| query | Tensor | 是 | ND | float16 | 查询张量,不支持非连续的 Tensor |
| key | Tensor | 是 | ND | float16 | 键张量,不支持非连续的 Tensor |
| weights | Tensor | 是 | ND | float16 | 分组权重张量,不支持非连续的 Tensor |
三个输入均为必选参数,数据格式统一支持 ND,数据类型统一为float16,且均不支持非连续(non-contiguous)的 Tensor,调用前需确保张量内存连续(可通过.contiguous()处理)。
结合 TileLang 实现 lightning_indexer.py,可以确认算子的实际张量形状设计为:
- Query:
(B, S1, N2, G * D),查询向量按 N2 个头、G 个分组组织,每个分组特征维度为 D; - KEY:
(B, S2, N2, D),键向量集合,每个键为 D 维; - QK_RES(中间结果 workspace):
(B, N2, S1, G, S2),存储 Query 与 Key 之间的相似度得分,数据类型为float(计算精度高于输入); - WEIGHTS:
(B, S1, N2, G),不同分组的加权得分权重; - OUT:
(B, N2, S1, TOP_K),最终输出的 Top-K 索引结果。
返回值说明
| 返回参数 | 类型 | 数据格式 | 说明 |
|---|---|---|---|
| out | Tensor | ND | 公式中的输出,即每个查询位置对应的 Top-K 键索引 |
输出数据类型为int32。从源码看,TileLang 内核内部先以"int"类型完成索引排序与选取(lightning_indexer.py),测试侧再统一.to(torch.int32)与 golden 结果对比。
约束说明
使用该算子时需注意以下约束:
- 该接口支持推理场景下使用;
- 该接口与 PyTorch 配合使用时,需要保证 CANN 相关包与 PyTorch 相关包的版本匹配;
- 参数 key、value 的 N 仅支持 1(即键侧仅支持单头,多头信息通过分组 G 表达);
- 参数 query 中的 D 与 key 的 D 值相等为 128。
最后一条约束对应测试中的实际参数:示例中D=64为测试用例取值,而算子约束声明 D 相等为 128 为通用约束,两者需要区分理解——测试用例用于验证算法正确性,实际部署形状需满足算子约束要求。
调用示例与精度验证
算子仓库提供了可直接运行的测试用例 test_lightning_indexer.py,其测试配置为:
B = 2 # 批次大小 N2 = 1 # KV 注意力头数量 G = 32 # 分组数量(满足 G × N2 = N1) S1 = 512 # Query 序列长度 S2 = 4096 # Key 序列长度 D = 64 # 每组的特征维度 TOP_K = 1024 # 需要返回的最相似结果数量算子实例化
from lightning_indexer import indexer func = indexer(B, N2, G, S1, S2, D, TOP_K, 256, 16, 64, 64, 64)其中后五个整数参数依次对应 TileLang 内核的向量化与分块配置:VECTOR_BASEN=256(向量化处理的键位置基本单位)、VECTOR_BASEG=16(向量化处理的分组基本单位)、BLOCK_M=64(Query 序列分块)、BLOCK_N=64(Key 序列分块)、BLOCK_K=64(特征维分块)。
输入构造
q = torch.randn(B, S1, N2, G, D).half() k = torch.randn(B, S2, N2, D).half() weights = torch.randn(B, S1, N2, G, 1).float() q_npu = q.view(B, S1, N2, -1).npu() k_npu = k.npu() weights_npu = weights.npu() torch.npu.synchronize() npu_out = func(q_npu, k_npu, weights_npu).to(torch.int32) torch.npu.synchronize()Golden 参考实现
测试用例通过 PyTorch 原语复现公式语义,作为正确性基准(test_lightning_indexer.py):
def index_golden(q, k, weights): score_1 = torch.einsum("bsmgd, btmd->bmsgt", q, k) # query · key score_1 = score_1.relu() # relu 激活 score = score_1.permute(0, 2, 1, 3, 4) mul_res = score * weights # 分组加权 reduce_res = torch.sum(mul_res, dim=3) # 分组求和 golden_out = torch.topk(reduce_res, TOP_K, dim=3, largest=True, sorted=True) return score_1.float(), golden_out.indices.to(torch.int32).permute(0, 2, 1, 3)精度判定
测试对输出索引按行做集合差比对(count_mismatches_last_dim,统计两行元素在多集意义上的差异数量),当索引匹配率(1 - mismatches / (B * S1 * N2 * TOP_K)) > 0.99时判定通过并打印Test passed!。这里采用“多集相等”而非“逐位相等”的判定方式,是因为 Top-K 排序结果在得分并列时索引顺序允许存在等价差异,体现了对浮点计算微小误差的工程容忍。
运行方式
在配置好 NPU TileLang 环境后,进入示例目录运行即可:
cd ops/tilelang/ds_v32/examples python3 test_lightning_indexer.py成功后会打印:
Test passed!源码级实现解析:Cube 核与 Vector 核的流水协作
LightningIndexer 的 TileLang 实现(lightning_indexer.py)采用两阶段异构流水设计:第一阶段由 Cube 核完成大规模矩阵乘(相似度计算),第二阶段由 Vector 核完成加权、归约与 Top-K 排序。内核通过T.Kernel(B * N2, is_npu=True)以 batch × 头数粒度并行启动,并利用T.set_cross_flag("FIX", 0)/T.wait_cross_flag(0)实现 Cube 核到 Vector 核的跨流水线同步。
阶段一:Cube 核相似度计算
Cube 核负责relu(query · key)的分块矩阵乘,其内存设计充分利用 NPU 多级存储层次:
- L1 缓存:作为主要数据暂存区,存储当前计算的 Query 与 Key 数据块(
Q_L1、K_L1); - L0C 缓存:作为矩阵乘累加器(
C_L0),承接 Cube 单元的计算结果。
内核通过T.annotate_address手动规划 L1/L0C 地址(Q_L1 从地址 0 开始、K_L1 从地址 16384 开始),确保数据块间无地址冲突并提升访问局部性。以BLOCK_M=128, BLOCK_N=128, BLOCK_K=128为例,Query 子块占用128 × 128 × sizeof(half) = 32768字节,与源码中 K_L1 的偏移 16384(即 BLOCK_M × BLOCK_K × 2 字节的一半)呼应了地址规划的精细度。
核内采用四重循环结构分块计算:
- 外层循环(注意力头 n2):遍历每个注意力头;
- 第二层循环(分组 g):遍历 Query 向量的每个分组,与完整 Key 进行匹配;
- 第三层循环(Query 序列分块 m):将 S1 按 BLOCK_M 分块;
- 内层循环(Key 序列分块 n):将 S2 按 BLOCK_N 分块。
每个内层迭代执行:T.copy将 Query/Key 块搬入 L1 →T.gemm_v0(Q_L1, K_L1, C_L0, transpose_B=True, init=True)执行矩阵乘 →T.copy将 L0C 结果写回全局 QK_RES,并利用enable_relu=True在搬运路径上完成 ReLU。计算完成后通过T.set_cross_flag("FIX", 0)通知 Vector 核开始消费中间结果。
阶段二:Vector 核加权归约与 Top-K
Vector 核接收 QK_RES((B, N2, S1, G, S2))与 WEIGHTS,完成“加权 → 分组归约 → Top-K 排序”的索引生成:
- 并行负载均衡:总任务数
N2 * S1平均分配给两个 Vector 核(total_process_num // 2),每个核通过vid计算自己的处理区间s1_start_idx至s1_end_idx; - 加权累加:按
VECTOR_BASEG × VECTOR_BASEN分块加载相似度与权重到 UB,逐行执行T.tile.mul加权,再通过T.tile.add跨分组累加到reduce_tmp_ub; - 分组归约:
T.reduce_sum(reduce_tmp_ub, reduce_g_ub, 0)沿分组维度求和,得到每个查询位置相对每个键的最终得分; - Top-K 增量归并排序:这是算子最核心的算法设计,采用增量式归并排序策略:
- 对每个
VECTOR_BASEN大小的键块调用T.tile.sort排序,并用T.tile.gather_mask(..., "P1010")提取排序索引; - 以
merge_sort_times = TOP_K // VECTOR_BASEN确定归并块数,将各块排序结果写入topk_global_ub1; - 每凑齐
merge_sort_times个块后执行T.tile.merge_sort归并;首次归并直接产生结果,后续归并通过T.tile.topk将结果集裁剪回 TOP_K 大小,从而全程只需维护 TOP_K 规模的全局候选集,避免了对全量 S2 得分排序的开销;
- 对每个
- 结果写回:
T.tile.cast(output_ub, topk_global_ub1_flat, "CAST_ROUND", TOP_K)将排序索引四舍五入转回 int 类型,写回OUT[cid, n2_id, s1_id, 0:TOP_K]。
该实现充分体现了 TileLang“调度空间与数据流解耦”的设计理念:开发者只描述数据流(搬运、计算、排序原语),线程绑定、L1/L0 布局、流水同步等底层优化由编译器结合 NPU 硬件自动完成。
算子在模型推理中的接入
在模型侧,DeepSeek-V3.2-Exp 推理实现中提供了配套的 Indexer 模块,负责将 LightningIndexer 算子(以及 PyTorch 侧torch_npu相关接口)接入整网推理流程,并通过custom_params(如enable_multi_streams)与执行模式(ge_graph/npugraph_ex)配合调度。这表明该算子已具备从“单算子验证”到“整网推理”的落地路径,读者可在模型推理指南(deepseek_v3.2_exp_inference_guide.md)中查看完整接入方式。
运行环境准备
运行该 TileLang 算子需要基于 Atlas A3 推理系列产品的 NPU 环境,并预装 NPU TileLang 及其依赖。环境准备可参考 ops/tilelang/README.md:
- 获取运行镜像:仓库提供预装 TileLang 代码仓及其全部依赖的 docker 镜像(含昇腾 CANN 运行环境),镜像内已支持全量基础 AscendC 算子 API,可直接运行代码;
- 拉起容器:将 NPU 设备(
/dev/davinci*、/dev/davinci_manager等)与宿主机驱动、数据目录挂载进容器; - 设置环境变量:
source /usr/local/Ascend/ascend-toolkit/set_env.sh- (可选)自行安装 TileLang:从源码编译安装
tilelang-ascend,执行bash install_ascend.sh后source set_env.sh完成环境配置。
算子在 JIT 编译时会自动设置ACL_OP_INIT_MODE=1并关闭 TileLang 缓存(tilelang.disable_cache()),确保每次按当前输入维度与 NPU 硬件状态即时生成并编译适配的 AscendC 代码(见 lightning_indexer.py)。
总结
LightningIndexer 是 CANN 推理优化样例中“稀疏注意力”技术路线的重要一环,通过“相似度计算(Cube)→ 加权归约(Vector)→ 增量 Top-K 归并排序”的两段式设计,在 Atlas A3 推理系列产品上高效完成了索引生成任务。本文从算子语义、参数约束、测试验证到 TileLang 源码实现逐层展开,读者既可以将其作为单算子开发与调用的实战参考,也可以结合 SparseFlashAttention 与 DeepSeek-V3.2-Exp TileLang 算子开发指南 进一步理解其在超长序列稀疏注意力推理中的完整应用链路。
【免费下载链接】cann-recipes-infer本项目针对LLM与多模态模型推理业务中的典型模型、加速算法,提供基于CANN平台的优化样例项目地址: https://gitcode.com/cann/cann-recipes-infer
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考