资讯动态

AT-12 QK^T MatMul + Scale:PyPTO 注意力分数计算的 Cube 侧核心模式与 FP8 变体实战

发布时间:2026/9/19 23:21:06 来源:尧图企业网站定制
AT-12 QK^T MatMul ScalePyPTO 注意力分数计算的 Cube 侧核心模式与 FP8 变体实战【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gymAttention自注意力前向计算中第一步是用 Query 与 Key 做矩阵乘得到原始注意力分数。在 PyPTO 算子设计中这一步被抽象为局部计算模式Atom PatternAT-12: QK^T MatMul Scale它描述了Q K^T的 Cube矩阵乘计算流、scale 缩放语义、实例化参数以及 FP8 输入下的反量化变体。本文以 AT-12-qk-matmul.md 为主体骨架结合本仓库的 Flash Attention 真实实现 flash_attention_mha_impl.py 与设计工作流 pypto-op-design讲解该模式的排布定位、计算流、参数化方法、FP8 变体以及它在 C1-V1-C2 注意力骨架中的编排要点帮助读者在设计阶段直接套用并写出可验证的 QK^T 计算段。模式定位Attention 计算链上的第一个 Cube 节点在 PyPTO 的算子设计体系中计算模式被划分为**骨架SkeletonSK与局部原子模式AtomAT**两层SK 描述整个算子的循环组织与阶段编排AT 描述其中一段具体的局部计算。AT-12 属于 AT 索引 中的 matmul 类模式flow_pattern为CCube即它完全落在 Cube 计算单元上不涉及 Vector 操作。AT-12 的典型使用场景在模式卡片的examples字段中写得很明确所有 Attention 算子。它在 Flash Attention / IFA / PageAttn / SparseAttn 等注意力算子中扮演的是 SK-01Online Flash Attention 骨架中C1 阶段的角色处于如下计算链的起点C1(QK^T) → V1(Softmax) → C2(PV)前置阶段Q/K 经过 AT-03 RMSNorm 或 RoPE 后进入 QK^T后继阶段AT-12 的输出scores直接喂给 AT-01 Online SoftmaxV1 阶段做逐块最大值/指数和累积下游镜像softmax 归一化后的概率 P 再与 V 做矩阵乘即姊妹模式 AT-13 PV MatMulC2 阶段。理解 AT-12 的关键在于它输出的不是最终的注意力权重而是原始分数scores——数值上可能很大也可能为负必须随后经过 scale 缩放与 softmax 归一化才具备概率语义。核心计算流Q K^T 再乘 1/√dAT-12 模式卡片给出的标准计算流为scores matmul(Q, K, dtypeFP32, b_transTrue) scores_scaled mul(scores, scale) # scale 1/√d两个关键语义点b_transTrue转置第二个操作数Q 与 K 通常以相同布局存放[seq, head_dim]的二维视图而注意力分数需要的是Q K^T因此矩阵乘接口对第二个输入做转置输出 shape 为[q_tile, k_tile]。与之对照AT-13 的 PV 不需要转置matmul(P, V)这是因为 P 与 V 的 shape 天然匹配。FP32 累积输出dtypeFP32表示矩阵乘在 Cube 上以 FP32 累加即使输入是 BF16/FP16。这是注意力分数精度保证的起点——后续 exp 操作对数值误差非常敏感若在 BF16 下累积误差会被指数放大。仓库实现佐证Flash Attention MHA 的 C1 段在 flash_attention_mha_impl.py 中AT-12 以几乎一一对应的形式落地。先由输入 shape 推导缩放系数scale 1.0 / (head_dim ** 0.5)再在 KV tile 循环内通过pypto.view取 Q/K 分块后执行 QK^Tscores pypto.matmul(q_tile_view, k_tile_view, out_dtypepypto.DT_FP32, b_transTrue) scores_scaled pypto.mul(scores, scale)其中q_tile_view/k_tile_view是对输入张量[total_seq, N*D]经pypto.reshape(..., inplaceTrue)后按 head 偏移与 tile 偏移切出的[q_tile, head_dim]/[k_tile, head_dim]视图带valid_shape处理序列边界b_transTrue使输出成为[q_tile, k_tile]的分数矩阵。文件中的 dtype 流转注释也印证了模式卡片的描述scores Q(BF16) K^T(BF16) → FP32 (matmul out_dtypeFP32) scores_scaled scores * scale → FP32 (mul)scale 的取值在 AT-01 Online Softmax 卡片中有常见参考值0.125d64、0.0625d256即1/√head_dim的标准注意力缩放本仓库实现统一用1.0 / (head_dim ** 0.5)动态推导head_dim 改变时无需手改常量。实例化参数dtype、反量化与 scale 的决策AT-12 模式卡片的实例化参数表是设计时必须逐项确认的参数说明q_dtypeQ 的输入 dtype (BF16/FP16/FP8)dequant_afterFP8 输入时是否需要在 matmul 后反量化q_scale / k_scaleFP8 模式下的反量化 scaleq_dtype决定 Q、K通常与 Q 同 dtype送入 Cube 的精度。BF16/FP16 是标准路径FP8 则进入下述 FP8 变体。设计时需要与 API 约束 C-API-02 核对matmul的两侧输入必须满足目标 API 的 dtype 配对要求不支持的配对会在编译期失败因此 dtype 选择必须落在目标版本文档明确支持的组合内并在计算图中显式标注转换位置。dequant_after与q_scale / k_scaleFP8 输入时Cube 输出的整数累加结果需要反量化回浮点这两个参数控制是否反量化以及用什么 scale 反量化。注意k_scale在模式卡片中写作k_scale_T——因为 K 被转置其 per-token scale 的轴方向也随之转置需要与分数矩阵的轴对齐后才能做逐元素反量化。FP8 变体量化输入的 QK^T 与动态反量化当 Q、K 以 FP8 存储时例如 PageAttn FP8 场景AT-12 的计算流变为scores_int matmul(Q_fp8, K_fp8, dtypeFP32, b_transTrue) scores dequant_dynamic(scores_int, q_scale, k_scale_T) scores_scaled mul(scores, scale)与标准路径的差异Cube 输入为 FP8Q_fp8/K_fp8是经 AT-08 FP8 Quantization对称 per-token 量化FP8 E4M3scale 上限 448.0量化后的数据每个 token 伴随一个 FP32 scale。先反量化再缩放dequant_dynamic用q_scale与转置后的k_scale_T把 FP32 整数累加结果还原为浮点分数之后才执行mul(scores, scale)的 1/√d 缩放。顺序不能颠倒——对整数结果直接乘 scale 无法正确还原量化损失。dequant_after开关对应matmul 后是否需要反量化。若下游如 online softmax可以直接消费 FP8 域数值或反量化已被融合进后续算子可将该开关置为关闭由设计文档明确记录并复核精度。这种量化→Cube matmul→动态反量化的链式结构与 AT-13 的 FP8 变体P_fp8, P_scale AT-08(P_fp32)后matmul再dequant_dynamic保持一致的机制约定量化 scale 随张量同行反量化发生在矩阵乘之后。设计时建议同时阅读 AT-08 与 AT-13保持三个模式在量化路径上的 scale 语义统一。在 C1-V1-C2 骨架中的编排与合图边界AT-12 不是孤立存在的在 SK-01 Online Flash Attention 中QK^T 处于 KV tile 循环体内紧接着是 online softmax 与 PV形成C1 → V1 → C2的循环模式。编排时有几个必须遵守的规则TileShape 分阶段配置每个阶段前必须调用对应的 tile 配置 API且矩阵乘的 m/k/n 各轴使用[L0, L1]两段式配置满足0 L0 L1且L1 % L0 0Tiling 约束 C-TILE-05Vector 配置不能替代 Cube 配置。本仓库 Flash Attention 默认c1_cube_tile [[128, 128], [128, 128], [128, 128]]即 C1 阶段 m/k/n 均为[128, 128]pypto.set_cube_tile_shapes( tile_config.c1_cube_tile[0], tile_config.c1_cube_tile[1], tile_config.c1_cube_tile[2]) scores pypto.matmul(q_tile_view, k_tile_view, out_dtypepypto.DT_FP32, b_transTrue)后续 V1 阶段才切换到pypto.set_vec_tile_shapes(...)。SK-01 强调C1/V1/C2 各阶段前必须set_cube/vec_tile_shapes不切换会导致表达式上限突破等编译问题。子图合图边界sg_set_scopeCube 与 Vector 的交替会产生跨子图的 GM 落地与调度气泡。AT-12 所在的 QK^T 段通常处于默认 scope-1之外而紧随其后的 softmax vec 链mul/amax/sub/exp/sum/cast被 AT-21 Attention 分阶段子图合图 用pypto.set_pass_options(sg_set_scope...)划为独立子图# (A) QK^TCubescope 外默认 -1 sij_full pypto.matmul(qi, kj, ...) # (B) softmax vec 链sg_set_scope2mul/amax/sub/exp/sum/cast 合为一子图 pypto.set_pass_options(sg_set_scope2) sij pypto.mul(sij_full, scale) ... pypto.set_pass_options(sg_set_scope-1)AT-21 明确约束Cube 与 Vec 操作不得在同一 scope混置会报F41007 OP_SCOPE_ERROR不同阶段用不同正整数 IDsg_set_scope-1用于结束 scope 并切回默认。AT-12 的 QK^T 作为 Cube 阶段因此必须与 softmax 的 Vec 链保持 scope 隔离。与 online softmax 的衔接QK^T 输出的scores_scaled直接作为 AT-01 的输入scores: Tensor[M, K_tile]FP32。在仓库实现中V1 阶段紧随其后计算mij pypto.amax(scores_scaled, dim-1, keepdimTrue) pij pypto.exp(pypto.sub(scores_scaled, mij)) lij pypto.sum(pij, dim-1, keepdimTrue)这解释了为何 AT-12 的 scale 缩放必须在 matmul 内/紧邻完成online softmax 的逐块行最大值、指数和都建立在已缩放的分数之上scale 提前应用可避免重复乘法。工程要点与扩展方向Cube 性能配置AT-12 作为 Attention 的两个 Cube 节点之一其性能直接决定整体吞吐。SK-01 的实测配置表给出 FA 类算子的关键选项pass_options.cube_l1_reuse_settingC1(QK^T) 与 C2(PV) 的权重/激活 L1 复用FA 的核心 cube 优化{-1: 2~4}或分阶段{-1: 2, 0: 8}pass_options.cube_nbuffer_settingCube 双缓冲掩盖 K/V 加载延迟{-1: 2}起步pass_options.vec_nbuffer_settingsoftmax 阶段向量算子多{-1: 4}起步。这些配置在 flash_attention_mha_impl.py 中有多套对应实现如 910 平台的{cube_l1_reuse_setting: {0: 8, 1: 1}, ...}可作为不同架构下的参考取值。K-Split大 K 投影的 QK^T 场景延伸当 head_dim 较大或中间量超 UB 限制时可参考 AT-22 K-Split MatMul 的思路沿 K 轴用pypto.view将大投影拆分为两部分各自 matmul 到 FP32 后相加。该模式主要用于大 K 线性投影如 MLA 的 K7168 拆半但对 QK^T 的 K 轴head_dim 或 KV 序列维度同样适用且其FP32 累加 FP32 相加、不依赖enable_split_k的确定性要求与 AT-12 的 FP32 累积策略一致——K 拆分只改变归约顺序数值差异处于 FP32 噪声内适合确定性优先的场景。精度与边界处理QK^T 全程 FP32 累积是精度底线mi/li/oi等 online softmax 累积器必须 FP32仅最终 cast 输出 dtypeSK-01 强制项序列边界用viewvalid_shape处理仓库实现中valid_shape[k_tile_len, head_dim]保证非整除 tile 的尾块不越界scale 采用1.0 / (head_dim ** 0.5)运行时推导避免硬编码不同 head_dim 下的常量表。小结AT-12 QK^T MatMul Scale 是 Attention 算子设计中最基础、出现频率最高的 Cube 原子模式以matmul(Q, K, dtypeFP32, b_transTrue)完成转置矩阵乘以mul(scores, scale)完成 1/√d 缩放以dequant_dynamic支撑 FP8 量化输入的动态反量化。设计落地时将其嵌入 SK-01 的 C1 位置、遵守 TileShape 分阶段配置与 Cube/Vec scope 隔离即可与 online softmax、PV 无缝衔接写出结构清晰且可验证的注意力分数计算段。【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价