尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

PyPTO-Gym 算子实战:DeepSeek V32 MLA 的 Sparse Attention FP8 反量化算子(Ascend 950PR)解析与验证

PyPTO-Gym 算子实战:DeepSeek V32 MLA 的 Sparse Attention FP8 反量化算子(Ascend 950PR)解析与验证 PyPTO-Gym 算子实战DeepSeek V32 MLA 的 Sparse Attention FP8 反量化算子Ascend 950PR解析与验证【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym本指南围绕 PyPTO-Gym 仓库中面向 DeepSeek V32 Multi-head Latent AttentionMLA架构的Sparse Attention Anti-Quantization FP8算子展开完整讲解其产品支持范围、KV Cache 布局、API 签名、5 层嵌套循环计算流程、在线反量化与 Online Softmax 的源码级实现以及测试用例与精度验证方法。读完本文你将掌握该算子在 Ascend 950PR 上的数据编排方式、PyPTO 算子编写范式以及如何复现并扩展其测试。算子背景与定位Sparse Attention Anti-Quantization FP8 是 ops_transformer 目录 下的一个实验性 Transformer 算子专为DeepSeek V32 MLA架构设计运行于华为昇腾 NPU实现Sparse Flash Attention with FP8 Anti-QuantizationFP8 反量化。其核心思路是基于PagedAttention机制通过 top-k 索引从分页 KV cache 中 gather 选定的 KV 条目对 FP8 量化的 key-nope 进行在线反量化per-group FP32 scales组装完整的 Q/K 后执行标准 Attention 计算O softmax(Q K^T / sqrt(d)) V。在 MLA 架构中KV cache 存储的是压缩后的 latent 表示且Value 直接复用反量化后的 key-nope512 维这是该算子与普通 MHA/Flash Attention 在数据流上的根本差异。产品支持情况产品形态支持情况Ascend 950PR支持Atlas A3 训练系列产品 / Atlas A3 推理系列产品不支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品不支持从测试文件的pytest.mark.soc(950)标记及pypto.platform.npuarch DAV_3510判断见 test_sparse_attention_antiquant_fp8.py该算子实际面向 DAV_3510Ascend 950平台验证运行。核心参数与 KV Cache 布局核心参数参数值说明kv_lora_rank512KV latent 维度qk_rope_dim64RoPE 维度head_dim576完整 head 维度512 64nq128Query head 数量n_kv1KV head 数量GQAtopk2048每个 token 选取的 top-k KV 数量block_size128PagedAttention block 大小其中kv_lora_rank512与qk_rope_dim64是面向 DeepSeek V32 MLA 的固定参数决定了 cache 中每个 token 的字节布局与反量化分组方式。KV Cache 布局nope_cache每个 token 在nope_cache中占656 字节padding 至672 字节对齐从 FP8 视角观察其数据编排如下字节偏移内容逻辑 dtype维度[0:512]kv_nope量化后FP8_E4M3512[512:640]key_ropeBF16以 FP8 视角存储64128 bytes[640:656]dequant scalesFP32以 FP8 视角存储416 bytes反量化方式将 512 维 kv_nope 按 4 组每组 128 个元素分组每组乘以对应的 FP32 scale。该布局在实现中体现为nope_cache的形状(block_num * block_size, kv_lora_rank rope_dim*2 4*4)见 sparse_attention_antiquant_fp8_impl.py 的 docstring672 512 64*2 4*4与 README 的字节表完全对应。同时 672 字节的 padding 对齐属于硬编码要求调整该值需同步修改实现与 Golden 参考。API 签名与 Decode / Prefill 差异算子提供两个入口sparse_attention_antiquant_dDecode与sparse_attention_antiquant_pPrefill两者签名一致sparse_attention_antiquant_d( query_nope, # (t*nq, 512) BF16 query nope 部分 query_rope, # (t*nq, 64) BF16 query rope 部分 nope_cache, # (block_num*bs, 672) FP8 分页 KV cache含 kn kr scales topk_indices, # (t, n_kv*topk) INT32 top-k 选取的 token 索引 block_table, # (b, max_blocknum) INT32 PagedAttention block 映射表 kv_act_seqs, # (b,) INT32 每个 batch 的实际序列长度 attention_out, # (b*s*nq, 512) BF16 输出 nq, n_kv, softmax_scale, topk, block_size, max_blocknum_perbatch, tile_config )在仓库当前实现中实际对外入口是 sparse_attention_antiquant_fp8_impl.py 中经pypto.frontend.jit装饰的sparse_attention_antiquant_fp8_high它通过pypto.experimental.set_operation_options(combine_axisTrue)开启轴合并后调用核心计算函数sparse_attention_antiquant_compute_950。该入口的张量签名中query_nope/query_rope/topk_indices/attention_out首维为pypto.DYNAMIC动态轴nope_cache与其余维度为pypto.STATIC静态轴并通过nq, n_kv, softmax_scale, topk, block_size, max_blocknum_perbatch, tile_config作为编译期标量参数传入。Decode vs Prefill 差异配置项Decode (_d)Prefill (_p)vec_nbuffer_setting{-1: 2, 0: 4}{-1: 4, 0: 4}cube_l1_reuse_setting{-1: 2}{-1: 4}device_sched_mode3未设置Decode 与 Prefill 的差异本质来自二者访存与计算特征的差异Prefill 侧 KV 序列更长、计算密度更高因此配置更大的 nbuffer{-1: 4}与 L1 复用{-1: 4}以摊薄访存开销Decode 侧则侧重低延迟调度device_sched_mode3。当前实现中sparse_attention_antiquant_fp8_high的 jit 默认配置为ooo_sched_modeGAPMIN、cube_l1_reuse_setting{-1: 1}、cube_nbuffer_setting{-1: 1}并设置stitch_function_max_num128、device_sched_mode1、ready_on_host_tensors[block_table, kv_act_seqs]、max_workspace_kb1607648等运行时选项说明 Decode/Prefill 差异配置可通过不同的 jit 装饰器参数注入。计算流程5 层嵌套循环算子采用 5 层嵌套循环结构遍历数据L0: batch_idx — 遍历 batch L1: slc_idx — 遍历 query 序列 L2: n_kv_idx — 遍历 KV head L3: group_idx — 遍历 query groupnq / n_kv / g_tile L4: s2_idx — 遍历 KV 序列 tile在 sparse_attention_antiquant_fp8_impl.py 中L0/L1/L2/L3 分别对应LOOP_L0_idx、LOOP_L1_s1_SA、LOOP_L2_n_kv_SA、LOOP_L3_g_SAL4 为LOOP_L4_s2_SA且带unroll_list[4]展开。每次 L4 迭代计算一个(cur_group_tile, dn)的输出分块dn 即kv_lora_rank512。每次 L4 迭代内部的计算步骤与 README 一致V0 Gather通过gather_in_ub根据 topk_indices block_table 从 nope_cache 中取出选定的 KV 条目Dequant提取 FP8 kn → cast FP32 → 乘以 per-group scales → cast BF16组装 K拼接 kn(512) kr(64) → kj(576)组装 Q拼接 qn(512) qr(64) → qi(576)C1 MatMulsij qi kj^TFP32 累加shape(g_tile, s2_tile)V1 Softmaxscale → amax → sub → exp → sum → div → cast BF16C2 MatMulq1 softmax vj其中 vj kn512 维输出 BF16写出将 q1 写入 attention_out。关键实现细节源码级动态序列裁剪L1 内通过cur_seq (cur_act_seq - s1_sym 1 slc_idx).max(0).min(topk)计算当前 query 实际可 attend 的 KV 数量再据此推导bn_per_batch避免越界GQA 分组group nq // n_kvL3 的迭代次数为g_loop_sym ceil(group / g_tile)每轮通过cur_group_tile min(group, group_tile)处理尾部Q/K 组装使用pypto.assemble将 qn/qr 拼入预分配的qi、将 kn/kr 拼入kj其中kj与vj共享同一份反量化后的 kn 数据vj是kj_view的前 512 列视图这正是Value 复用 key-nope的 MLA 特性在代码层面的落点。在线反量化与 Online Softmax 的实现细节FP8 在线反量化Dequant实现中反量化并非一次性完成整行 512 维而是分两步利用向量指令对齐kn_quant_fp32 pypto.cast(kn_quant, pypto.DT_FP32) # FP8 - FP32 kn_quant_fp32_tmp pypto.reshape(kn_quant_fp32, [s2_tile * 4, 128]) # 拆成 4 组 kn_scale_tmp pypto.reshape(kn_scale, [s2_tile * 4, 1]) # 每组 1 个 scale kn_fp32 pypto.mul(kn_quant_fp32_tmp, kn_scale_tmp) # 逐组乘 scale kn pypto.cast(kn_fp32, dtype) # FP32 - BF16即将(s2_tile, 512)的 FP8 量化 kn 重排为(s2_tile*4, 128)与重排为(s2_tile*4, 1)的 FP32 scale 逐元素相乘再还原回(s2_tile, 512)并 cast 为 BF16——与 README 的4 组 × 128 元素反量化描述严格一致。该流程的数值语义与测试 Golden 中的slc_kv_fp32 slc_kv_fp8.reshape(-1, 128).to(torch.float)与slc_kv slc_kv_fp32 * slc_kv_scales完全对应。Online Softmax 的跨 tile 增量更新虽然 L4 步骤 6 描述为scale → amax → sub → exp → sum → div但源码实际实现了online softmaxflash 式增量更新当topk2048大于单个s_kv_tile时需要跨多个 L4 迭代逐步累积输出。源码中维护oi_update、sum_update、max_update三个 FP32 状态张量并在每次 L4 迭代通过pypto.cond(pypto.is_loop_begin(s2_idx))/pypto.cond(pypto.is_loop_end(s2_idx))区分首块、中间块与末块首块oi_update q1记录sum_update与max_update中间块max_new max(max_update, tilda_mij)按exp(max_old - max_new)与exp(tilda_mij - max_new)分别缩放旧/新累积值后相加末块oi_update oi_tmp / sum_update使用pypto.PrecisionType.INTRINSIC精度的除法cast 为 BF16 后pypto.assemble写回attention_out。从测试源码看Golden 参考函数compute_attention_aq同样实现了mi/li/oi的 online softmax 增量更新见 test_sparse_attention_antiquant_fp8.py 中s2_idx 0分支与mi_new/scale_old/scale_new更新逻辑可据此推断当前测试 Golden 已采用 flash/online softmax 语义与 README 中Golden 使用 per-tile softmax的历史约束说明可能存在版本差异精度对齐时需以实际测试代码为准。Tiling 配置Tiling 通过SaTileShapeConfig数据类配置dataclass class SaTileShapeConfig: g_tile: int # Group tile 大小如 128 s_kv_tile: int # KV 序列 tile 大小如 2048 c1_tile_shape: list # 6 个 intC1 MatMul cube tile v1_tile_shape: list # 2 个 intV1 Softmax vector tile c2_tile_shape: list # 6 个 intC2 MatMul cube tile v2_tile_shape: list # 2 个 int已定义但未使用README 给出的典型配置值为g_tile128, s_kv_tile2048, cube tiles 128x128, vector tiles 8x2048 / 64x128而当前测试文件实际使用的配置为SaTileShapeConfig( g_tile128, s_kv_tile512, # 与 README 典型值 2048 不同 c1_tile_shape[128, 128, 256, 256, 128, 128], v1_tile_shape[64, 128], c2_tile_shape[128, 128, 128, 128, 256, 256], v2_tile_shape[128, 128] )其中c1/c2_tile_shape的 6 个 int 对应 cube 指令的三组 tile 形状M、N、K 方向通过pypto.set_cube_tile_shapes注入 C1/C2 MatMulv1_tile_shape通过pypto.set_vec_tile_shapes控制 softmax 的向量分块。实现中还存在固定的v2_2_tile [32, 512]用于 online softmax 更新阶段的部分向量运算。需要说明s_kv_tile 需要与 topk 的整除关系匹配它直接决定 L4 的迭代次数bn_per_batch ceil(cur_seq / s2_tile)。测试用例与精度验证运行方式README 提供的运行命令为pytest deepseekv32_sparse_attention_antiquant_fp8.py -v在当前仓库中对应的测试文件实际位于 tests/ops/experimental/ops_transformer/sparse_attention_antiquant_fp8/test_sparse_attention_antiquant_fp8.py可直接执行pytest tests/ops/experimental/ops_transformer/sparse_attention_antiquant_fp8/test_sparse_attention_antiquant_fp8.py -v测试通过TILE_FWK_DEVICE_ID环境变量指定 NPU 设备默认 0使用torch.npu.set_device完成设备切换。测试矩阵README 记载的测试矩阵用例名(b, nq, n_kv, s_q)actual_seq模式sfa_bf16_b4_s2_seq64K_total_fp8_d(4, 128, 1, 2)[65536, 16381, 666, 15]Decodesfa_bf16_b4_s2_seq64K_per_fp8_d(4, 128, 1, 2)[65536]*4Decode性能测试默认 skipsfa_bf16_b1_s256_seq64K_fp8_p(1, 128, 1, 256)[65536]Prefill默认 skip当前测试源码中的实际用例配置略有差异均标记pytest.mark.soc(950)用例名(b, nq, n_kv, s_q)actual_seq备注test_sfa_bf16_b4_s2_seq64k_per_fp8_d(4, 128, 1, 2)[65536]*4pytest.mark.skip大用例test_sfa_bf16_b64_s2_seq64k_per_fp8_d(64, 128, 1, 2)[65536]*64默认执行test_sfa_bf16_b64_s2_seq64k_uniform_per_fp8_d(64, 128, 1, 2)8 种长度均匀分布16384~114688pytest.mark.skip大用例其中sfa_bf16_b64_s2_seq64k_uniform_per_fp8_d的 actual_seq 为[16384]*8 [32768]*8 [49152]*8 [65536]*16 [81920]*8 [98304]*8 [114688]*8用于验证不均匀变长序列下的正确性。测试用例以do_test_sfa_entry为统一入口先生成 Goldengen_gather_select_attention_golden包含 FP8 量化、随机 block_table 打乱、top-k 索引构造等步骤再调用算子入口对比输出。精度验证标准atol 0.0001rtol 0.005max_error_count 100验证通过 common_utils/compare 中的比较工具完成对(b, s1, nq, kv_lora_rank)形状的输出张量逐元素比对。支持芯片Ascend 910Ascend 950注意与产品支持情况中 Ascend 950PR 支持的表述对应测试标记限定soc(950)。约束与注意事项Golden 参考实现的 softmax 语义README 说明 Golden 参考实现使用 per-tile softmax非 flash/online softmax当 topk 超过单个 s2_tile 时精度对齐可能存在偏差从当前测试源码看compute_attention_aq已实现 online softmax 增量更新实际对齐时以测试代码为准672 字节 padding 对齐是硬编码要求nope_cache单行 672 字节的布局512 128 16贯穿实现、Golden 与量化流程不可随意变更固定参数算子专为 DeepSeek V32 MLA 架构定制kv_lora_rank512与qk_rope_dim64为固定参数Value 复用Value 复用反量化后的 key-nopekv_lora_rank512维这是 MLA 架构的特性与普通 attention 的独立 V 投影不同平台限制算子仅在 Ascend 950PRDAV_3510上验证A2/A3 系列产品不支持。相关算子对比kv_split 变体同一目录下存在姊妹算子 sparse_attention_antiquant_kv_split二者算法与循环结构一致关键差异在于cache 的组织方式本算子将 kn、kr、scales 打包进单个nope_cache672 字节/行而 kv_split 变体将其拆分为三个独立张量kn_quantFP8, 512、krBF16, 64、kn_scalesFP32, 4并相应放宽了672 字节硬编码约束。对比阅读两份 README 有助于理解数据布局选择对算子访存模式的影响。总结Sparse Attention Anti-Quantization FP8 算子完整展示了 PyPTO 在分页稀疏注意力 FP8 反量化 MLA Value 复用场景下的编写范式以SaTileShapeConfig驱动 5 层循环与 cube/vector 双流水 tiling用gather_in_ub完成分页 gather用 FP32 中间态保证反量化与 online softmax 的数值精度最终以 BF16 输出。其测试体系Golden 参考 精度容差 多 batch/变长序列矩阵可作为在 PyPTO-Gym 中新增算子验证用例的参考模板。若需进一步理解 PyPTO 的 jit 配置与向量/立方指令约束可查阅 PyPTO 算子开发相关文档 与 ops_transformer 目录下的其他 README。【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表