CANN 推理优化实践:LightningIndexer 算子的 TileLang 实现与使用指南
2026/9/18 20:00:41 网站建设 项目流程

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) $$

公式语义分步拆解为:

  1. 相似度计算query · key计算查询与键之间的点积相似度矩阵;
  2. ReLU 激活relu(·)将负相似度置零,从源码实现看(T.copy(..., enable_relu=True))该操作在数据搬运回写过程中利用 Fixpipe 算子原生能力完成,不额外占用计算指令;
  3. 分组加权broadcast_vmul(·, weights)将分组权重以广播方式逐元素乘到相似度得分上,等价于对不同分组(Group)的相似度做加权;
  4. 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。

参数类型必选数据格式数据类型说明
queryTensorNDfloat16查询张量,不支持非连续的 Tensor
keyTensorNDfloat16键张量,不支持非连续的 Tensor
weightsTensorNDfloat16分组权重张量,不支持非连续的 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 索引结果。

返回值说明

返回参数类型数据格式说明
outTensorND公式中的输出,即每个查询位置对应的 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_L1K_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 字节的一半)呼应了地址规划的精细度。

核内采用四重循环结构分块计算:

  1. 外层循环(注意力头 n2):遍历每个注意力头;
  2. 第二层循环(分组 g):遍历 Query 向量的每个分组,与完整 Key 进行匹配;
  3. 第三层循环(Query 序列分块 m):将 S1 按 BLOCK_M 分块;
  4. 内层循环(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 排序”的索引生成:

  1. 并行负载均衡:总任务数N2 * S1平均分配给两个 Vector 核(total_process_num // 2),每个核通过vid计算自己的处理区间s1_start_idxs1_end_idx
  2. 加权累加:按VECTOR_BASEG × VECTOR_BASEN分块加载相似度与权重到 UB,逐行执行T.tile.mul加权,再通过T.tile.add跨分组累加到reduce_tmp_ub
  3. 分组归约T.reduce_sum(reduce_tmp_ub, reduce_g_ub, 0)沿分组维度求和,得到每个查询位置相对每个键的最终得分;
  4. 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 得分排序的开销;
  5. 结果写回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:

  1. 获取运行镜像:仓库提供预装 TileLang 代码仓及其全部依赖的 docker 镜像(含昇腾 CANN 运行环境),镜像内已支持全量基础 AscendC 算子 API,可直接运行代码;
  2. 拉起容器:将 NPU 设备(/dev/davinci*/dev/davinci_manager等)与宿主机驱动、数据目录挂载进容器;
  3. 设置环境变量
source /usr/local/Ascend/ascend-toolkit/set_env.sh
  1. (可选)自行安装 TileLang:从源码编译安装tilelang-ascend,执行bash install_ascend.shsource 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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询