CANN ops-transformer LightningIndexerV2 算子使用指南:基于 QK 相关性 Top-k 稀疏索引的 PyTorch API 实战
【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer
导读
lightning_indexer是 CANN ops-transformer 中LightningIndexerV2算子的 PyTorch 调用入口,用于在 Attention 计算之前,为每个 Query token 从 Key 序列中快速挑选出相关性最高的 Top-k 个位置,从而把后续 Attention 从全量计算收敛为稀疏计算。本文以 torchapi_lightning_indexer.md 为核心,完整介绍该接口的算子原理、函数原型、全部参数与返回值语义、约束条件,并结合仓库内的算子定义、Infershape、Tiling、metadata 结构与 golden 测试实现,深入解析其背后的计算与负载均衡机制,最后给出可复制的单算子模式与 TorchAir(aclgraph)图模式调用示例。
阅读完本文,你将掌握:LightningIndexerV2 的完整计算流程与公式、metadata 前置算子与主算子如何配合、BSND/TND/PA_BBND 三种 layout 组合规则、cmp_ratio 压缩与 mask_mode 稀疏模式的协同用法,以及在不同 NPU 产品上的参数差异与规避策略。
产品支持情况
lightning_indexer接口(LightningIndexerV2 算子)在不同硬件产品上的支持情况如下:
| 产品 | 支持情况 |
|---|---|
| Ascend 950PR / Ascend 950DT | 支持 |
| Atlas A3 训练系列产品 / Atlas A3 推理系列产品 | 支持 |
| Atlas A2 训练系列产品 / Atlas A2 推理系列产品 | 支持 |
| Atlas 200I/500 A2 推理产品 | 不支持 |
| Atlas 推理系列产品 | 不支持 |
| Atlas 训练系列产品 | 不支持 |
该算子当前仅面向推理(Inference)场景,并同时支持单算子模式与 TorchAir(aclgraph)图模式两种调用方式。
功能说明:算子在做什么
接口分工
该功能由两个配套接口组成:
lightning_indexer_metadata(前置接口):用于生成一个任务列表(metadata),包含每个 AI Core 上 Attention 计算任务的起止点的 batch、head,以及 Q 和 K 的分块索引,供后续lightning_indexer算子使用。它负责的是负载均衡计算——先把任务按 AI Core 切分好,主算子才能按图索骥地并行执行。lightning_indexer(主接口):基于一系列操作得到每一个 token 对应的 Top-k 个位置,输出稀疏的索引与取值,作为后续 Attention 的输入。
三步计算流程
- 将某个 token 对应的输入参数
q($Q_{index}\in\R^{g\times d}$)乘以给定上下文k($K_{index}\in\R^{S_{k}\times d}$),得到相关性。 - 通过激活函数 $ReLU$ 过滤无效负相关信号后,得到当前 Token 与所有前序 Token 的相关性分数向量。
- 将其与权重系数
w($W$)相乘后,沿 g 的方向,选取前 $Top-k$ 个索引值得到输出sparseIndices,并输出对应的sparseValues,作为 Attention 的输入。
计算公式
$$ Top-k \left{ \left[ 1 \left] \mathop{{}}\nolimits_{{1 \times \text{ }g}}\text{@} \left[ \left( W\text{@} \left[ 1 \left] \mathop{{}}\nolimits_{{1\text{ } \times \text{ }S\mathop{{}}\nolimits_{{k}}}} \left) \text{ } \odot \text{ }ReLU \left( Q\mathop{{}}\nolimits_{{index}}\text{@}K\mathop{{}}\nolimits_{{T}}^{{index}} \left) \left] \right} \right. \right. \right. \right. \right. \right. \right. \right. \right. \right. $$
该公式与仓库 LightningIndexerV2 README 中的算子描述一致,可互相印证。
从实现角度看,相关性计算(QK^T 矩阵乘)与 Top-k 选择分别在 Cube 与 Vector 单元上执行:仓库 kernel 侧按架构划分了 arch22 实现 与 arch35 实现,其中 arch35 侧还包含 lightning_indexer_v2_topk.h 等专用 Top-k 向量实现,可以推断这是"Cube 计算相关性 + Vector 执行 Top-k 筛选"的异构流水设计。
函数原型
调用lightning_indexer接口之前,必须先调用前置接口lightning_indexer_metadata,完成负载均衡计算(metadata 生成)。两个接口的原型如下。
lightning_indexer_metadata
cann_ops_transformer.lightning_indexer_metadata( num_heads_q, num_heads_k, head_dim, topk, *, cu_seqlens_q=None, cu_seqlens_k=None, seqused_q=None, seqused_k=None, cmp_residual_k=None, batch_size=0, max_seqlen_q=-1, max_seqlen_k=-1, layout_q="BSND", layout_k="BSND", mask_mode=0, cmp_ratio=1 ) -> Tensorlightning_indexer
cann_ops_transformer.lightning_indexer( q, k, w, topk, *, cu_seqlens_q=None, cu_seqlens_k=None, seqused_q=None, seqused_k=None, cmp_residual_k=None, block_table=None, output_idx_offset=None, metadata=None, max_seqlen_q=-1, layout_q="BSND", layout_k="BSND", mask_mode=0, cmp_ratio=1, return_value=0 ) -> (Tensor, Tensor)上述原型与仓库 torch_extension/lightning_indexer.py 中通过OpBuilder注册的算子 schema 完全一致,其背后对应 C++ 侧 lightning_indexer_v2_def.cpp 中定义的LightningIndexerV2算子原型。
参数说明
在展开参数表之前,先统一说明维度符号的含义:**b(batch size)**表示输入样本批量大小;q_s / k_s表示 q / k 的输入样本序列长度;q_n / k_n表示 q / k 的多头数;**d(head dim)**表示注意力头的维度;q_t / k_t表示 q / k 输入样本序列长度的累加和。参数 q 中的 d 和参数 k 中的 d 值相等,当前仅支持 128。
lightning_indexer_metadata 参数
| 参数名 | 参数类型 | 可选/必选 | 描述 | 数据类型 | 维度(shape) |
|---|---|---|---|---|---|
| num_heads_q | int | 必选 | 表示 q 的 head 个数。 | int64 | - |
| num_heads_k | int | 必选 | 表示 k 的 head 个数,当前仅支持 1。 | int64 | - |
| head_dim | int | 必选 | 表示注意力头的维度,当前仅支持 128。 | int64 | - |
| topk | int | 必选 | 表示为每个 q token 保留的 Key token 索引个数,当前支持 [1, 8192]。 | int64 | - |
| cu_seqlens_q | Tensor | 可选 | 表示不同 batch 中 q 的有效 Sequence Length 的累加前缀和,仅 layout_q 为 TND 场景下必传,第一个值固定为 0。数据格式为 ND,支持非连续的 Tensor。 | int32 | (b+1, ) |
| cu_seqlens_k | Tensor | 可选 | 表示不同 batch 中 k 的有效 Sequence Length 的累加前缀和,仅 layout_k 为 TND 场景下必传,第一个值固定为 0。数据格式为 ND,支持非连续的 Tensor。 | int32 | (b+1, ) |
| seqused_q | Tensor | 可选 | 表示不同 batch 中 q 实际参与运算的 Sequence Length。数据格式为 ND,支持非连续的 Tensor。 | int32 | (b, ) |
| seqused_k | Tensor | 可选 | 表示不同 batch 中 k 实际参与运算的 Sequence Length。数据格式为 ND,支持非连续的 Tensor。 | int32 | (b, ) |
| cmp_residual_k | Tensor | 可选 | 表示不同 batch 中 cmp_kv 压缩前 Sequence Length 除以 cmp_ratio 的余数,配合 cmp_ratio 实现 cmp_kv 部分的 mask 和负载计算。cmp_ratio 不为 1 且 mask_mode 为 3 场景下必传。数据格式为 ND,支持非连续的 Tensor。 | int32 | (b, ) |
| batch_size | int | 可选 | 表示 batch 数量,默认值为 0。 | int64 | - |
| max_seqlen_q | int | 可选 | 表示 q 的最长 Sequence Length,-1 表示任意可能长度,默认值为 -1。 | int64 | - |
| max_seqlen_k | int | 可选 | 表示 k 的最长 Sequence Length,-1 表示任意可能长度,默认值为 -1。 | int64 | - |
| layout_q | str | 可选 | 表示 q 的排列格式,支持 BSND、TND,默认值为 BSND。 | string | - |
| layout_k | str | 可选 | 表示 k 的排列格式,支持 BSND、TND、PA_BBND,默认值为 BSND。 | string | - |
| mask_mode | int | 可选 | 表示 sparse 模式,0 表示 No mask,3 表示 rightDownCausal 模式,默认值为 0。 | int64 | - |
| cmp_ratio | int | 可选 | 表示 k 的压缩率,取值范围 [1, 128],默认值为 1,表示无压缩。 | int64 | - |
lightning_indexer 参数
| 参数名 | 参数类型 | 可选/必选 | 描述 | 数据类型 | 维度(shape) |
|---|---|---|---|---|---|
| q | Tensor | 必选 | 公式中的输入 Q。不支持空 tensor。数据格式为 ND。 | bfloat16、float16 | layout_q 为 BSND 时 shape 为 (b,q_s,q_n,d);layout_q 为 TND 时 shape 为 (q_t,q_n,d) |
| k | Tensor | 必选 | 公式中的输入 K。不支持空 tensor。数据格式为 ND,支持非连续的 Tensor(仅 PA_BBND 场景下 0 轴支持非连续)。 | bfloat16、float16 | layout_k 为 BSND 时 shape 为 (b,k_s,k_n,d);layout_k 为 TND 时 shape 为 (k_t,k_n,d);layout_k 为 PA_BBND 时 shape 为 (block_num,block_size,k_n,d) |
| w | Tensor | 必选 | 公式中的输入 W。不支持空 tensor。数据格式为 ND。 | float | layout_q 为 BSND 时 shape 为 (b,q_s,q_n);layout_q 为 TND 时 shape 为 (q_t,q_n) |
| topk | int | 必选 | topK 阶段需要保留的 Key token 索引数量,当前支持 [1, 8192]。 | int64 | - |
| cu_seqlens_q | Tensor | 可选 | 当前 batch 及前序 batch 中 q 的有效 token 数的累加和。仅 layout_q 为 TND 场景下必传,第一个值固定为 0。数据格式为 ND。 | int32 | (b+1,) |
| cu_seqlens_k | Tensor | 可选 | 当前 batch 及前序 batch 中 k 的有效 token 数的累加和。仅 layout_k 为 TND 场景下必传,第一个值固定为 0。数据格式为 ND。 | int32 | (b+1,) |
| seqused_q | Tensor | 可选 | 不同 batch 中 q 的真实使用长度。数据格式为 ND。 | int32 | (b,) |
| seqused_k | Tensor | 可选 | 不同 batch 中 k 的真实使用长度。数据格式为 ND。 | int32 | (b,) |
| cmp_residual_k | Tensor | 可选 | 表示 k 压缩前 token 数量除以 cmp_ratio 的余数。需要在 mask_mode 等于 3、cmp_ratio 不等于 1 的场景下使用。数据格式为 ND。 | int32 | (b,) |
| block_table | Tensor | 可选 | 表示 PageAttention 中 KV 存储使用的 block 映射表。不支持空 tensor。数据格式为 ND。 | int32 | (b, k_s_max/block_size) |
| output_idx_offset | Tensor | 可选 | 表示 topK 结果输出索引所需要加上的偏移。值必须大于等于 0,加上偏移后 topK index 不能超过 int32 最大值。数据格式为 ND。 | int32 | layout_q 为 BSND 时 shape 为 (b,q_s,k_n);layout_q 为 TND 时 shape 为 (q_t,k_n) |
| metadata | Tensor | 可选 | 由 lightning_indexer_metadata 得到的分核信息,包含使用核数、分块大小以及每个核处理数据的起始点等内容。不支持空 tensor。数据格式为 ND。 | int32 | (1024,) |
| max_seqlen_q | int | 可选 | q 的最大序列长度。-1 表示任意可能长度,默认值为 -1。 | int64 | - |
| layout_q | str | 可选 | 用于标识输入 q 的数据排布格式,支持 BSND、TND,默认值为 BSND。 | string | - |
| layout_k | str | 可选 | 用于标识输入 k 的数据排布格式,支持 BSND、TND、PA_BBND,默认值为 BSND。 | string | - |
| mask_mode | int | 可选 | 表示 mask 的模式,0 代表 defaultMask 模式,3 代表 rightDownCausal 模式,默认值为 0。 | int64 | - |
| cmp_ratio | int | 可选 | 用于稀疏计算,表示 k 的压缩倍数。支持 1-128,默认值为 1。 | int64 | - |
| return_value | int | 可选 | 代表是否需要返回 Indices 对应的 Values 值。0 代表不返回,1 代表返回值,默认值为 0。 | int64 | - |
返回值说明
lightning_indexer_metadata
| 参数名 | 参数类型 | 可选/必选 | 描述 | 数据类型 | 维度(shape) |
|---|---|---|---|---|---|
| metadata | Tensor | 必选 | 每个 AI Core 的 Attention 计算任务的 batch、head、以及 Q 和 K 的分块的索引。数据格式为 ND,不支持非连续的 Tensor。 | int32 | shape 为 (1024,) |
lightning_indexer
| 参数名 | 参数类型 | 可选/必选 | 描述 | 数据类型 | 维度(shape) |
|---|---|---|---|---|---|
| sparse_indices | Tensor | 必选 | 公式中的 Indices 输出。不支持空 tensor。无效部分填 -1。数据格式为 ND。 | int32 | layout_q 为 BSND 时 shape 为 (b,q_s,k_n,topk);layout_q 为 TND 时 shape 为 (q_t,k_n,topk) |
| sparse_values | Tensor | 条件输出 | 公式中的 Indices 对应的 Values 输出。当 return_value 为 1 时输出对应值;当 return_value 为 0 时输出 shape 为 [0] 的空 tensor。无效部分填 -inf。数据格式为 ND。 | float | layout_q 为 BSND 时 shape 为 (b,q_s,k_n,topk);layout_q 为 TND 时 shape 为 (q_t,k_n,topk);return_value 为 0 时 shape 为 (0,) |
输出 shape 的推导依据
输出的sparse_indicesshape 由算子注册的 Infershape 实现静态推导:lightning_indexer_v2_infershape.cpp 中按 layout_q 分支计算:
- layout_q 为 BSND 时,
sparse_indicesshape 为(q.dim0, q.dim1, k.dim2, topk),即(b, q_s, k_n, topk); - layout_q 为 TND 时,shape 为
(q.dim0, k 的 N 维, topk),其中 k 的 N 维在 layout_k 为 PA_BBND 时取第 2 维、否则取第 1 维,即(q_t, k_n, topk); sparse_values在return_value为真时与sparse_indices同 shape,否则 shape 为(0,)。
PyTorch 侧的 Meta 实现 lightning_indexer.py 也以同样的规则计算输出 shape,两者保持一致。
约束说明
通用约束
- 该接口支持推理场景下使用。
- 该接口支持单算子模式和 TorchAir(aclgraph)图模式调用。
- lightning_indexer_metadata 接口需与 lightning_indexer 算子配套使用。
- b(batch)表示输入样本批量大小。
- 参数 cu_seqlens_q、cu_seqlens_k 要求其值为当前 batch 与前序 batch 有效 token 数的累加值,第一个元素必须为 0,且后一个元素的值必须大于等于前一个元素的值。
- 参数 seqused_q、seqused_k 要求其值表示每个 batch 中的有效 token 数。
- 参数 cmp_residual_k 需满足 cmp_residual_k[i] < cmp_ratio。
- mask_mode 所表示的 mask 模式的详细介绍见 sparse_mode 参数说明。
- pa_kv_cache 支持 0 轴非连续;pa_block_size 支持 [16, 1024],且是 16 的倍数。
- 参数 q、k 的数据类型应保持一致。
- 该接口的 TopK 排序过程对 NaN 排序是未定义行为。
- 当 layout_q 为 BSND 时,不支持传入 cu_seqlens_q;当 layout_k 为 BSND 或 PA_BBND 时,不支持传入 cu_seqlens_k。
- 当传入的 cmp_ratio > 1 且 mask_mode = 3 时,必须传入 cmp_residual_k,其余情况不传入。
- sparse_indices 无效部分填 -1;sparse_values 无效部分填 -inf。
按产品的差异化约束
Atlas A3 训练/推理系列产品、Atlas A2 训练/推理系列产品:
- topk 取值范围当前仅支持 [1, 2048],以及 3072、4096、5120、6144、7168、8192。
- 当前 lightning_indexer 接口不支持 seqused_q、output_idx_offset、max_seqlen_q 功能,不建议传入这些参数。
- 支持 num_heads_q = 1~64、q_n = 1~64。
- layout_k 支持 BSND、TND、PA_BBND;非 PA_BBND 场景下 layout_q 和 layout_k 必须一致。
- layout_k 为 PA_BBND 时必须传入 seqused_k;为 BSND 或 TND 时可选传入。
- cmp_ratio 支持 [1, 128]。
- 支持 return_value 功能,0 代表不返回对应值,1 代表返回 FLOAT32 类型的 sparse_values。
Ascend 950PR/Ascend 950DT:
- topk 取值范围当前仅支持 [1, 8192]。
- 支持 num_heads_q = 1~64、q_n = 1~64。
- cmp_ratio 支持 [1, 128]。
- 当传入 output_idx_offset 时,支持大于等于 0 的索引偏移值;且应满足约束:加上传入的索引偏移值后,得到的 sparseIndice 值不超过 INT32 的最大值。
- 当 layout_q 为 TND 时,必须传入 cu_seqlens_q,如果也传入 seqused_q,应保证由 seqused_q 传入的各个 batch 的 q 长度不超过根据 cu_seqlens_q 计算出的各个 batch 的 q 序列长度。当某个 batch 由 seqused_q 传入的 q 序列长度 seqlen1 小于由 cu_seqlens_q 计算出的 q 长度 seqlen2 时,会启用 TND Padding 功能,将该 batch 的 seqlen2 与 seqlen1 的差值部分的 q 输出的 sparse_indices 和 sparse_values 全部置为无效值。部分长序列场景下,如果需要填充的无效数据过多,由于硬件限制可能会导致 aicore 执行超时,可以通过 (seqlen2 - seqlen1) * topk 来计算需要填充的数据量,建议将这个数据量控制在 4 亿以内。
- 参数 metadata 必须传入,shape 为 (1024,)。
特性参数组
| 特性参数组 | 参数字段名称 |
|---|---|
| 公共参数组 | q、k、w、metadata、output_idx_offset、topk、layout_q、layout_k、sparse_indices、sparse_values |
| Mask 参数组 | mask_mode |
| SeqLens 参数组 | cu_seqlens_q、cu_seqlens_k、seqused_q、seqused_k、max_seqlen_q |
| 稀疏压缩参数组 | cmp_ratio、cmp_residual_k |
| Paged Attention 参数组 | block_table |
基准信息说明:参数校验规则详解
本算子在 host 侧具备完善的参数校验框架(见 op_host/checkers 目录下的 base/compression/seq_len/paged_attention 等 checker),以下为基准校验信息。
公共参数组
入参为空的场景处理:
- 空 Tensor 指必选输入和输出的 shape size 为 0,即有任意轴为 0。
- 触发空 tensor 的用例将全部拦截报错。
q、k、sparse_indices、sparse_values 校验:
| 参数 | 单参数校验 | 存在性校验 | 一致性校验 | 特性交叉校验 |
|---|---|---|---|---|
| q | tensor_type 支持 BFLOAT16 和 FLOAT16;BSND -> (b, q_s, q_n, d);TND -> (q_t, q_n, d) | 必须存在 | q、k 的数据类型需相同;Layout 校验规则见 layout 匹配关系表 | 轴校验:65536 > b > 0;q_t > 0;k_t > 0;q_n > 0;k_n = 1;q_s > 0;k_s > 0;d = 128 |
| k | tensor_type 支持 BFLOAT16 和 FLOAT16;BSND -> (b, k_s, k_n, d);TND -> (k_t, k_n, d);PA_BBND -> (num_blocks, block_size, k_n, d);1024 >= block_size >= 16,block_size % 16 == 0 | 必须存在 | 同上 | 同上 |
| sparse_indices | tensor_type 支持 INT32;layout_q 为 BSND 时 shape 为 (b, q_s, k_n, topk);layout_q 为 TND 时 shape 为 (q_t, k_n, topk) | 必须存在 | 同上 | 同上 |
| sparse_values | tensor_type 支持 FLOAT32;layout_q 为 BSND 时 shape 为 (b, q_s, k_n, topk);layout_q 为 TND 时 shape 为 (q_t, k_n, topk) | 必须存在 | 同上 | 同上 |
layout 匹配关系表:
| layout_q | layout_k | layout_out |
|---|---|---|
| BSND | BSND / PA_BBND | BSND |
| TND | TND / PA_BBND | TND |
即:输出 layout 与 layout_q 保持一致;layout_k 为 PA_BBND 时可与任意 layout_q 组合(Paged Attention 场景),否则 layout_k 必须与 layout_q 相同。
metadata 校验:
| 参数 | 单参数校验 | 存在性校验 | 一致性校验 | 特性交叉校验 |
|---|---|---|---|---|
| metadata | tensor_type 仅支持 INT32;shape 由 lightning_indexer_v2_metadata 动态计算;当前不支持不传入,未传入将发出拦截报警 | 可选参数 | 无 | 传入时需与 lightning_indexer_v2_metadata 生成的结果一致 |
关于 metadata 的 shape:主算子文档中描述为 (1024,),这与仓库 lightning_indexer_v2_metadata.h 中的定义一致——该文件定义了LI_V2_METADATA_TOTAL_SIZE = 1024,其中前36 * 8个 int32 用于存放最多 36 个 AI Core(AIC)的任务描述(每个核 8 个字段:core 使能标志、bn2/m/s2 的起始与结束索引等),其余空间用于存放最多 72 个 AIV Core 的 load 任务元数据。因此 metadata 实际上是一个"分核任务表",驱动主算子按核并行执行。
mask_mode 参数解释:
- mask_mode=0,全计算模式(默认值)。
- mask_mode=3,Causal 模式(rightDownCausal,即以右顶点为划分的下三角场景)。
| 参数 | 单参数校验 | 存在性校验 |
|---|---|---|
| mask_mode | data_type 支持 INT;支持输入范围仅为 0、3,默认值为 0 | 可选输入,如果不传该参数,默认值为 0 |
关于 rightDownCausal 的具体含义,可参考 sparse_mode 参数说明 中 sparseMode=3 的矩阵遮蔽示意。
SeqLengths 参数组
| 参数 | 单参数校验 | 存在性校验 | 一致性校验 | 特性交叉校验 |
|---|---|---|---|---|
| seqused_q | tensor_type 支持 INT32;tensor_shape 为 (b,);仅支持非负整数;seqused_q 中的值需小于等于 q_s;seqused_k 中的值需小于等于 k_s | 可选参数 | 无 | 无 |
| seqused_k | 同上 | 可选参数 | 无 | 无 |
| cu_seqlens_q | tensor_type 支持 INT32;tensor_shape 为 (b+1,);值仅支持非负整数;其值应非递减(大于等于前一个值)排列,第一个元素为 0 且最后一个元素等于 q_t | 可选参数 | 无 | 当 layout_q 为 TND 时,必须传入;当 layout_q 不为 TND 时,不支持传入 |
| cu_seqlens_k | tensor_type 支持 INT32;tensor_shape 为 (b+1,);值仅支持非负整数;其值应非递减排列,第一个元素为 0 且最后一个元素等于 k_t | 可选参数 | 无 | 当 layout_k 为 TND 时,必须传入;当 layout_k 不为 TND 时,不支持传入 |
| max_seqlen_q | data_type 支持 INT;暂不生效,仅支持 -1;默认值为 -1 | 暂不生效,仅支持传入 -1 | 无 | 无 |
Paged Attention 参数组
当 block_table 不为空时,开启 Paged Attention。
| 参数 | 单参数校验 | 存在性校验 | 一致性校验 | 特性交叉校验 |
|---|---|---|---|---|
| block_table | tensor_type 仅支持 INT32;tensor_shape 为 (b, max_num_blocks_per_seq);值只能为正整数 | 可选参数 | 无 | PagedAttention 开启情况下,必须传入 seqused_k;PagedAttention 开启情况下,block_table 必须不为空 |
确定性计算
默认支持确定性计算(确定性的 Top-k 排序),这也与 aclnn 接口文档 aclnnLightningIndexerV2.md 中"默认确定性实现"的说明一致。从测试框架 pytest/README.md 看,其 golden 实现与 NPU 输出采用 CPU/NPU 精度对比的方式验证算子正确性,确定性语义是这种对比验证能够成立的前提。
调用示例
单算子模式调用
import torch import torch_npu import numpy as np from cann_ops_transformer.ops import lightning_indexer, lightning_indexer_metadata B = 2 S1 = 64 S2 = 130 N1 = 16 N2 = 1 D = 128 topk = 32 cmp_ratio = 4 # 计算压缩后S2长度及余数 S2_compressed = S2 // cmp_ratio res_k = S2 % cmp_ratio # 构造输入 q = torch.randn(B, S1, N1, D, dtype=torch.float16).npu() k = torch.randn(B, S2_compressed, N2, D, dtype=torch.float16).npu() w = torch.randn(B, S1, N1, dtype=torch.float32).npu() cmp_residual_k = torch.tensor([res_k] * B, dtype=torch.int32).npu() # 生成metadata metadata = lightning_indexer_metadata( num_heads_q=N1, num_heads_k=N2, head_dim=D, topk=topk, cmp_residual_k=cmp_residual_k, batch_size=B, max_seqlen_q=S1, max_seqlen_k=S2_compressed, layout_q="BSND", layout_k="BSND", mask_mode=3, cmp_ratio=cmp_ratio ) # 执行lightning_indexer sparse_indices, sparse_values = lightning_indexer( q, k, w, topk, cmp_residual_k=cmp_residual_k, metadata=metadata, max_seqlen_q=-1, layout_q="BSND", layout_k="BSND", mask_mode=3, cmp_ratio=cmp_ratio, return_value=0 ) print(f"sparse_indices shape: {sparse_indices.shape}")示例中的关键点说明:
- 由于 S2=130、cmp_ratio=4,压缩后 k 的实际参与长度为
130 // 4 = 32,余数res_k = 2通过cmp_residual_k传入,以配合 mask_mode=3 完成压缩后尾部残余 token 的 mask 与负载计算(对应 压缩相关 checker 的校验逻辑)。 cmp_residual_k的每个元素必须小于 cmp_ratio,即2 < 4成立。return_value=0时sparse_values返回 shape 为(0,)的空 tensor。
TorchAir(aclgraph)图模式调用
import torch import torch_npu import numpy as np import torchair from cann_ops_transformer.ops import lightning_indexer, lightning_indexer_metadata B = 2 S1 = 64 S2 = 130 N1 = 16 N2 = 1 D = 128 topk = 32 cmp_ratio = 4 S2_compressed = S2 // cmp_ratio res_k = S2 % cmp_ratio q = torch.randn(B, S1, N1, D, dtype=torch.float16).npu() k = torch.randn(B, S2_compressed, N2, D, dtype=torch.float16).npu() w = torch.randn(B, S1, N1, dtype=torch.float32).npu() cmp_residual_k = torch.tensor([res_k] * B, dtype=torch.int32).npu() class LightningIndexerNetwork(torch.nn.Module): def __init__(self): super(LightningIndexerNetwork, self).__init__() def forward(self, q, k, w, cmp_residual_k): metadata = torch.ops.cann_ops_transformer.lightning_indexer_metadata( num_heads_q=N1, num_heads_k=N2, head_dim=D, topk=topk, cmp_residual_k=cmp_residual_k, batch_size=B, max_seqlen_q=S1, max_seqlen_k=S2_compressed, layout_q="BSND", layout_k="BSND", mask_mode=3, cmp_ratio=cmp_ratio ) return torch.ops.cann_ops_transformer.lightning_indexer( q, k, w, topk, cmp_residual_k=cmp_residual_k, metadata=metadata, max_seqlen_q=-1, layout_q="BSND", layout_k="BSND", mask_mode=3, cmp_ratio=cmp_ratio, return_value=1 ) from torchair.configs.compiler_config import CompilerConfig config = CompilerConfig() config.mode = "reduce-overhead" npu_backend = torchair.get_npu_backend(compiler_config=config) torch._dynamo.reset() npu_mode = torch.compile(LightningIndexerNetwork(), fullgraph=True, backend=npu_backend, dynamic=False) sparse_indices, sparse_values = npu_mode(q, k, w, cmp_residual_k) print(f"sparse_indices shape: {sparse_indices.shape}")图模式下的实现支撑:仓库 graph_convert_lightning_indexer.py 注册了torch.ops.cann_ops_transformer.lightning_indexer_metadata的 GE 转换器(当前实现会显式抛出"不支持"的 RuntimeError,即 metadata 算子不进入 GE 图,而是通过torch.compiler.allow_in_graph保留在计算图中执行),而lightning_indexer主算子通过 TorchAir 的 aclgraph 后端编译执行。因此图模式示例中return_value=1,可以直接拿到与sparse_indices同 shape 的sparse_values。
与其他调用方式的关系
本接口对应的算子除 PyTorch API 外,还提供 C/C++ 的 aclnn 两段式接口aclnnLightningIndexerV2,其完整原型与 C++ 调用示例见 aclnnLightningIndexerV2.md,可运行的完整样例位于 examples/test_aclnn_lightning_indexer_v2.cpp。两种调用方式共用同一套 host 侧算子实现(lightning_indexer_v2_tiling.cpp 负责 tiling 计算),因此本文档的参数语义、约束与校验规则同样适用于 aclnn 接口。
小结
- LightningIndexerV2 通过"QK 相关性 + ReLU 过滤 + W 加权 + Top-k 选取"为每个 token 筛选稀疏的 Key 索引,是后续稀疏 Attention 的前置组件;
- 使用前必须先用
lightning_indexer_metadata生成 (1024,) 的 metadata 分核任务表,metadata 中编码了 AI Core 级任务的 batch/head/Q/K 分块起止索引(可对照 lightning_indexer_v2_metadata.h 中的字段布局); - 参数选择上需要重点把握三组联动关系:layout 组合(BSND/TND/PA_BBND)、mask_mode 与 cmp_ratio/cmp_residual_k 的组合约束、以及 Paged Attention 场景下 block_table 与 seqused_k 的强制搭配;
- 不同产品(Atlas A2/A3 与 Ascend 950)在 topk、可选参数(seqused_q、output_idx_offset、max_seqlen_q)与 TND Padding 行为上存在差异,落地时需按目标产品核对约束。
【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考