ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

pypto-gym 算子实战:GatherPaKvCache 的 Norm ND 分支设计与 PyPTO 实现解析

pypto-gym 算子实战:GatherPaKvCache 的 Norm ND 分支设计与 PyPTO 实现解析 pypto-gym 算子实战GatherPaKvCache 的 Norm ND 分支设计与 PyPTO 实现解析【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym导读本文以 GatherPaKvCache 目录下的 DESIGN.md 为骨架系统讲解 PyPTO-Gym 中gather_pa_kv_cache算子如何将 AscendCGatherPaKvCache的Norm ND分支映射为 PyPTO 的view assemble拷贝流水包括宿主侧 wrapper 的职责划分、两套 JIT kernel 的分工与切换阈值、tiling 与循环结构设计、以及 CPU golden 的精度验证链路。读完本文你将掌握分页 KV CachePagedAttention 的 KV 页如何被收集为连续输出的完整实现思路并能直接使用仓库内的测试入口在 NPU 上复现精度验证。该算子位于仓库 src/pypto_gym/ops/pypto_tensor/experimental/vector/GatherPaKvCache与其配套的 README.md、SPEC.md 与 API_REPORT.md 共同构成了该算子的完整文档闭环。1. 算子定位与活跃分支Active Branchgather_pa_kv_cache解决的是 PagedAttention / KV Cache 分页存储场景下的经典问题KV cache 按物理块block离散存放在key_cache/value_cache中需要通过block_tables把每个逻辑序列的 token 收集gather成连续输出的key_out/value_out。DESIGN.md 明确声明当前实现与 AscendCGatherPaKvCache的Norm分支对齐数据布局为NDkey_cache [num_blocks, block_size, key_num_heads, key_dim] value_cache [num_blocks, block_size, value_num_heads, value_dim] key_out [total_tokens, key_num_heads, key_dim] value_out [total_tokens, value_num_heads, value_dim]其中block_tables的形状为[Q, block_table_cols]Q为序列条数seq_lens为[Q 1]的累加长度cumsum形式seq_offset为[Q]的可选 token 偏移。需要特别注意的是PA_NZ布局在当前 PyPTO 代码中未实现属于明确的 out-of-scope 项这在 SPEC.md 的 Unsupported 章节 中同样得到印证。1.1 产品支持情况与支持范围按 README.md 的记录本算子支持的硬件平台为Ascend 950PR支持Atlas A3 训练系列产品 / Atlas A3 推理系列产品支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品支持支持范围Supported Scope要点如下仅支持cache_modeNormkey_cache与value_cache的数据类型限定为torch.bfloat16采用 ND 缓存布局如上节张量契约block_tables、seq_lens、seq_offset的数据类型限定为torch.int32seq_lens需为 cumsum 形式[Q 1]wrapper 默认is_seq_lens_cumsumTrue非 cumsum 形式会被拒绝seq_offsetNone可接受wrapper 会自动归一化为全零张量。不支持项包括cache_modePA_NZ、非 BF16 缓存张量、INT64 索引张量以及 CPU/sim 执行路径当前测试入口仅支持 NPU 模式。2. 宿主侧 Wrapper 职责Python wrapper对应 gather_pa_kv_cache_impl.py 中的gather_pa_kv_cache_out与gather_pa_kv_cache_wrapper承担以下六大职责dtype 与 rank 校验_validate_cache_pair强制 cache 张量为 rank 4 且 dtype 为 BF16并校验 K/V 的num_blocks、block_size一致_validate_index_tensor强制索引张量block_tables、seq_lens、seq_offset为 INT32。cache_mode Norm校验非 Norm 直接抛出NotImplementedError。seq_lens归一化为 cumsum 形式_build_seq_lens_cumsum校验seq_lens形状必须为[Q 1]且首元素为 0、各段长度非负、总 token 数不超INT32_MAX对应源码常量INT32_MAX 2**31 - 1。seq_offsetNone归一化为全零在gather_pa_kv_cache_out与 wrapper 中均有if seq_offset is None: seq_offset torch.zeros((q_count,), dtypetorch.int32, devicedevice)分支。输出分配当key_ref/value_ref未提供时_checked_ref按(total_tokens, heads, dim)自动torch.empty分配输出缓冲。block table 越界检查host 侧_validate_block_tables在 CPU 上逐行检查每条序列所需 block 数是否超过block_tables列数、被使用的物理块 ID 是否落在[0, num_blocks)。kernel 分派选择_select_gather_tile_config依据heads * dim是否超过 4096 选择默认或 large-token 的 tile 配置并最终调用对应的 JIT kernel。关于 shape 参数传递DESIGN.md 强调wrapper 从张量形状推导全部布局维度block_size、head 数、head 维度不向 JIT kernel 单独传递 Python 标量 shape 参数。这些维度在 PyPTO 特化specialization中是静态的只有 token 相关轴num_blocks、Q、total_tokens、block_table_cols在 JIT 签名中是动态的与 API_REPORT.md 的Dynamic and Static Axes章节一致。3. JIT Kernel 与分派逻辑当前实现包含两个活跃 kernel_gather_pa_kv_cache_nd_kernel_npu _gather_pa_kv_cache_nd_large_token_kernel_npu其中 large-token kernel 的选取条件为key_num_heads * key_dim 4096 或 value_num_heads * value_dim 4096该阈值对应源码中的 tile 配置选择函数_select_gather_tile_configgather_pa_kv_cache_impl.py#L308-L311def _select_gather_tile_config(key_num_heads, key_dim, value_num_heads, value_dim): if key_num_heads * key_dim 4096 or value_num_heads * value_dim 4096: return LARGE_TOKEN_GATHER_TILE_CONFIG return DEFAULT_GATHER_TILE_CONFIG这一分支覆盖了手动验证过的[64,128,64,128]大 head 场景即 README 中的heads64_k128_v128用例key_cache [64,128,64,128]、value_cache [64,128,64,128]、输出[6,64,128]而正常 kernel 覆盖网络 sweep 用例及较小 token 载荷。3.1 两个 kernel 共有的五步执行流程DESIGN.md 给出两套 kernel 完全一致的算法骨架将 cache reshape 为[num_blocks * block_size, heads, dim]外层循环遍历Q序列维度对每条序列循环其逻辑 cache block从block_tables读取physical_blockview出有效 token 范围并assemble到key_ref/value_ref。3.2 JIT 编译配置以 gather_pa_kv_cache_impl.py#L70-L79 的pypto.frontend.jit装饰器为例可以看到该算子在 NPU 上的关键运行时配置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_mode: pypto.RunMode.NPU强制 NPU 执行路径与仅支持 NPU的规格一致device_sched_mode: 1设备侧调度模式stitch_function_max_num: 128算子融合stitch时可拼接的最大函数数量上限ready_on_host_tensorsblock_tables、seq_lens、seq_offset这些 host 侧张量在 kernel 启动前即就绪用于运行时 shape/循环边界计算valid_shape_optimize: 1开启有效形状优化配合view的valid_shape参数减少无效计算vec_nbuffer_setting: {-2: 1, -1: 8}向量指令 buffer 数量设置。3.3 Kernel 签名中的动态/静态轴JIT kernel 签名gather_pa_kv_cache_impl.py#L80-L89使用pypto.DYNAMIC与pypto.STATIC标记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, ):即num_blocks、Q、total_tokens、block_table_cols为动态轴block_size、key_num_heads、key_dim、value_num_heads、value_dim为静态特化轴kernel 内部通过key_cache.shape[1]等方式读取。4. Tiling 设计4.1 两套 tile 配置DESIGN.md 给出两套 kernel 的向量 tile 形状Normal kernelK tile: [16, key_num_heads, key_dim] V tile: [32, value_num_heads, value_dim]Large-token kernelK tile: [8, key_num_heads, key_dim] V tile: [8, value_num_heads, value_dim]这两套配置与源码顶部的常量一一对应DEFAULT_GATHER_TILE_CONFIG [16, 32] # K tile, V tile LARGE_TOKEN_GATHER_TILE_CONFIG [8, 8]在 kernel 内部通过pypto.set_vec_tile_shapes(tile_config[0], key_num_heads, key_dim)与pypto.set_vec_tile_shapes(tile_config[1], value_num_heads, value_dim)gather_pa_kv_cache_impl.py#L120-L123在 K、V 两条拷贝链路上分别设置不同的向量任务粒度。4.2 large-token tile 的引入动机DESIGN.md 记录了重要性能决策large-token tile 之所以引入是因为大 head 场景下若沿用[1, 64, 128]形状会产生过多细小任务而[8, 64, 128]能使 BF16 tile 保持在约 128 KiB附近兼顾任务粒度与片上带宽利用。这是将性能特化与语义正确分离的典型做法——API_REPORT 的 Risks 章节也明确说明large-token 分支是性能特化而非语义要求。5. 循环结构设计Q 循环使用pypto.loop_unroll(..., unroll_list[2, 1])实现 2/1 混合展开block 循环使用标准动态pypto.loop。DESIGN.md 给出的伪代码如下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): copy valid token range for K copy valid token range for V对应的实际源码gather_pa_kv_cache_impl.py#L99-L125将这一结构落实为完整实现关键行语义如下for q_base, q_unroll in pypto.loop_unroll( 0, block_tables.shape[0], 1, namegather_q_loop, idx_nameq_idx, unroll_list[2, 1], ): for q_inner in range(q_unroll): q_idx q_base q_inner out_start seq_lens[q_idx] # 输出基准 cumsum[q] seq_len seq_lens[q_idx 1] - out_start # 序列长度 cumsum[q1] - cumsum[q] table_offset seq_offset[q_idx] // block_size block_count pypto.ceildiv(seq_len, block_size) for block_idx in pypto.loop(block_count, namegather_block_loop, idx_nameblock_idx, unroll_list[1]): physical_block block_tables[q_idx, table_offset block_idx] token_offset block_idx * block_size valid_tokens (seq_len - token_offset).min(block_size) out_offset out_start token_offset cache_offset physical_block * block_size ... pypto.set_vec_tile_shapes(tile_config[0], key_num_heads, key_dim) key_tile pypto.view(key_cache_3d, key_shape, [cache_offset, 0, 0], valid_shapekey_valid) pypto.assemble(key_tile, [out_offset, 0, 0], key_ref) ...这里的核心语义与 SPEC.md 的 Semantics 章节 完全一致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, :, :]值得注意的实现细节cache 先被pypto.reshape(..., inplaceTrue)展平为[num_blocks * block_size, heads, dim]因此view的源偏移直接用physical_block * block_size计算避免多维索引每个 block 的尾部可能不是满块valid_tokens (seq_len - token_offset).min(block_size)用于精确裁剪有效 token 数再通过view的valid_shape参数传入unroll_list[1]的 block 循环保持动态执行因为block_count依赖运行时seq_lens属动态边界。6. 已知性能要点Known Performance NotesDESIGN.md 明确记录了两个方向的性能结论未做夸大小 token 形状算子退化为少量微小拷贝任务几乎无 VF向量融合可做的算术量因此AICore 利用率可能偏低这类形状属于调度受限scheduling-bound而非计算受限大 token 形状得益于更大的向量 tile当前使用 large-token kernel 承接如[64,128,64,128]大 head 场景。API_REPORT.md 的 Risks 章节 对上述结论给出了同样的定位并补充PA_NZ是有意排除在范围之外而非能力缺失。7. 验证产物与精度验证链路7.1 验证工件清单主测试入口tests/ops/experimental/vector/GatherPaKvCache/test_gather_pa_kv_cache.pyCPU goldentests/ops/experimental/vector/GatherPaKvCache/gather_pa_kv_cache_golden.py网络 sweep 用例tests/ops/experimental/vector/GatherPaKvCache/test_cases.json最近一次手动记录的 target-shape swimlane 产物output/output_20260530_111415_350728_1001774_C0A96050/merged_swimlane.json output/output_20260530_111420_833310_1001774_C0A96050/merged_swimlane.json7.2 网络 sweep 用例构成test_cases.json包含十个 ND sweep 用例level0~level9基本形状为key_cache [5513,128,1,512] value_cache [5513,128,1,64] blockTables [Q,8] seqLens [Q 1] # cumsum 形式 seqOffset [Q] key_ref [T,1,512] value_ref [T,1,64] cache_mode Norm dtype BF16sweep 覆盖Q4, T in {6,27,31,35,39,43,47,51,55}与Q3, T44的 token 数变化。从 test_cases.json 中可以看到每个 level 以seed区分、以total_tokens/q_count描述规模的完整字段结构如level0为Q4, total_tokens6seq_lens_shape[5]is_seq_lens_cumsumtrue。另有两个手动验证过的 target shapeheads1_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即触发 large-token kernel64 * 128 8192 4096的验证用例。7.3 Golden 实现与精度判定gather_pa_kv_cache_golden.py 提供纯 CPU 参考实现核心为_gather_kv_pages函数按序列逐 block 计算physical_block与valid_tokens以切片赋值完成key_out[out_start:out_end] key_cache[physical_block, :valid_tokens]的等价收集逻辑。make_case函数则负责按给定total_tokens/q_count/ 块表列数等参数生成随机 BF16 cache 与 INT32 block table含按lengths构造 cumsumseq_lens的逻辑供测试与 golden 共用。由于该算子本质是纯 BF16 拷贝/收集操作精度判定使用torch.equal对 NPU 输出与 CPU golden 做逐元素严格相等比较SPEC.md 的 Precision 章节 记录为atol: 0.0 / rtol: 0.0。测试流程在 test_gather_pa_kv_cache.py#L134-L172 的_run_case中完成构造 CPU 用例 → 搬运到 NPU → 调用 wrapper → 同步 → golden 对照 → 分别对 K/V 输出断言torch.equal全部通过后主流程打印[PRECISION_PASS]。7.4 运行方式运行全部精度用例以仓库 README 记录的实测环境为例source /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运行指定用例如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列出全部用例用例 ID 缺失或未知时会抛出明确错误保证每个level0~level9都有对应用例REQUIRED_LEVELS校验。8. 2026-06-11 整改同步DESIGN.md 末尾记录了最近一次整改要点三份配套文档README / SPEC / API_REPORT均带同名整改同步章节形成一致口径Host wrapper 现在要求 cumsumseq_lens形状为[Q 1]_build_seq_lens_cumsum中is_seq_lens_cumsumFalse直接抛错seq_lens必须以 0 开头、长度非负wrapper 与 golden 的默认值均为is_seq_lens_cumsumTrueKernel 路径不变仍是同一套 ND gather 路径本次整改只是删除了 host 侧非 cumsum 的隐式归一化分支避免语义二义性测试用例与 golden 统一改用 cumsumseq_lenstest_cases.json中所有 level 的seq_lens_shape已从[Q]调整为[Q 1]make_case通过_finalize_seq_lens_tensor生成[0] cumsum(lengths)的序列长度张量。这也解释了为何 DESIGN.md 第 2 节特别强调seq_lensnormalization to cumsum form是 wrapper 的既有职责——整改后该职责从接受两种形式并归一化收敛为只接受 cumsum 形式并校验kernel 侧逻辑保持稳定。9. 总结从 DESIGN 到源码的完整映射把 DESIGN.md 的八节结构与仓库源码一一对应可以形成一张完整的设计→实现对照表DESIGN.md 章节源码/文档落点Active BranchNorm NDSPEC.md Tensor Contract 与 README Supported ScopeWrapper Responsibilitiesgather_pa_kv_cache_out/gather_pa_kv_cache_wrappergather_pa_kv_cache_impl.py#L314-L385JIT Kernels_gather_pa_kv_cache_nd_kernel_npu及 tile 选择函数gather_pa_kv_cache_impl.py#L80-L125TilingDEFAULT_GATHER_TILE_CONFIG[16,32]/LARGE_TOKEN_GATHER_TILE_CONFIG[8,8]常量Loop Structureloop_unroll(..., unroll_list[2,1])loop(block_count)Performance NotesAPI_REPORT Risks小 shape 调度受限、大 shape 性能特化Validation Artifactstest_gather_pa_kv_cache.py gather_pa_kv_cache_golden.py test_cases.json2026-06-11 整改全链路强制 cumsumseq_lens删除 host 隐式归一化对于希望在 PyPTO 上实现类似分页数据收集类算子的开发者本算子提供了一套可复用的范式host 侧完成全部校验与形状推导 → JIT 内以reshape view assemble表达按块收集 → 用set_vec_tile_shapes按载荷规模切换 tile 粒度 → 以 CPU golden torch.equal做零容忍精度验证。在此基础上把 block 表映射为 ND 索引、把尾部有效 token 用valid_shape裁剪即可推广到其他页式拷贝/压缩类算子场景。【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进