ops-transformer select_attention_operators:基于 Quest 在 Ascend 910B 上实现稀疏注意力块预测的加速内核
【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer
本模块是 CANN ops-transformer 仓库 experimental 目录下的一组 Ascend 910B 高性能向量内核,用于 LLM 解码阶段的稀疏注意力模式预测(Quest 算法)。读完本文,你将掌握其环境初始化与内核编译流程、quest_prefill_metadata与quest_block_select_paged系列三个 Python 接口的完整参数与硬性约束,并能依据仓库给出的开发规范为自己新增一个内核。
模块定位与目录结构
稀疏注意力预测的目标是:在逐 token 解码时,根据当前 query 向量与 KV-cache 的块级元数据(每块 K 向量的逐通道 max/min 向量),快速预测出 top-k 个“重要”的 KV 块索引,从而让解码只读取少数 KV 块,大幅降低显存带宽消耗。算法依据 Quest(ICML 2024)论文实现,并在原始方案上扩展了 GQA 支持(query 头按 KV 头分组取均值后预测)。
模块目录结构如下(引自 README):
. |-- experiments - per kernel: test (functional correctness) and benchmark (time and bandwidth) | |-- 2_quest_prefill_metadata - constructing metadata after prefill | |-- 3_quest_block_select_paged - quest sparse mask predictor using metadata | |-- 4_quest_block_select_paged_w - quest sparse mask predictor using metadata with extra sink+window features |-- kernels - python packages, each having one or more ascendc kernels and a single torch interface | |-- select_attn_ops - predictor kernels (quest predictors of sparse pattern during LLM decoding) `-- scripts |-- build_kernels.sh - builds all kernels `-- init_cann.sh - initialize the environment and Ascend device version- experiments/:每个内核一个目录,含
ref_*.py(Python 参考实现)、gen_data_*.py(输入数据生成)、test_*.py(正确性测试)与benchmark_*.py(时延/带宽基准); - kernels/select_attn_ops/:Python 包,内含 AscendC 内核源码(
quest_prefill_metadata.cpp、quest_block_select_paged.cpp)、torch 接口(torch_interface.cpp)与编译脚本; - scripts/:环境初始化与批量构建脚本。
运行环境要求
README 明确给出的已验证环境组合:
- 硬件:Ascend 910B2、Ascend 910B4;
- CANN 版本:8.0.RC3.beta1、8.2.RC2、8.3.RC1;
- Python 3.11.10;
- TorchNPU 2.4.0 / 2.5.1.post1;
- 其余依赖见 requirements.txt,其中实际固定了
torch==2.5.1、torch-npu==2.5.1.post1、pytest==9.0.2、scipy==1.14.1等版本。
创建 conda 环境
conda create -n sa python=3.11.10 -y conda activate sa pip install -r requirements.txt初始化 CANN 环境并编译内核
在 conda 环境内执行:
source scripts/init_cann.sh Ascend910B4 # change Ascend910B4 to your card model bash scripts/build_kernels.sh结合源码可以看到这两条命令的底层行为:
- init_cann.sh 会校验传入的 SOC 版本是否在
num_cores_map.sh的映射表中,随后依次执行conda activate sa、source /usr/local/Ascend/ascend-toolkit/set_env.sh、追加驱动LD_LIBRARY_PATH,并导出SOC_VERSION与NUM_CORES(单芯片 Davinci 核数)两个环境变量,最后调用 check_cann.sh 做环境自检。也就是说,NUM_CORES正是内核启动时用于确定核数上限(blockDim)的来源。 - build_kernels.sh 的逻辑很简单:遍历
kernels/*/子目录并对每个目录执行./build.sh,因此新增内核包时也要提供对应的build.sh。对于select_attn_ops包,目录中的compile.sh负责编译 AscendC 内核,setup.py负责构建 Python 扩展包;按 内核包 README 的 Good practices,仅改动内核.cpp时重跑 compile.sh 即可生效,改动torch_interface.cpp则必须用 build.sh 重建整个 Python 包。
三个核心算子:接口、参数与硬性约束
Python 侧通过from select_attn_ops import ...使用三个解码预测接口和一个预填充元数据接口:
import torch_npu from select_attn_ops import quest_block_select_paged_in_out_w help(quest_block_select_paged_in_out_w)内核自带完整的 pybind11 docstring(定义于 torch_interface.cpp 的PYBIND11_MODULE段),help()会打印参数说明与限制。以quest_block_select_paged_in_out_w为例,其签名与关键说明为:
quest_block_select_paged_in_out_w(query, maxblocks, minblocks, metadata_block_tables, seq_lens, tokens_since_metadata_update, selected_indices) -> Nonequery:[B, H, D],fp16/bf16,当前解码 token 的 query 向量;maxblocks/minblocks:[num_meta_blocks, BLOCK_SIZE, N, D],fp16/bf16,每个 KV 块的逐通道最大/最小元数据向量。注意:不存在的 KV 块的元数据位置必须填 0;metadata_block_tables:[B, MMBPR],int32,请求到元数据块的映射表;seq_lens:[B],int32,每个请求的序列长度;tokens_since_metadata_update:int,距上次元数据更新以来解码的 token 数(更新只发生在 BLOCK_SIZE 的整数倍位置);返回的块索引是“按序列内枚举 0..块数”的序号,而非 KV-cache 物理块号;selected_indices:[B, N, k],int32,预分配输出。
该内核因 910B 片上缓冲设计存在硬性限制(docstring 与源码TORCH_CHECK一致可查):
| 约束 | 要求 | 源码依据 |
|---|---|---|
| 头维度 | D == 128 | torch_interface.cpp 中TORCH_CHECK(D == DIM128, ...) |
| 块大小 | BLOCK_SIZE == 128 | TORCH_CHECK(BLOCK_SIZE == DIM128, ...) |
| GQA 分组 | H % N == 0且H / N <= BLOCK_SIZE | 同上 |
| 元数据块数/请求 | MMBPR <= 6(docstring 补充:低于 5 最稳定) | #define MAXMBPR 6与TORCH_CHECK(MMBPR <= MAXMBPR, ...) |
| top-k | k % 8 == 0(in_out 接口) | 由BYTES_ASCEND_DATA_BLOCK=32与 int32 索引 4 字节推出,k == k_round校验 |
| 数据类型 | query/maxblocks/minblocks 必须同为 fp16 或同为 bf16 | is_bfloat16三向一致性校验 |
| 窗口参数 | 0 <= tokens_since_metadata_update <= BLOCK_SIZE(_w接口) | 两条TORCH_CHECK边界校验 |
三个接口的差异(均出自 torch_interface.cpp):
quest_block_select_paged(query, maxblocks, minblocks, metadata_block_tables, seq_lens, k):内核自动分配输出[B, N, k_round],函数内部将k_round向上取整到 32 字节对齐后,再用slice裁剪回原始k返回;启动时tokens_since_metadata_update固定传 -1(禁用窗口特性);quest_block_select_paged_in_out(..., selected_indices):预分配输出,省去每次分配,同样以 -1 禁用窗口特性;quest_block_select_paged_in_out_w(..., tokens_since_metadata_update, selected_indices):在预分配输出基础上启用“window”特性——内核根据距上次更新的 token 数与序列长度,决定把 sink 块(块 0)与局部窗口块的索引强制加入选择结果。README 指出这一做法对实际精度有正面影响。
quest_prefill_metadata:元数据构建原理
预填充完成后需要为 K-cache 建立元数据,quest_prefill_metadata(k_cache, block_tables, seq_lens, metadata_block_tables, maxblocks, minblocks)完成这件事:
- 对每个 KV 块,在 token 维(BLOCK_SIZE=128 个 token)上做逐通道 reduce-max/reduce-min,得到一个 D 维向量;
- BLOCK_SIZE(128)个这样的 KV 块共 16384 个 token,其 128 个元数据向量打包成一个“元数据块”,写入
metadata_block_tables指定的maxblocks/minblocks区域; - 由于元数据块与 KV 块同尺寸,可以直接复用 vLLM 的 paged KV-cache 页表,甚至可以让
num_meta_blocks == num_kv_blocks、把maxblocks = k_cache、minblocks = v_cache原地存放(V 槽恰好闲置)。
quest_prefill_metadata.cpp 的实现要点(向量核,1 核处理 1 个 (batch, kv-head) 任务):
- 并行划分:以
B * N为任务空间,GetBlockIdx()步长遍历,启动核数取min(B*N, NUM_CORES)(与 init_cann.sh 导出的NUM_CORES呼应); - 带跨度的搬运:从 4D 张量
[num_kv_blocks, BLOCK_SIZE, N, D]中取某 KV 头切片时,用DataCopyParams的srcStride跳过其他头的行、dstStride=0压紧写入 UB,最后一个 KV 块按seq_len只搬运有效 token 数; - 对数式归约:
ReduceTokenDim模板函数在 UB 内做“成对 Max/Min + 尾部搬运”的迭代归约,log2(BLOCK_SIZE)轮把(BLOCK_SIZE, D)压缩为(1, D),避免引入额外存储; - 尾块零填充:请求 KV 块数不足 128 的末尾元数据块,用
Duplicate(0.0f)将未用行清零——这正是解码侧“不存在的块元数据必须为 0”约束的来源; - 写回:以
dstStride跳过其他 KV 头,把[BLOCK_SIZE, D]元数据块写回maxblocks/meta_blk_id[:, h, :]。
quest_block_select_paged:解码期 top-k 预测算法
内核包 README 给出了完整算法步骤:
1. For every batch, reduce-mean the query tensor across H dimension such that every group of H/N vectors of shape D are reduced to one vector of shape D denoted as "grouped_query[b,n]". 2. For each batch b and KV-head n: 2.1. Use metadata_block_tables[b] to locate the relevant metadata blocks for the sequence 2.2. For each metadata block in the sequence: 2.2.1. product_max, product_min = Elementwise-multiply grouped_query[b,n] with each maxblock and minblock vector 2.2.2. channel_max_product = Elementwise-max between the two products 2.2.3. block_scores = Reduce-sum the last dimension of approx_attention (D to 1) 2.3. selected_values, selected_indices = Find the top-k indices across all relevant blocks 3. Return selected_indices即利用max(q·maxblock, q·minblock)作为该块真实注意力分数的上界近似(Quest 论文的核心思想),做 D 维求和得到块分数,再跨所有元数据块做 top-k。两个实现细节值得注意:
- 该算子是纯向量(vector-only)内核,可处理任意数量的元数据块,天然适配变长序列;
- README 的 Implementation Notes 提醒:当
k > seq_lens[r] // BLOCK_SIZE(KV-cache 中实际块数不足 k 个)时内核仍会运行并返回 top-k,但其中部分索引是基于元数据块填充区计算出的“垃圾值”,用户必须保证 k 不大于实际 KV 块数; - 输出索引是序列内枚举号(0 起),而非 KV-cache 物理块号,接入 paged attention 时需要再经
block_tables做一次映射。
使用示例(与 内核包 README 一致的完整可运行片段):
import torch import torch_npu from select_attn_ops import quest_block_select_paged, quest_block_select_paged_in_out, quest_block_select_paged_in_out_w # Create dummy inputs (ideally minblocks and maxblocks should be generated from KV cache # to have proper 0 paddings) B, H, N, BLOCK_SIZE, D = 20, 32, 8, 128, 128 MMBPR = 1 k = 8 dtype = torch.bfloat16 device = "npu:0" num_meta_blocks = B * MMBPR max_seq_len = MMBPR * BLOCK_SIZE * BLOCK_SIZE query = torch.empty(B, H, D, dtype=dtype).uniform_(-1, 1).to(device).contiguous() maxblocks = torch.empty(num_meta_blocks, BLOCK_SIZE, N, D, dtype=dtype).uniform_(-1, 1).to(device).contiguous() minblocks = torch.empty(num_meta_blocks, BLOCK_SIZE, N, D, dtype=dtype).uniform_(-1, 1).to(device).contiguous() metadata_block_tables = torch.randint(0, num_meta_blocks, (B, MMBPR), dtype=torch.int32).to(device).contiguous() seq_lens = torch.randint(max_seq_len, max_seq_len+1, (B,), dtype=torch.int32).to(device).contiguous() # Option 1: let the kernel allocate the output ids = quest_block_select_paged(query, maxblocks, minblocks, metadata_block_tables, seq_lens, k) # Output shape: [B, N, k] # Option 2: pre-allocated output tensor ids = torch.zeros((B, N, k), dtype=torch.int32, device=device) # k must be a multiple of 8 quest_block_select_paged_in_out(query, maxblocks, minblocks, metadata_block_tables, seq_lens, ids) # Option 3: preallocated output + sink/window blocks forced in (good practical accuracy) tokens_since_metadata_update = 0 quest_block_select_paged_in_out_w(query, maxblocks, minblocks, metadata_block_tables, seq_lens, tokens_since_metadata_update, ids)而quest_prefill_metadata的调用(同文件给出的完整示例):
import torch import torch_npu from select_attn_ops import quest_prefill_metadata device = torch.device('npu:0') dtype_ind, dtype_val = torch.int32, torch.float16 B, N, BLOCK_SIZE, D = 4, 8, 128, 128 MKBPR = 200 # number of kv blocks in every request MMBPR = (MKBPR + BLOCK_SIZE - 1) // BLOCK_SIZE num_kv_blocks = B * MKBPR num_meta_blocks = B * MMBPR max_seq_len_per_req = BLOCK_SIZE * MKBPR seq_lens = torch.tensor([max_seq_len_per_req]*B, dtype=dtype_ind, device=device) k_cache = torch.randn(num_kv_blocks, BLOCK_SIZE, N, D, dtype=dtype_val, device=device) perm_kv_blk_ids = torch.randperm(num_kv_blocks, device=device)[:num_kv_blocks] block_tables = perm_kv_blk_ids.reshape((B, MKBPR)).to(dtype=dtype_ind, device=device) perm_meta_blk_ids = torch.randperm(num_meta_blocks, device=device)[:num_meta_blocks] metadata_block_tables = perm_meta_blk_ids.reshape((B, MMBPR)).to(dtype=dtype_ind, device=device) maxblocks = torch.zeros(num_meta_blocks, BLOCK_SIZE, N, D, dtype=dtype_val, device=device) minblocks = torch.zeros(num_meta_blocks, BLOCK_SIZE, N, D, dtype=dtype_val, device=device) # the outputs are filled into (maxblocks, minblocks) quest_prefill_metadata(k_cache, seq_lens, block_tables, metadata_block_tables, maxblocks, minblocks)涉及的尺寸参数(两个算子共用):B批大小、NKV 头数、BLOCK_SIZE每块 token 数(128)、D头维度(128)、MKBPR每请求最大 KV 块数、MMBPR每请求最大元数据块数、k每 KV 头返回的重要块数。
生产实践建议
README 给出的当前最佳实践(在 vllm-ascend 中验证):预填充结束后用quest_prefill_metadata()建立元数据,此后每 128 个 token 更新一次元数据,解码时每步用quest_block_select_paged_in_out_w()给定当前 query 预测重要 KV 块索引。
测试与性能验证
在 conda 环境下运行全部实验:
pytest -v experiments也可以按实验目录单独执行(如 experiments/2_quest_prefill_metadata/):
pytest -k basic -v # 仅基础扫描 pytest . # 该目录下全部测试 python test_quest_prefill_metadata.py # 内核 vs Python 参考实现的单场景对比并 dump 输出 python benchmark_quest_prefill_metadata.py # 时延与带宽基准仓库文档附带了在 x86 主机 + 910B4 上的实测数据(引自 实验 2 README):fp16 下quest_prefill_metadata达到约 0.52–0.57 TB/s,相对 910B4 标称 0.80 TB/s 全局内存带宽约 69% 利用率,且各配置下与参考实现输出一致(Outputs_equal=yes)。解码侧 实验 3 README 的基准显示,quest_block_select_paged/_in_out在 H=32、N=8、B=10~32、MMBPR=1~6(最长 98304 token)范围内时延约 15~300 μs、带宽 0.3~0.58 TB/s;文档同时说明 bf16 因向量单元需先转 float32、缓冲更大,带宽低于 fp16。每个元数据块覆盖128×128 = 16384个 token,故 MMBPR=10 即可支撑 160k 级序列(受MMBPR<=6限制,当前内核上限对应约 98k token)。
开发工作流:为模块新增一个内核 "OP"
README 定义了标准开发流程,可完整继承:
- 在
kernels/目录以下列两种方式之一添加内核实现:- 并入现有 Python 包,如
kernels/select_attn_ops/:新增一个OP.cpp,在 compile.sh 中加一行编译命令,并在torch_interface.cpp中注册 torch 接口; - 新建 Python 包
kernels/OP/,内含OP.cpp、torch_interface.cpp、compile.sh、build.sh(build_kernels.sh 依赖 build.sh 自动发现并构建);
- 并入现有 Python 包,如
- 创建专用实验目录
experiments/5_OP,实现四类程序:ref_OP.py:先写 Python 参考实现保证正确性基准;gen_data_OP.py:生成输入张量集合的数据生产函数;test_OP.py:先做单输入冒烟测试,再扩展为覆盖宽形状/数据类型范围的 pytest 自动化测试;benchmark_OP.py:测量时延与带宽。
这套“参考实现 → 数据生成 → 正确性测试 → 基准”的四件套与experiments/下现有三个实验目录的文件构成完全对应,是复现与扩展本模块的既定模板。
小结
select_attention_operators 以 Quest 论文为算法基础,用两个纯向量 AscendC 内核把“预填充后建元数据 + 解码期 top-k 块预测”压缩为带宽友好的轻量算子:quest_prefill_metadata以对数式 UB 归约构建 max/min 元数据并零填充尾块,quest_block_select_paged系列则用逐通道上界近似加 top-k 完成稀疏掩码预测,_w变体进一步强制纳入 sink 与窗口块以提升精度。全部接口、约束(D=128、BLOCK_SIZE=128、H/N≤128、MMBPR≤6、k 为 8 的倍数)均可在 torch_interface.cpp 的校验逻辑与 pybind docstring 中逐条对应;结合experiments/中的参考实现、测试与基准,读者既可直接接入 vLLM-Ascend 类推理框架,也能按仓库规范为同一框架贡献新的预测内核。
【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考