PyPTO 实现 GLM-4.5 Paged Attention FP8 量化算子:分页 KV Cache 与 E4M3 混合精度 Kernel 实战解析
发布时间:2026/9/18 3:03:11来源:尧图网络
PyPTO 实现 GLM-4.5 Paged Attention FP8 量化算子分页 KV Cache 与 E4M3 混合精度 Kernel 实战解析【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym本文以 PyPTO-Gym 仓库中的 Page Attention Quant FP8 算子文档 为主体结合 Kernel 实现与测试源码系统讲解如何在 Ascend NPU 上基于 PyPTO 框架实现面向 GLM-4.5 的 Paged Attention FP8 量化算子。读完本文你将掌握 Paged KV Cache 的 block_table 分页管理、FP8 E4M3 的 per-token/per-channel 量化策略、Flash Attention Online Softmax 的增量式累加以及 PyPTO Kernel 的循环结构、tiling 配置与精度校验全流程。算子背景与产品支持情况该算子是基于 PyPTO 框架实现的 Paged Attention FP8 量化 Kernel运行于 Ascend NPU服务 GLM-4.5 模型的推理场景。其核心思路借鉴操作系统内存分页思想将 KV cache 按固定 block 切分并以非连续内存布局管理从而高效处理变长序列与动态批次同时引入 FP8 E4M3 量化降低内存占用与计算开销。产品支持情况与仓库文档一致产品支持情况Ascend 950PR支持Atlas A3 训练系列产品 / Atlas A3 推理系列产品不支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品不支持平台层面Kernel 面向DAV_3510Ascend 950架构测试用例通过pytest.mark.soc(950)标注适用芯片参见 测试文件 与 pytest.ini 中 soc marker 的定义conftest.py 会根据当前设备的 soc_version 自动过滤不适配的用例。文件说明文件说明page_attention_quant_fp8_impl.pyKernel 实现 量化/反量化辅助函数test_page_attention_quant_fp8.py测试用例 Golden reference 实现仓库中二者分处src/与tests/两个目录实现位于src/pypto_gym/ops/pypto_tensor/experimental/ops_transformer/page_attention_quant/测试位于tests/ops/experimental/ops_transformer/page_attention_quant/符合仓库算子实现与测试一一对应的目录约定见 README.md 中的目录结构说明。算法概述算子实现 GLM-4.5 模型的 Paged Attention 机制采用分页内存管理策略高效处理变长序列和动态批次同时使用 FP8 (E4M3) 量化降低内存占用和计算开销。语义约定Q 侧 (s1_size)Query 序列长度生成阶段通常为 1KV 侧 (s2_size)KV cache 序列长度可变支持长序列Paged KVKV cache 按 blockblock_size128分页管理通过 block_table 映射FP8 Quantization使用 FP8 E4M3 格式量化scale 采用 per-token 或 per-channel 方式GQA (Grouped Query Attention)nq 个 query head 共享 nkv 个 KV headgroup nq // nkvG_TILEGroup tile 大小每次处理的 query head group 数量S2_TILEKV 序列分块大小用于迭代计算。从源码看这些语义在 Kernel 实现 中均有对应group nq // nkv、n2_sym nkv、g_tile tile_config.g_tile、s2_tile tile_config.s2_tileKernel 通过kv_act_seqs[b_idx]动态获取每个 batch 的实际序列长度。循环结构batch_loop (parallel) — 遍历 batch从 kv_act_seqs 动态获取实际序列长度 s1_loop — 遍历 s1 (query 序列维度) n2_loop — 遍历 KV head (n2_sym nkv) g_loop — 遍历 group每次处理 g_tile 个 query head s2_loop (unroll 8,4,2,1) — KV 序列按 s2_tile 分块迭代源码中对应实现为pypto.loop的五层嵌套LOOP_b、LOOP_s1、LOOP_n2、LOOP_g、LOOP_s2其中 s2 循环通过unroll_list[8, 4, 2, 1]指定展开策略。值得注意的是batch 维度虽然标注为(parallel)实际 Kernel 中 B 是动态轴pypto.DYNAMIC通过pypto.loop(b_scalar)展开遍历s2_loop的迭代次数由当前 batch 的实际序列长度动态决定s2_loop (cur_seq s2_tile - 1) // s2_tile。计算流程 (per batch, per s1, per group)1. 查询分块 Q_tile [g_tile, D] FP8 E4M3 scale 2. 组装 KV 分块 (通过 block_table) - K_assemble [s2_tile, D] FP8 E4M3 scale - V_assemble [s2_tile, D] FP8 E4M3 scale 3. MM1: S_tile Q_tile K_assemble^T [g_tile, s2_tile] FP8 matmul → FP32 4. Dequant: S_fp32 S_tile * Q_scale * K_scale^T [g_tile, s2_tile] FP32 5. Flash Attention Online Softmax: - Scale: S_scaled S_fp32 * softmax_scale [g_tile, s2_tile] - Update: M_new max(M_old, max(S_scaled)) [g_tile, 1] - Exp: P exp(S_scaled - M_new) [g_tile, s2_tile] - Accum: Sum_new Sum_old * exp(M_old - M_new) sum(P) [g_tile, 1] 6. Quant: P_fp8, P_scale quant_per_token(P) [g_tile, s2_tile] FP8 scale 7. MM2: O_tile P_fp8 V_assemble [g_tile, D] FP8 matmul → FP32 8. Dequant: O_fp32 O_tile * P_scale * V_scale [g_tile, D] FP32 9. Accum: O_accum O_accum * exp(M_old - M_new) O_fp32 [g_tile, D] 10. Final: O_final O_accum / Sum_new [g_tile, D] FP32 → BF16源码中第 2 步组装 KV 分块通过kj_assemble/vj_assemble张量配合pypto.view完成每个 s2_tile 由block_num s2_tile // block_size个 block 组成逐个从block_table[b_idx, idx i]读取全局 block_id非法值为 -1 时用block_idx.max(0)收敛为 0再以pypto.view(k_2d, [block_size, dn], [block_idx_valid * block_size, n2_idx * dn])切出对应 KV head 的数据最后通过valid_shape[actual_s2_tile, dn]只保留有效 token 部分从而支持变长序列。Online Softmax 的 M/Sum 状态max_update、sum_update、oi_update均为[g_tile, 1]/[g_tile, dn]大小的 FP32 张量跨 s2_tile 增量更新最终在最后一个 s2 tile 归一化后 cast 为 BF16 写入atten_out。Kernel 签名与源码对应关系原文档给出的 Kernel 签名为ifa_func_kernel( q, # [B*S1, NQ, D] FP8E4M3 — Query 输入 q_scale, # [B, NQ, 1] FP32 — Query scale (per-token) k, # [num_blocks, block_size, NKV, D] FP8E4M3 — Key cache (paged) k_scale, # [num_blocks, block_size, NKV, 1] FP32 — Key scale (per-token) v, # [num_blocks, block_size, NKV, D] FP8E4M3 — Value cache (paged) v_scale, # [B, NKV, D] FP32 — Value scale (per-channel) block_table, # [B, max_blocks_per_query] INT32 — Block mapping table kv_act_seqs, # [B] INT32 — 每个 batch 的实际 KV 序列长度 atten_out, # [B*S1, NQ, D] BF16 — Attention 输出 )需要说明的是当前仓库源码中对应的 JIT Kernel 函数实际命名为pfa_func_kernel_v2_bound见 page_attention_quant_fp8_impl.py其参数语义与上述签名完全一致且所有张量维度声明为动态轴pypto.DYNAMICdtype 由pypto.DT_FP8E4M3、pypto.DT_FP32、pypto.DT_INT32、pypto.DT_BF16约束Kernel 末尾额外携带softmax_scale与tile_config两个配置参数。阅读源码时请以实际函数名为准。其中B(batch_size) 为动态轴 (pypto.DYNAMIC)S1、NQ、NKV、D为静态轴。Kernel 入口处将三维/四维输入 reshape 为二维形式以便于 view 操作源码中通过pypto.reshape(..., inplaceTrue)完成Q:[B*S1*NQ, D]q_2d_shape (b_scalar * s1_scalar * nq, dn)K/V:[num_blocks*block_size, NKV*D]k_2d_shape (block_num_scalar * block_size, n2_sym * dn)Kernel 通过传入的kv_act_seqs反推 batch 数b_scalar进而由s1_scalar bs_scalar // b_scalar得到 Query 序列长度实现对动态 shape 的推导。Dtype 转换流程与量化策略Kernel 内部严格控制 FP8/FP32/BF16 转换以平衡精度和性能阶段操作Dtype输入Q/K/VFP8 E4M3输入Q_scale/K_scale/V_scaleFP32MM1Q K^T → S_quantFP8 → FP32Dequant MM1S_quant * Q_scale * K_scale^TFP32 全程Softmaxscale/max/expFP32 全程Quant PP → P_fp8FP32 → FP8 E4M3MM2P_fp8 V → O_quantFP8 → FP32Dequant MM2O_quant * P_scale * V_scaleFP32 全程AccumO_accum (跨 s2_tile 累加)FP32 全程输出O_final cast BF16FP32 → BF16上述流程在源码中的关键节点均有直接对应MM1 为pypto.matmul(qi, kj_assemble, pypto.DT_FP32, a_transFalse, b_transTrue)反量化由dequant_dynamic完成先 cast 到 FP32再依次乘以两个 scalesoftmax 部分依次执行mul(softmax_scale)、amax、sub、exp、sumP 的量化复用symmetric_quantization_per_token_fp8_e4m3MM2 为pypto.matmul(tilda_pij_fp8_e4m3, vj_assemble, pypto.DT_FP32)输出在pypto.div(..., precision_typepypto.PrecisionType.INTRINSIC)归一化后 cast 为输出 dtype。量化策略Query 量化 (per-token)每个 query token 独立计算 scalescale 448.0 / max(|Q|)FP8 E4M3 最大值为 448输出: Q_fp8, Q_scaleKey 量化 (per-token)每个 KV token 独立计算 scalescale 448.0 / max(|K|)输出: K_fp8, K_scaleValue 量化 (per-channel)每个 channel (head_dim) 共享 scalescale 448.0 / max(|V|, dimseq)输出: V_fp8, V_scaleAttention P 量化 (per-token)每个 softmax 输出 token 独立计算 scalescale 448.0 / max(|P|)输出: P_fp8, P_scale从测试 Golden 实现可以印证测试文件分别提供quant_fp8e4m3_per_tokenQ/P、quant_fp8e4m3_per_token_keyK、quant_fp8e4m3_per_channel_valueV按dim1即序列维度求最大值三个量化函数与 Kernel 内symmetric_quantization_per_token_fp8_e4m3的实现保持一致先求max(|x|)用448.0 / max得到量化 scale乘到数据上再 cast 为torch.float8_e4m3fn反量化 scale 为1.0 / scale。测试用例与精度校验原文档描述测试通过独立的test_ifa_XX函数定义每个函数指定不同的batch_size、s1_size、s2_size用例batchs1s2说明test_ifa_011618192默认配置 (大批次)test_ifa_02818192小批次 (skip)需要说明的是当前仓库的 测试文件 已演进为pfa_test_impl(case_name)test_pfa_for_950()的结构入口函数pfa_test_impl通过get_case_config(case_name)从 Kernel 实现 内置的 6 个 case 配置中选取参数test_pfa_for_950则一次性串行跑完全部 6 个用例覆盖如下形状组合case_namebatchs1s2nqnkvpfa_fp8_b1_s1_2_s2_2048_nkv_8122048488pfa_fp8_b16_s1_1_s2_8195_nkv_21618195122pfa_fp8_b16_s1_1_s2_8k_nkv_11618192121pfa_fp8_b16_s1_1_s2_8k_nkv_21618192122pfa_fp8_b2_s1_1_s2_1k211024121pfa_fp8_b16_s1_1_s2_256_nkv_4161256164测试同时验证了两个约束atten_cfg.b len(atten_cfg.actual_seq)b 与 actual_seq 长度相等以及所有 actual_seq 值均小于等于 s2。此外kv_num_blocks b * ceil(s2 / block_size)用于分配分页 KV cache 的 block 总量。精度校验使用numpy.testing.assert_allclose进行精度验证rtol 0.0078125 # 1/128考虑 FP8/BF16 混合精度误差 atol 0.001测试文件中的实际实现为assert_allclose(..., rtol0.0078125, atol0.0001)atol 较文档更严格一个数量级以源码为准。Golden referencepfa_flash_torch模拟完整的 Paged Attention FP8 量化流程按 block_table 组装 KV cachegen_block_table生成随机打乱的 block 映射表未使用部分以 -1 填充kv_cache_concat_bsnd将分页格式还原为 BSND 连续布局fp8_bsnd_to_pa_format再把 KV 数据按全局 block_id 写回[total_blocks, 128, n, d]的分页张量执行 QK matmul softmax PV matmul_assemble_kv_for_s2_tile按 s2_tile 组装 K/K_scale/V_process_s2_tile完成 MM1 反量化、softmax、P 量化、MM2 反量化在关键节点进行 FP8 量化/反量化与 kernel 内部流程一致Online Softmax 状态更新由_flash_update_state复现与 Kernel 中is_loop_begin/is_loop_end分支逻辑一一对应。运行方式# 设置设备 ID export TILE_FWK_DEVICE_ID0 # 运行全部测试用例 python test_page_attention_quant_fp8.py注意测试文件通过向上查找含src目录的仓库根路径并插入sys.path随后以from experimental.ops_transformer.page_attention_quant.page_attention_quant_fp8_impl import ...导入 Kernel因此建议在仓库根目录下运行直接执行时入口会检查pypto.platform.npuarch DAV_3510后再调用test_pfa_for_950()。使用 pytest 运行特定用例# 运行 test_pfa_for_950 pytest -v tests/ops/experimental/ops_transformer/page_attention_quant/test_page_attention_quant_fp8.py::test_pfa_for_950 # 运行全部跳过不适配当前 SoC 的用例 pytest -v tests/ops/experimental/ops_transformer/page_attention_quant/test_page_attention_quant_fp8.py仓库 conftest.py 支持--device N参数覆盖TILE_FWK_DEVICE_ID例如pytest ... --device 1pytest.ini中testpaths tests/ops定义了默认测试路径socmarker 用于按芯片筛选用例。添加新用例原文档给出了通过ifa_test_impl添加用例的模板pytest.mark.soc(950) def test_ifa_03(): batch4, s12, s24096 ifa_test_impl(b4, s12, s24096)ifa_test_impl支持的可选参数参数说明默认值b批次数量16s1Query 序列长度1s2KV cache 序列长度8192对照当前源码测试入口已重构为按 case_name 驱动扩展新用例的推荐做法是在 page_attention_quant_fp8_impl.py 的get_case_config字典中新增一个create_config(b..., s1..., s2..., nq..., nkv..., qd..., block_size...)条目再在测试文件中将 case_name 加入test_pfa_for_950的列表即可参数校验b 与 actual_seq 长度一致、actual_seq 不超过 s2会自动生效。注意b必须与atten_cfg.actual_seq的长度相等所有actual_seq值必须小于等于s2。配置说明原文档给出AttentionConfig与AttentionTileConfig两个 dataclass当前源码中对应的类名分别为PfaConfig与PfaTileShapeConfig见 page_attention_quant_fp8_impl.py字段语义一致dataclass class PfaConfig: b: int 8 # batch_size s1: int 1 # query 序列长度 s2: int 16384 # KV cache 最大序列长度 nq: int 12 # query head 数量 nkv: int 1 # KV head 数量 (GQA) qd: int 128 # head dimension kvd: int 128 # KV head dimension block_size: int 128 # 分页 block 大小 softmax_scale: float # softmax scale 1/sqrt(d) kv_layout: str PA_BSND # KV cache 布局格式 actual_seq: torch.Tensor # 每个 batch 的实际序列长度dataclass class PfaTileShapeConfig: g_tile: int # Group tile 大小源码 create_config 中动态计算为 nq // nkv s2_tile: int # KV 序列 tile 大小s28192 时取 1024否则默认 128 c1_tile_shape: list # MM1 cube tile shapes v1_tile_shape: list # Vector tile shapes (softmax) c2_tile_shape: list # MM2 cube tile shapes v2_tile_shape: list # Vector tile shapes (output)build_pfa_config负责将 case 配置落地为PfaConfig实例softmax_scale qd ** -0.5kv_num_blocks b * ceil(s2 / block_size)max_num_blocks_per_query ceil(s2 / block_size)actual_seq默认全部填充为 s2 的 int32 张量置于npu:${TILE_FWK_DEVICE_ID}设备上。默认的 tile shapes 为c1/c2_tile_shape [[128, 128], [128, 128], [128, 128]]v1_tile_shape [128, s2_tile]v2_tile_shape [128, 256]。关键特性Paged KV CacheBlock Size: 128 tokens per blockBlock Table: 每个 batch 维护一个 block 映射表支持不连续内存布局测试中gen_block_table用torch.randperm随机打乱全局 block 编号以验证非连续场景无效位置以 -1 填充动态序列长度: 通过kv_act_seqs支持变长序列Kernel 内以valid_shape裁剪每个 s2_tile 的无效 token。FP8 Quantization低精度计算: 使用 FP8 E4M3范围 [-448, 448]降低内存占用混合精度策略: Matmul 使用 FP8累加/softmax 使用 FP32 保证精度量化位置选择:Q/K/V 输入量化Q/K 为 per-tokenV 为 per-channelAttention P 中间量化降低 MM2 计算开销per-token。Flash Attention Online SoftmaxOnline Algorithm: 采用增量式 softmax避免存储完整的注意力矩阵内存优化: 只维护[g_tile, 1]的 max 和 sum不存储[g_tile, s2_tile]的完整 P 矩阵精度保证: FP32 累加最终输出 BF16。性能优化选项Kernel 通过pypto.frontend.jit的runtime_options与pass_options配置编译期优化策略源码中的实际配置为pypto.frontend.jit( runtime_options{stitch_function_max_num: 512, device_sched_mode: 0}, # 当子图大小达到上界不允许与其他子图合并 pass_options{ # Q 常驻 L1-1 代表所有子图8 代表 8 次 matmul 合并 cube_l1_reuse_setting: {-1: 8}, vec_nbuffer_setting: {-1: 4}, # Vector 四缓冲 cube_nbuffer_setting: {-1: 4} # Cube 四缓冲 }, verify_options{ enable_pass_verify: False, pass_verify_save_tensor: False, }, )原文档示例中的参数含义可对照理解stitch_function_max_num控制可拼接子图数量上限过大时会显著增加 workspace 占用仓库 README.md 常见问题一节给出workspace totalSlot × (stitch_function_max_num 1) × parallelism的估算公式内存不足时可下调该值cube_l1_reuse_setting: {0: 4}表示让 Q 常驻 L1、合并 4 次 matmulvec_nbuffer_setting/cube_nbuffer_setting配置 Vector/Cube 的多缓冲深度。注意事项Block Table 管理: block_table 中的 block_id 必须 0 且 num_blocks-1 表示无效 blockKernel 内通过block_idx.max(0)将 -1 收敛为 0 以规避越界读取但语义上仍要求有效位置正确填充序列长度一致性: kv_act_seqs 的值不能超过 s2最大序列长度量化 scale 非零: 量化前需确保输入 tensor 非零避免 scale 为无穷大内存布局: 所有输入 tensor 必须为 ND 格式非 NZ。测试中的check_args会逐一校验 query/key_cache/value_cache/block_tables/actual_seqs/attn_res 的维度、ND 格式与 dtypeFP8E4M3、INT32、BF16不符合条件直接抛出ValueError。依赖与平台支持运行本算子需要Python 3.xPyTorch torch_npuPyPTOpypto包pypto.frontend.jit与pypto.loop、pypto.matmul、pypto.view等 Tile 编程原语NumPyGolden reference 与精度断言pytest用例执行与 soc 筛选平台支持DAV_3510Ascend 950对应产品为 Ascend 950PRAtlas A2/A3 系列当前不支持。小结本文以 PyPTO-Gym 仓库的 Page Attention Quant FP8 算子为实例完整梳理了从算法语义、五层循环结构、十步计算流程、Kernel 签名与 dtype 流转到量化策略、测试 Golden 校验、tiling 配置与编译期优化选项的全链路。该算子同时展示了 PyPTO 支撑长序列 LLM 推理的三类关键能力分页 KV cache 的非连续内存管理、FP8 E4M3 低比特量化与 FP32 精度保障的混合精度设计以及 Flash Attention Online Softmax 的增量式内存优化可作为在 Ascend NPU 上开发同类 Transformer 注意力算子的参考模板。【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
网站建设高端定制企业官网