PyPTO-Gym 分页缓存散射写入模式(AT-16)实战:Paged Cache Scatter/Update 算子设计与实现

发布时间:2026/9/19 12:39:34
PyPTO-Gym 分页缓存散射写入模式(AT-16)实战:Paged Cache Scatter/Update 算子设计与实现 PyPTO-Gym 分页缓存散射写入模式AT-16实战Paged Cache Scatter/Update 算子设计与实现【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym分页 KV 缓存Paged KV Cache是长序列推理与 Paged Attention 架构的核心数据结构而将每步新算出的 K/V 写入缓存中指定物理位置则是其中最高频的写路径。本文以 PyPTO-Gym 仓库中的原子模式卡片 AT-16: Paged Cache Scatter/Update 为主线结合仓库内scatter_pa_kv_cache完整算子实现、GLM / DeepSeek 系列真实调用点与配套测试用例系统讲解在 PyPTO 框架下如何用scatter_update在纯 Vector 排布中完成 KV cache 散射更新。读完本文你将掌握物理 block 寻址、2D reshape loop valid_shape 的动态 shape 处理套路、pypto.scatter_update的调用约束以及围绕该模式进行精度验证与性能调优的完整方法。一、模式卡片核心AT-16 的定位与计算流在 PyPTO-Gym 的算子设计体系中原子模式Atom Pattern被集中收录于 patterns/atoms/index.md共 23 个。AT-16 属于其中的散射类scatter写路径模式与 AT-17 Block Table Gather读路径构成 Paged KV 缓存读写的一对镜像原子。1.1 卡片原始定义AT-16 卡片的核心信息如下标题Paged Cache Scatter/Update描述将计算结果写入分页 KV 缓存的指定物理位置tagsscatterflow_patternV纯 Vector 排布examplesGLMAttnFusion、MLAPrologQuant、Compressor卡片给出的计算流伪代码为# 将 cache_index 映射到物理 block 位置 physical_idx block_table[batch_idx, logical_block] cache_4d reshape(cache, [total_blocks, block_size, N, D]) src_4d reshape(src, [1, 1, N, D]) # scatter_update: 在 axis-2 方向将 src 写入 cache 的 physical_idx 位置 scatter_update(cache_4d, axis-2, indexphysical_idx, srcsrc_4d) cache_4d.move() # 确保写回三个要点值得展开纯 V 排布KV cache 更新是典型的索引驱动写操作不需要 Cube 参与整条数据流落在 Vector 单元上因此 flow_pattern 标为 V。axis-2 写入scatter_update固定沿倒数第二维block 内 offset 维做散射index 决定目标物理槽位src 是被写入的数据。.move()写回PyPTO 的 buffer 生命周期模型要求对原地更新的 cache 显式调用.move()确保修改真正落回原始 tensor。1.2 与相邻模式的关系在骨架模式 SK-05 Fused Pre-Attn 中AT-16 被明确嵌入两阶段融合注意力骨架的 Phase1 末尾KV Cache scatter 推荐放 Phase1 末AT-16写入紧贴 RoPE 之后与 Phase2 的 KV 读完全解耦。其逻辑在于预处理阶段Norm → Quant Linear → Dequant/Split → RoPE产出的 K/V 一旦就绪应立即散射写入 cache而 Phase2 的 Flash Attention 再通过 AT-17 Block Table Gather 按 block_table 零搬运拼装出连续 KV 块进行读取从而让读写两条路径在时间与 buffer 上彻底解耦。此外 AT-20 Tail Block 卡片指出尾块 valid_shape 声明正是 AT-16 / AT-17 写回阶段的前置形状声明三者经常在同一 kernel 中协同出现。二、分页缓存的物理寻址模型要正确实现散射写入必须先厘清 Paged KV Cache 的物理布局与索引语义。仓库中scatter_pa_kv_cache算子的 README 给出了精确的数学定义key_cache[block_idx, block_offset, :, :] key[i, :, :] value_cache[block_idx, block_offset, :, :] value[i, :, :] 其中 block_idx slot_mapping[i] // block_size block_offset slot_mapping[i] % block_size即 cache 是四维张量[num_blocks, block_size, num_heads, head_size]block_idx物理 block 编号由slot_mapping[i] // block_size得到block_offsetblock 内偏移由slot_mapping[i] % block_size得到每个 token 在 cache 中的槽位由slot_mappingvLLM 风格的 slot 映射等价于卡片中的physical_idx唯一确定。在 AT-16 卡片伪代码中block_table[batch_idx, logical_block]是逻辑块到物理块的映射表——同一个 batch 序列的逻辑 KV 块在物理内存中可以不连续这正是分页的意义所在按需分配物理块、减少显存碎片、提升并发序列的缓存利用率。三、完整算子实现scatter_pa_kv_cache 源码精读仓库提供了一个可直接运行的完整示例算子 scatter_pa_kv_cache_impl.py它是 AT-16 模式最忠实、注释最详尽的落地。下面按实现步骤逐段拆解。3.1 动态 shape 注解与 JIT 配置pypto.frontend.jit(runtime_options{device_sched_mode: 0}, pass_options{vec_nbuffer_setting: {-2: 1, -1: 8}}) def scatter_pa_kv_cache_kernel( key: pypto.Tensor([pypto.DYNAMIC, pypto.STATIC, pypto.STATIC], pypto.DT_BF16), key_cache: pypto.Tensor([pypto.DYNAMIC, pypto.STATIC, pypto.STATIC, pypto.STATIC], pypto.DT_BF16), slot_mapping: pypto.Tensor([pypto.DYNAMIC], pypto.DT_INT32), value: pypto.Tensor([pypto.DYNAMIC, pypto.STATIC, pypto.STATIC], pypto.DT_BF16), value_cache: pypto.Tensor([pypto.DYNAMIC, pypto.STATIC, pypto.STATIC, pypto.STATIC], pypto.DT_BF16), ):关键设计决策num_tokens是唯一真正的动态轴pypto.DYNAMIC取值范围 1~16384num_blocks、block_size、num_heads、head_size均为编译期静态常量示例中num_blocks9760, block_size128, num_heads2, head_size256。num_blocks在注解中被标为 DYNAMIC 但在 kernel 内通过key_cache.shape[0]读取README 明确这是已知限制之一注解为动态、实际按静态 shape 编译。slot_mapping使用DT_INT32是索引张量的标准 dtype。JIT 选项vec_nbuffer_setting用于控制 Vector buffer 数量属于后续性能调优的旋钮。3.2 reshape 降维4D → 2Dkv_dim num_heads * head_size # 512 cache_2d_shape [num_blocks * block_size, kv_dim] # [1249280, 512] key_cache_2d pypto.reshape(key_cache, cache_2d_shape, inplaceTrue) value_cache_2d pypto.reshape(value_cache, cache_2d_shape, inplaceTrue) key_2d pypto.reshape(key, [num_tokens, kv_dim], inplaceTrue) value_2d pypto.reshape(value, [num_tokens, kv_dim], inplaceTrue) slot_mapping_2d pypto.reshape(slot_mapping, [num_tokens, 1], inplaceTrue)这一步把四维物理 cache 展平为二维[num_blocks * block_size, kv_dim]——第一维恰好是扁平化的槽位号slot_mapping直接取值这正是scatter_update沿axis-2散射时使用的索引空间。三个输入统一降到二维后src[tokens, kv_dim]、index[tokens, 1]、cache[slots, kv_dim]三者形状语义对齐。实现注释还记录了两条 PyPTO 约束经验inplaceTrue的 reshape 输出不能是函数输出参数此处 key_cache/value_cache 不是输出参数规避了该限制view 的 shape 参数必须是纯 Python int不能是符号表达式。3.3 TileShape 与 Tiling 策略tile_tokens 32 if kv_dim 512: tile_tokens 32 elif kv_dim 4096: tile_tokens 16 else: tile_tokens 4 pypto.set_vec_tile_shapes(tile_tokens, kv_dim)TileShape 的维度数必须与 src 维度数一致2D即[tile_tokens, kv_dim]tile_tokens随kv_dim增大而减小32 → 16 → 4本质是UB 容量约束下的反比关系kv_dim 越大单次 tile 能容纳的 token 越少尾轴约束kv_dim 512 16满足 BF16 对齐要求。3.4 loop valid_shape 处理动态边界num_tokens_loop (num_tokens tile_tokens - 1) // tile_tokens for loop_idx in pypto.loop(num_tokens_loop, namescatter_loop, idx_nameloop_idx, unroll_list[1]): offset loop_idx * tile_tokens actual_tokens (num_tokens - offset).min(tile_tokens) index_view pypto.view(slot_mapping_2d, [tile_tokens, 1], [offset, 0], valid_shape[actual_tokens, 1]) key_view pypto.view(key_2d, [tile_tokens, kv_dim], [offset, 0], valid_shape[actual_tokens, kv_dim]) value_view pypto.view(value_2d, [tile_tokens, kv_dim], [offset, 0], valid_shape[actual_tokens, kv_dim])这是 PyPTO 动态 shape 的标准三件套pypto.loop真实遍历动态轴trip count 是符号表达式(num_tokens tile_tokens - 1) // tile_tokens不能用静态 Python 循环替代pypto.view切 tileshape 参数为编译期常量[tile_tokens, kv_dim][offset, 0]为起始偏移valid_shape标注尾块最后一个 tile 的真实有效 token 数为actual_tokens min(num_tokens - offset, tile_tokens)valid_shape[actual_tokens, kv_dim]让编译器只对有效区域执行计算避免越界访问。3.5 scatter_update 与 .move() 写回key_cache_2d_result pypto.scatter_update(key_cache_2d, -2, index_view, key_view) value_cache_2d_result pypto.scatter_update(value_cache_2d, -2, index_view, value_view) key_cache.move(key_cache_2d_result) value_cache.move(value_cache_2d_result)pypto.scatter_update(input, dim, index, src)dim-2沿倒数第二维散射index形状[tile_tokens, 1]src形状[tile_tokens, kv_dim]返回更新后的 input原地语义不支持 broadcastsrc 与 index 形状必须严格匹配这是 scatter_update 的重要约束.move()负责把 2D 结果写回原始 4D cache tensorshape 转换由.move()自动处理同时确保 buffer 生命周期正确。3.6 Wrapper 与数据流全景wrapper 函数scatter_pa_kv_cache_wrapper直接调用 JIT kernel 并返回原地更新后的(key_cache, value_cache)无需额外创建输出 tensor。README 给出了完整数据流4D key_cache [num_blocks, block_size, num_heads, head_size] ↓ reshape(inplaceTrue) 2D key_cache_2d [num_blocks * block_size, num_heads * head_size] ↓ scatter_update(dim-2, index, src) 2D key_cache_2d_result ↓ .move() 4D key_cache原地更新四、真实算子中的模式应用AT-16 卡片标注的三个 example 在仓库中均有对应源码印证了该模式的普遍性。4.1 GLMAttnFusion融合预注意力glm_attention_fusion_impl.py 在 Phase1 末尾紧贴 RoPE 完成 cache 写入b_ofs bs_idx * bs_tile b_valid (b_scalar - bs_idx * bs_tile).min(bs_tile) index_view pypto.view(index, [bs_tile], [b_ofs], valid_shape[b_valid]) index_view pypto.reshape(index_view, [bs_tile, 1], valid_shape[b_valid, 1]) pypto.set_vec_tile_shapes(bs_tile, 128) key_cache.move(pypto.scatter_update(key_cache_2d, -2, index_view, k_res)) value_cache.move(pypto.scatter_update(value_cache_2d, -2, index_view, v_res))该实现与独立算子版本结构完全一致唯一区别是 index 先做view切 batch tile 再reshape成[bs_tile, 1]然后直接在.move()内联调用 scatter_update——这是更紧凑的写法语义等价。这正是 SK-05 骨架所描述的AT-16 写入紧贴 RoPE 之后、与 Phase2 KV 读完全解耦。4.2 MLAPrologQuantMLA 投影 量化DeepSeek MLA 架构需要把 rope 分量、nope 分量甚至量化 scale 分别散射到不同的 cache 中。mla_prolog_quant_impl.py 展示了一个 loop 内多次 scatter_update的写法index pypto.view(k_cache_index_2d, [tile_bs, 1], [bs_offset, 0]) kr_cache_out[:] pypto.scatter_update(kr_cache, -2, index, k_rope_4d) kv_cache_out[:] pypto.scatter_update(kv_cache, -2, index, k_nope_4d) k_scale_cache_out[:] pypto.scatter_update(k_scale_cache, -2, index, k_scale_4d)此处cache_index为[t, ]INT64 的散射索引对应卡片中的physical_idx同一份 index 复用三次分别写入 rope cache、nope cache 与 scale cache——散射更新天然支持同索引多目标且索引 dtype 从 INT32 到 INT64 都可见于真实代码。4.3 Compressor压缩器状态更新compressor_impl.py 将 scatter_update 封装成 3D 工具函数def scatter_update_3d(input_tensor, index, src): output pypto.scatter_update(input_tensor, -2, index, src)配合pypto.arange(block_size)生成块内索引序列、valid_shape[1, ratio - pos]处理压缩尾部见该文件 L416-L545是 AT-16 在块内滑动窗口式散射场景的变体应用index 不再是全局槽位而是块内列位置沿axis-2逐列覆盖。五、精度验证与测试体系AT-16 模式的正确性由 test_scatter_pa_kv_cache.py 与 scatter_pa_kv_cache_golden.py 共同保障。5.1 Golden 参考实现golden 用纯 PyTorch 索引赋值描述散射语义block_indices slot_mapping_cpu // block_size block_offsets slot_mapping_cpu % block_size golden_key_cache[block_indices, block_offsets, :, :] key_cpu这正是 README 中数学公式的直接翻译作为 kernel 的对照基准。测试用numpy.testing.assert_allclose对比精度标准为atol 0.0001, rtol 0.0078125BFLOAT16 合理容差输出三态标记[PRECISION_PASS]/[PRECISION_FAIL]。5.2 测试用例矩阵与运行方式README 列出三个用例config1_performance_p0num_tokens2633性能、config2_function_p0num_tokens7902功能、config3_boundary_p0num_tokens16384边界公共参数block_size128, num_heads2, head_size256。测试入口支持多种运行模式# 设置空闲 NPU device ID export TILE_FWK_DEVICE_ID0 # 运行所有测试用例默认 NPU 模式 python test_scatter_pa_kv_cache.py # 运行单个测试用例 / 列出用例 python test_scatter_pa_kv_cache.py config1_performance_p0 python test_scatter_pa_kv_cache.py --list # 无 NPU 环境使用 sim 模式 python test_scatter_pa_kv_cache.py --run_mode sim测试数据构造也值得借鉴当num_tokens num_blocks * block_size时用torch.randperm生成无重复的随机槽位映射模拟真实 Paged Attention 中 token 散落在不同物理槽位的场景超出时退化为randint。六、约束、已知限制与性能调优方向6.1 实现约束清单综合源码注释与 README实现 AT-16 模式需遵守以下约束约束项规则dim参数固定-2沿倒数第二维散射不可改为其他维broadcastscatter_update不支持 broadcastsrc 与 index 形状必须严格匹配TileShape 维度TileShape 维度数 src 维度数2D尾轴对齐kv_dim需 16 以满足 BF16 对齐要求view shapeview 的 shape 参数必须全部是 Python intreshape 约束inplaceTruereshape 的输出不能是函数输出参数动态轴loop 必须真实遍历动态轴pypto.loop 符号 trip count配合valid_shape处理尾块6.2 已知限制当前仓库实现README 明确记录四点 P0 限制暂不支持compress_lens_optional、compress_seq_offset_optional、seq_lens_optional压缩特性参数SPEC 中标记为 P2 优先级仅支持 BFLOAT16 单一 dtypenum_heads2, head_size256, block_size128为编译期常量kernel 注解内num_blocks在注解中标为动态、实际按静态值编译。6.3 性能调优方向README 给出的调优旋钮包括Tiling 配置tile_tokens按kv_dim分段调整32/16/4减少大 kv_dim 下的 loop 迭代次数Loop unrollpypto.loop支持unroll_list如[8, 4, 2, 1]GLMAttnFusion 即如此使用展开可降低循环开销需核对循环依赖与尾块UB buffer 管理通过 JIT 的vec_nbuffer_setting调整 Vector buffer 数量提升 UB 内存利用率性能目标README 标注目标为首跑精度成功性能的 2 倍可结合 pypto-op-perf-tune 的经验体系进一步压榨。七、与 AT-17 Block Gather 的读写闭环最后将视野拉回模式体系AT-16 解决写AT-17 Block Table Gather 解决读。AT-17 在 loop 外分配拼装缓冲区kj_assembleloop 内逐块用view(k_cache, [BS, D], [block_idx_valid*BS, offset])零搬运拼装连续 KV 块再用valid_shape[actual_len, D]处理尾块。其实现警示值得所有 PyPTO 开发者牢记本模式的结构机制是 view 拼装零搬运禁止替换为gather_in_l1/gather_in_ub——后者是显式 GM→L1/UB 搬运指令功能等价但机制冲突会引入真实搬运开销。这条警示反向衬托出 AT-16 写路径的特殊性散射更新本质是必须发生的真实写入数据要落盘到物理 cache因此不存在零搬运优化空间而读路径则要极力避免搬运。读AT-17 view 拼装 写AT-16 scatter_update .move()两条路径配合构成了 Paged KV Cache 在 PyPTO 中的完整存取闭环是 GLMAttnFusion、MLAPrologQuant、PageAttnFP8、SparseCompressFA 等生产级算子的共同基石。延伸阅读原子模式总索引见 patterns/atoms/index.mdAT-16 嵌入融合注意力的完整骨架见 SK-05 Fused Pre-Attn尾块 valid_shape 的前置声明规则见 AT-20 Tail Block完整算子实现与测试见 scatter_pa_kv_cache_impl.py、test_scatter_pa_kv_cache.py。【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考