news 2026/9/19 12:39:31

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
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 缓存的指定物理位置
  • tags:scatter
  • flow_pattern:V(纯 Vector 排布)
  • examples:GLMAttnFusion、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, index=physical_idx, src=src_4d) cache_4d.move() # 确保写回

三个要点值得展开:

  1. 纯 V 排布:KV cache 更新是典型的索引驱动写操作,不需要 Cube 参与,整条数据流落在 Vector 单元上,因此 flow_pattern 标为 V。
  2. axis=-2 写入scatter_update固定沿倒数第二维(block 内 offset 维)做散射,index 决定目标物理槽位,src 是被写入的数据。
  3. .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_offset:block 内偏移,由slot_mapping[i] % block_size得到;
  • 每个 token 在 cache 中的"槽位"由slot_mapping(vLLM 风格的 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~16384;num_blocksblock_sizenum_headshead_size均为编译期静态常量(示例中num_blocks=9760, block_size=128, num_heads=2, head_size=256)。
  • 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 → 2D

kv_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, inplace=True) value_cache_2d = pypto.reshape(value_cache, cache_2d_shape, inplace=True) key_2d = pypto.reshape(key, [num_tokens, kv_dim], inplace=True) value_2d = pypto.reshape(value, [num_tokens, kv_dim], inplace=True) slot_mapping_2d = pypto.reshape(slot_mapping, [num_tokens, 1], inplace=True)

这一步把四维物理 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 约束经验:

  • inplace=True的 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_tokenskv_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, name="scatter_loop", idx_name="loop_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 的标准三件套:

  1. pypto.loop真实遍历动态轴:trip count 是符号表达式(num_tokens + tile_tokens - 1) // tile_tokens,不能用静态 Python 循环替代;
  2. pypto.view切 tile:shape 参数为编译期常量[tile_tokens, kv_dim][offset, 0]为起始偏移;
  3. 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(原地语义);
  • 不支持 broadcast:src 与 index 形状必须严格匹配,这是 scatter_update 的重要约束;
  • .move()负责把 2D 结果写回原始 4D cache tensor,shape 转换由.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(inplace=True) 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 MLAPrologQuant(MLA 投影 + 量化)

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.0078125(BFLOAT16 合理容差),输出三态标记[PRECISION_PASS]/[PRECISION_FAIL]

5.2 测试用例矩阵与运行方式

README 列出三个用例:config1_performance_p0(num_tokens=2633,性能)、config2_function_p0(num_tokens=7902,功能)、config3_boundary_p0(num_tokens=16384,边界),公共参数block_size=128, num_heads=2, head_size=256。测试入口支持多种运行模式:

# 设置空闲 NPU device ID export TILE_FWK_DEVICE_ID=0 # 运行所有测试用例(默认 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不支持 broadcast,src 与 index 形状必须严格匹配
TileShape 维度TileShape 维度数 = src 维度数(2D)
尾轴对齐kv_dim需 > 16 以满足 BF16 对齐要求
view shapeview 的 shape 参数必须全部是 Python int
reshape 约束inplace=Truereshape 的输出不能是函数输出参数
动态轴loop 必须真实遍历动态轴(pypto.loop+ 符号 trip count),配合valid_shape处理尾块

6.2 已知限制(当前仓库实现)

README 明确记录四点 P0 限制:

  1. 暂不支持compress_lens_optionalcompress_seq_offset_optionalseq_lens_optional压缩特性参数(SPEC 中标记为 P2 优先级);
  2. 仅支持 BFLOAT16 单一 dtype;
  3. num_heads=2, head_size=256, block_size=128为编译期常量(kernel 注解内);
  4. num_blocks在注解中标为动态、实际按静态值编译。

6.3 性能调优方向

README 给出的调优旋钮包括:

  • Tiling 配置tile_tokenskv_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_assemble,loop 内逐块用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.md;AT-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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/19 12:37:28

向日葵被控服务异常排查指南:从服务状态到网络侧修复完整手册

平时遇到过一次&#xff0c;你就知道“向日葵被控服务异常&#xff0c;暂时无法控制”这句话有多让人上头。本来人在外面&#xff0c;手机掏出来想连回家里或办公室电脑取个文件&#xff0c;结果设备列表里明明显示“在线”&#xff0c;点进去就是这句提示&#xff0c;再点还是…

作者头像 李华
网站建设 2026/9/19 12:35:43

npm cnpm淘宝更新镜像

切换淘宝新镜像 npm config set registry https://registry.npmmirror.com --global查看当前淘宝镜像 npm config get registry

作者头像 李华
网站建设 2026/9/19 12:31:26

风电基础模板厂家价格合理,河北鸿钢模具实力参考

河北鸿钢模具制造有限公司是保定本地专注钢模具生产的源头厂家&#xff0c;始终坚持不做塑料模具&#xff0c;深耕实体制造领域&#xff0c;核心业务覆盖各类混凝土预制钢模具的研发、生产与交付&#xff0c;是资质齐全的市政工程老牌源头厂家&#xff0c;主营检查井模具、风电…

作者头像 李华
网站建设 2026/9/19 12:28:22

Cloudera Manager运维实战:从CMS到Hadoop服务管理的关键操作与API实践

简介&#xff1a;Cloudera Manager是大数据集群统一管理平台&#xff0c;这份日常运维手册正是面向集群管理员、运维工程师及Hadoop生态初学者的实操型文档。文档以图文步骤为主线&#xff0c;完整覆盖登录Cloudera Manager、启停Management Service、批量启停Hadoop全部服务、…

作者头像 李华