:基于 view + assemble 的分页 KV Cache 收集算子实战解析)
PyPTO 实现 AscendC GatherPaKvCacheNorm/ND 分支基于 view assemble 的分页 KV Cache 收集算子实战解析【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym导读本文基于 pypto-gym 仓库中GatherPaKvCache算子的 API_REPORT深入解析如何用 PyPTO 编程框架复刻 AscendCGatherPaKvCache的NormND分支——即将分页PagedKV Cache 中不连续的物理块按block_tables索引收集为连续的 key/value 输出。读完本文你将掌握该算子的 PyPTO API 映射关系、动态/静态轴划分、JIT 核函数双层循环结构、tiling 配置策略、精度验证方法以及当前实现的能力边界与已知风险可直接对照仓库源码与测试用例复现全流程。1. 算子背景为什么需要 GatherPaKvCache在大模型 PagedAttention 类推理框架中KV Cache 通常以固定大小的物理块block分页存放在显存中每个序列通过一张逻辑块 → 物理块的映射表block_tables访问自己的缓存。GatherPaKvCache的职责正是把这些物理上离散的缓存块按序列索引收集gather成连续的 key/value 张量供后续注意力计算使用。pypto-gym 中该算子的定义见 SPEC.mdgather_pa_kv_cachegathers discontiguous paged KV cache blocks into contiguous key/value outputs.当前 PyPTO 实现对齐 AscendC 的Norm分支ND 布局用 PyPTO 的viewassemble两个张量视图/落盘原语来完成整个收集过程而非逐元素拷贝。整体数据流如下key_cache [num_blocks, block_size, key_num_heads, key_dim] -- view 读取 -- value_cache [num_blocks, block_size, value_num_heads, value_dim] key_out [total_tokens, key_num_heads, key_dim] -- assemble 写入 -- value_out [total_tokens, value_num_heads, value_dim]2. PyPTO API 映射总览API_REPORT.md 第 2 节给出了该算子的核心 API 映射表这是理解整个实现的第一把钥匙需求PyPTO API用途JIT 核函数pypto.frontend.jitNPU 执行动态循环pypto.loop、pypto.loop_unroll遍历序列Q轴与块block轴Cache 展平pypto.reshape(..., inplaceTrue)[num_blocks, block_size, H, D] - [num_blocks * block_size, H, D]源数据切片pypto.view读取 cache 块范围写回pypto.assemble将收集到的范围写入输出Tilingpypto.set_vec_tile_shapes控制向量任务粒度对应的实际实现位于 gather_pa_kv_cache_impl.py。其中 JIT 核函数_gather_pa_kv_cache_nd_kernel_npu的装饰器配置第 70-79 行值得展开说明pypto.frontend.jit( runtime_options{ run_mode: pypto.RunMode.NPU, device_sched_mode: 1, stitch_function_max_num: 128, ready_on_host_tensors: [block_tables, seq_lens, seq_offset], valid_shape_optimize: 1, }, pass_options{vec_nbuffer_setting: {-2: 1, -1: 8}}, )run_modepypto.RunMode.NPU强制 NPU 执行路径ready_on_host_tensors声明block_tables、seq_lens、seq_offset三个索引张量为宿主侧就绪让核函数可以按需读取valid_shape_optimize1与vec_nbuffer_setting配合view的valid_shape参数做向量 buffer 与有效形状优化。3. 张量契约与数据语义3.1 输入输出张量依据 SPEC.md 第 2 节张量契约如下输入名称ShapeDtype说明key_cache[num_blocks, block_size, key_num_heads, key_dim]BF16ND 布局的 K cachevalue_cache[num_blocks, block_size, value_num_heads, value_dim]BF16ND 布局的 V cacheblock_tables[Q, block_table_cols]INT32逻辑块到物理块的映射表seq_lens[Q]或[Q 1]INT32序列长度或累计长度key_ref[total_tokens, key_num_heads, key_dim]BF16可选输出 buffervalue_ref[total_tokens, value_num_heads, value_dim]BF16可选输出 bufferseq_offset[Q]或NoneINT32可选的 token 偏移用于定位block_tables起始列输出名称ShapeDtypekey_out[total_tokens, key_num_heads, key_dim]BF16value_out[total_tokens, value_num_heads, value_dim]BF163.2 收集语义SPEC 第 3 节对每个序列q和每个 token核心映射为logical_block token_in_seq // block_size slot token_in_seq % block_size physical_block block_tables[q, table_offset(q) logical_block] key_out[output_base(q) token_in_seq, :, :] key_cache[physical_block, slot, :, :] value_out[output_base(q) token_in_seq, :, :] value_cache[physical_block, slot, :, :]当is_seq_lens_cumsumTrue时当前实现唯一支持的模式seq_len(q) seq_lens[q 1] - seq_lens[q] output_base(q) seq_lens[q]当提供seq_offset时table_offset(q) seq_offset[q] // block_size否则 wrapper 自动创建全零偏移张量见实现第 333-335 行。3.3 形状约束SPEC 第 4 节key_cache/value_cache均为 rank 4且num_blocks、block_size必须一致key_ref.shape[1:] key_cache.shape[2:]value_ref同理block_tables.shape[0] Qcumsum 模式下seq_lens.shape [Q 1]且必须以 0 开头seq_offset非负且能被block_size整除被使用的block_tables条目必须在[0, num_blocks)区间内。这些约束不仅写在 SPEC 里也逐一落在 wrapper 的校验函数中实现第 140-290 行_validate_cache_pair负责 dtype/rank/设备/维度一致性_build_seq_lens_cumsum负责 cumsum 语义与 token 总数合法性_validate_seq_offsets与_validate_block_tables负责宿主侧host的偏移与块表越界检查。4. 动态轴与静态轴划分API_REPORT.md 第 3 节明确了 PyPTO 特化specialization下的轴属性划分动态轴随运行输入变化num_blocks Q total_tokens block_table_cols静态特化轴从张量 shape 内读取block_size key_num_heads key_dim value_num_heads value_dim这一点在 JIT 核函数签名中体现得非常直观实现第 80-88 行def _gather_pa_kv_cache_nd_kernel_npu( key_cache: pypto.Tensor([pypto.DYNAMIC, pypto.STATIC, pypto.STATIC, pypto.STATIC], pypto.DT_BF16), value_cache: pypto.Tensor([pypto.DYNAMIC, pypto.STATIC, pypto.STATIC, pypto.STATIC], pypto.DT_BF16), block_tables: pypto.Tensor([pypto.DYNAMIC, pypto.DYNAMIC], pypto.DT_INT32), seq_lens: pypto.Tensor([pypto.DYNAMIC], pypto.DT_INT32), key_ref: pypto.Tensor([pypto.DYNAMIC, pypto.STATIC, pypto.STATIC], pypto.DT_BF16), value_ref: pypto.Tensor([pypto.DYNAMIC, pypto.STATIC, pypto.STATIC], pypto.DT_BF16), seq_offset: pypto.Tensor([pypto.DYNAMIC], pypto.DT_INT32), tile_config: list, ):可以看到cache 的首维num_blocks与输出首维total_tokens是DYNAMIC其余 head/dim 维度全部STATIC。这意味着 PyPTO 会为固定的 head 数与 head 维度生成特化代码而每次调用只动态传入 token 数量——这是该算子在网络扫描network sweep场景下保持高性能的关键设计。正如 README 所述The wrapper derivesblock_size, head counts, and head dimensions from tensor shapes. These dimensions are static per PyPTO specialization; only token-related axes are dynamic in the JIT signature.5. Wrapper 职责与核函数结构5.1 Wrapper 职责DESIGN 第 2 节Python 侧 wrappergather_pa_kv_cache_wrapper/gather_pa_kv_cache_out实现第 314-385 行承担全部宿主侧工作dtype 与 rank 校验cache_mode Norm校验seq_lens归一化当前强制 cumsum 形式非 cumsum 直接报错seq_offsetNone归一化为全零张量未提供key_ref/value_ref时自动分配输出_checked_ref实现第 225-230 行block table 的宿主侧越界检查核函数分派选择_select_gather_tile_config实现第 308-311 行。特别值得注意wrapper 从张量 shape 推导全部布局维度不向 JIT 核函数传递任何独立的 Python 标量 shape 参数——tile_config是唯一的额外参数。5.2 双核函数分派DESIGN 第 3 节当前实现包含两个激活核函数_gather_pa_kv_cache_nd_kernel_npu 常规 token 负载 _gather_pa_kv_cache_nd_large_token_kernel_npu大 token 负载性能特化选择条件实现第 308-311 行key_num_heads * key_dim 4096 或 value_num_heads * value_dim 4096该阈值覆盖了人工验证的[64,128,64,128]大 head 数场景常规核函数覆盖网络扫描与较小 token 负载。两个核函数的共同流程DESIGN 第 3 节将 cache reshape 为[num_blocks * block_size, heads, dim]遍历Q对每个序列遍历其逻辑 cache 块从block_tables读取physical_blockview有效 token 范围assemble到key_ref/value_ref。5.3 核函数内部循环DESIGN 第 5 节Q 循环使用pypto.loop_unroll(..., unroll_list[2, 1])做 2 路展开块循环使用标准动态pypto.loop结构如下对应实现第 99-125 行for q_base, q_unroll in loop_unroll(Q, [2,1]): for q_inner in range(q_unroll): seq_len seq_lens_cumsum[q1] - seq_lens_cumsum[q] block_count ceildiv(seq_len, block_size) for block_idx in loop(block_count): physical_block block_tables[q, table_offset block_idx] valid_tokens min(seq_len - block_idx*block_size, block_size) set_vec_tile_shapes(K tile) key_tile view(key_cache_3d, [block_size, K_heads, K_dim], [cache_offset,0,0], valid_shapekey_valid) assemble(key_tile, [out_offset,0,0], key_ref) set_vec_tile_shapes(V tile) value_tile view(value_cache_3d, [block_size, V_heads, V_dim], [cache_offset,0,0], valid_shapevalue_valid) assemble(value_tile, [out_offset,0,0], value_ref)每轮循环只拷贝有效 token 数valid_tokens因此最后一个不完整块tail block也能被正确处理——这正对应 PyPTO 的valid_shape机制view声明实际有效范围避免越界读取。6. Tiling 配置两种粒度策略API_REPORT.md 与 DESIGN.md 第 4 节共同给出了两套向量 tiling 配置常规核函数K tile: [16, key_num_heads, key_dim] V tile: [32, value_num_heads, value_dim]大 token 核函数K tile: [8, key_num_heads, key_dim] V tile: [8, value_num_heads, value_dim]对应实现中的常量实现第 16-17 行DEFAULT_GATHER_TILE_CONFIG [16, 32] LARGE_TOKEN_GATHER_TILE_CONFIG [8, 8]为什么大 token 场景要单独调小 tileDESIGN 中的解释很直白[1,64,128]会产生太多细小任务而[8,64,128]可以把 BF16 tile 控制在约 128 KiB 左右兼顾任务粒度与内存占用。这是一个典型的任务数 × 单任务体积权衡token 负载越大越需要靠更大的向量 tile 摊薄调度开销。7. 约束与能力边界API_REPORT.md 第 4 节列出的约束必须严格执行仅支持cache_modeNorm仅支持 BF16 的 cache 与输出张量仅支持 INT32 索引张量block_tables、seq_lens、seq_offsetseq_offset必须能被block_size整除seq_lens必须是 cumsum 形式[Q 1]非 cumsum 的seq_lens会被 wrapper 直接拒绝实现第 194-195 行raise ValueError(only cumsum seq_lens [Q 1] is supported)。结合 SPEC.md 第 7 节明确不支持的范围还包括PA_NZcache 布局、INT8/FP16/FP32 cache、INT64 索引、任意非 ND 布局、非 cumsumseq_lens以及 CPU/sim 执行路径当前测试入口仅支持 NPU 模式见 README.md。产品支持情况README为Ascend 950PR、Atlas A3 训练/推理系列、Atlas A2 训练/推理系列。8. 精度验证从网络扫描到手工靶向 shape8.1 测试用例体系API_REPORT.md 第 5 节声明所有test_cases.json条目均由test_gather_pa_kv_cache.py验证通过。测试骨架如下测试入口test_gather_pa_kv_cache.pyCPU 黄金参考gather_pa_kv_cache_golden.py用例配置test_cases.json测试采用torch.equal精确比对该算子是纯 BF16 copy/gather无任何算术误差将 NPU 输出与 CPU golden 做逐位相等判断测试文件第 168-172 行。test_cases.json中十个 ND 扫描用例level0 ~ level9的公共 shape 如下key_cache [5513,128,1,512] value_cache [5513,128,1,64] blockTables [Q,8] seqLens [Q 1] # cumsum以 0 开头 seqOffset [Q] key_ref [T,1,512] value_ref [T,1,64] cache_mode Norm dtype BF16扫描覆盖Q4, T ∈ {6,27,31,35,39,43,47,51,55}以及Q3, T44的组合seq_lens_shape均为[Q 1]且is_seq_lens_cumsum: true整改后的统一形态。8.2 额外手工验证的靶向 shape除网络扫描外API_REPORT.md 记录了两组手工靶向 shape 的精度验证heads1_k128_v64: key_cache [64,128,1,128] value_cache [64,128,1,64] output [6,1,128], [6,1,64] heads64_k128_v128: key_cache [64,128,64,128] value_cache [64,128,64,128] output [6,64,128], [6,64,128]其中heads64_k128_v128恰好命中heads * dim 64 * 128 8192 4096的大 token 分支。最新一轮手工运行两组靶向 shape 均通过精度[CUSTOM_PRECISION_PASS]8.3 运行测试README 给出了完整的 NPU 精度测试命令注意环境为昇腾 NPU torch_npusource /mnt/workspace/gitCode/cann/pypto/env_setup.sh cd /mnt/workspace/zhangsr/pypto-gym-2 PYTHONPATH/mnt/workspace/zhangsr/pypto-gym-2/src:/tmp/pypto-wheel:${PYTHONPATH} \ TILE_FWK_DEVICE_ID0 \ /opt/buildtools/Python-3.11.4/bin/python3 tests/ops/experimental/vector/GatherPaKvCache/test_gather_pa_kv_cache.py --run-mode npu只运行指定 level例如 level0 与 level9PYTHONPATH/mnt/workspace/zhangsr/pypto-gym-2/src:/tmp/pypto-wheel:${PYTHONPATH} \ /opt/buildtools/Python-3.11.4/bin/python3 tests/ops/experimental/vector/GatherPaKvCache/test_gather_pa_kv_cache.py level0 level9 --run-mode npu测试脚本还支持--list参数列出全部用例。全部用例通过时主流程返回 0 并打印[PRECISION_PASS]否则打印[PRECISION_FAIL]测试文件第 206-212 行。此外测试用例构造器make_casegolden 文件第 384-449 行支持自定义total_tokens、q_count、num_blocks、block_size、head 数与 dim可直接用于复现上面的手工靶向 shape。9. 已知风险与性能特征API_REPORT.md 第 6 节与 DESIGN.md 第 6 节共同披露了三个关键风险小 shape 是调度受限而非计算受限小 token 形状下算子退化为几个微小的拷贝任务几乎无算术可供 VF 融合因此 AICore 利用率偏低PA_NZ有意不在范围内当前 ND 路径不实现 NZ 布局这与产品侧对PA_NZ的支持形成差距属于明确的 scope 边界大 token 分支是性能特化而非语义要求heads * dim 4096时切到大 tile 配置本质是性能优化不改变算子语义。10. 2026-06-11 整改要点API_REPORT.md 第 7 节记录了最近一次整改且整改已同步到 SPEC/DESIGN/README/测试/黄金实现各处wrapper 与 golden 的is_seq_lens_cumsum默认值统一为Truetest_cases.json所有网络扫描 level 的seq_lens_shape已统一记录为[Q 1]的 cumsum 形态非 cumsum 的seq_lens不再被该 ND 网络路径隐式归一化——宿主侧归一化分支被移除传入即报错。这一整改的直接后果是调用契约更严格也更明确任何接入方都必须自己维护 cumsum 形式的seq_lens以 0 开头、单调不减避免 wrapper 侧悄悄帮你归一化带来的歧义。11. 总结pypto-gym 的GatherPaKvCache是 PyPTO 张量视图原语在真实网络算子上的典型落地案例它用view读离散 cache 块assemble写连续输出reshape展平块表loop/loop_unroll双层动态循环set_vec_tile_shapes向量 tiling 控制五类 API完整复刻了 AscendCGatherPaKvCache的Norm/ND 分支。其设计要点可归纳为静态轴特化 动态 token 轴head 数与 dim 从 shape 读取并静态特化Q/token 相关轴保持动态双核函数分派heads * dim 4096时切换大 token 配置tile 8×8否则用默认配置K 16 / V 32严格宿主校验cumsumseq_lens、seq_offset整除性、block table 越界均在宿主侧完成纯拷贝语义 精确比对精度验证采用torch.equal测试覆盖 10 个网络扫描 level 与 2 组手工靶向 shape。如需深入可直接阅读仓库内配套文档与源码API_REPORT.md本文主体、SPEC.md张量契约、DESIGN.md设计细节、gather_pa_kv_cache_impl.py核心实现以及测试三件套test_gather_pa_kv_cache.py、gather_pa_kv_cache_golden.py、test_cases.json。【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考