AT-12 QK^T MatMul + Scale:PyPTO 注意力分数计算的 Cube 侧核心模式与 FP8 变体实战
2026/9/19 23:21:00 网站建设 项目流程

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

【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym

Attention(自注意力)前向计算中,第一步是用 Query 与 Key 做矩阵乘得到原始注意力分数。在 PyPTO 算子设计中,这一步被抽象为局部计算模式(Atom Pattern)AT-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 的算子设计体系中,计算模式被划分为**骨架(Skeleton,SK)局部原子模式(Atom,AT)**两层:SK 描述整个算子的循环组织与阶段编排,AT 描述其中一段具体的局部计算。AT-12 属于 AT 索引 中的 matmul 类模式,flow_patternC(Cube),即它完全落在 Cube 计算单元上,不涉及 Vector 操作。

AT-12 的典型使用场景在模式卡片的examples字段中写得很明确:所有 Attention 算子。它在 Flash Attention / IFA / PageAttn / SparseAttn 等注意力算子中扮演的是 SK-01(Online Flash Attention 骨架)中C1 阶段的角色,处于如下计算链的起点:

C1(QK^T) → V1(Softmax) → C2(P@V)
  • 前置阶段:Q/K 经过 AT-03 RMSNorm 或 RoPE 后进入 QK^T;
  • 后继阶段:AT-12 的输出scores直接喂给 AT-01 Online Softmax(V1 阶段)做逐块最大值/指数和累积;
  • 下游镜像:softmax 归一化后的概率 P 再与 V 做矩阵乘,即姊妹模式 AT-13 P@V MatMul(C2 阶段)。

理解 AT-12 的关键在于:它输出的不是最终的注意力权重,而是原始分数(scores)——数值上可能很大也可能为负,必须随后经过 scale 缩放与 softmax 归一化才具备概率语义。

核心计算流:Q @ K^T 再乘 1/√d

AT-12 模式卡片给出的标准计算流为:

scores = matmul(Q, K, dtype=FP32, b_trans=True) scores_scaled = mul(scores, scale) # scale = 1/√d

两个关键语义点:

  1. b_trans=True转置第二个操作数:Q 与 K 通常以相同布局存放([seq, head_dim]的二维视图),而注意力分数需要的是Q @ K^T,因此矩阵乘接口对第二个输入做转置,输出 shape 为[q_tile, k_tile]。与之对照,AT-13 的 P@V 不需要转置(matmul(P, V)),这是因为 P 与 V 的 shape 天然匹配。
  2. FP32 累积输出dtype=FP32表示矩阵乘在 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^T:

scores = pypto.matmul(q_tile_view, k_tile_view, out_dtype=pypto.DT_FP32, b_trans=True) scores_scaled = pypto.mul(scores, scale)

其中q_tile_view/k_tile_view是对输入张量[total_seq, N*D]pypto.reshape(..., inplace=True)后按 head 偏移与 tile 偏移切出的[q_tile, head_dim]/[k_tile, head_dim]视图(带valid_shape处理序列边界),b_trans=True使输出成为[q_tile, k_tile]的分数矩阵。文件中的 dtype 流转注释也印证了模式卡片的描述:

scores = Q(BF16) @ K^T(BF16) → FP32 (matmul out_dtype=FP32) scores_scaled = scores * scale → FP32 (mul)

scale 的取值在 AT-01 Online Softmax 卡片中有常见参考值:0.125(d=64)、0.0625(d=256),即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 模式下的反量化 scale
  • q_dtype:决定 Q、K(通常与 Q 同 dtype)送入 Cube 的精度。BF16/FP16 是标准路径,FP8 则进入下述 FP8 变体。设计时需要与 API 约束 C-API-02 核对:matmul的两侧输入必须满足目标 API 的 dtype 配对要求,不支持的配对会在编译期失败,因此 dtype 选择必须落在目标版本文档明确支持的组合内,并在计算图中显式标注转换位置。
  • dequant_afterq_scale / k_scale:FP8 输入时,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, dtype=FP32, b_trans=True) scores = dequant_dynamic(scores_int, q_scale, k_scale_T) scores_scaled = mul(scores, scale)

与标准路径的差异:

  1. Cube 输入为 FP8Q_fp8/K_fp8是经 AT-08 FP8 Quantization(对称 per-token 量化,FP8 E4M3,scale 上限 448.0)量化后的数据,每个 token 伴随一个 FP32 scale。
  2. 先反量化再缩放dequant_dynamicq_scale与转置后的k_scale_T把 FP32 整数累加结果还原为浮点分数,之后才执行mul(scores, scale)的 1/√d 缩放。顺序不能颠倒——对整数结果直接乘 scale 无法正确还原量化损失。
  3. dequant_after开关:对应"matmul 后是否需要反量化"。若下游(如 online softmax)可以直接消费 FP8 域数值或反量化已被融合进后续算子,可将该开关置为关闭,由设计文档明确记录并复核精度。

这种"量化→Cube matmul→动态反量化"的链式结构与 AT-13 的 FP8 变体(P_fp8, P_scale = AT-08(P_fp32)matmuldequant_dynamic)保持一致的机制约定:量化 scale 随张量同行,反量化发生在矩阵乘之后。设计时建议同时阅读 AT-08 与 AT-13,保持三个模式在量化路径上的 scale 语义统一。

在 C1-V1-C2 骨架中的编排与合图边界

AT-12 不是孤立存在的:在 SK-01 Online Flash Attention 中,QK^T 处于 KV tile 循环体内,紧接着是 online softmax 与 P@V,形成C1 → V1 → C2的循环模式。编排时有几个必须遵守的规则:

TileShape 分阶段配置

每个阶段前必须调用对应的 tile 配置 API,且矩阵乘的 m/k/n 各轴使用[L0, L1]两段式配置,满足0 < L0 <= L1L1 % L0 == 0(Tiling 约束 C-TILE-05);Vector 配置不能替代 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_dtype=pypto.DT_FP32, b_trans=True)

后续 V1 阶段才切换到pypto.set_vec_tile_shapes(...)。SK-01 强调:C1/V1/C2 各阶段前必须set_cube/vec_tile_shapes,不切换会导致表达式上限突破等编译问题。

子图合图边界(sg_set_scope)

Cube 与 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^T:Cube,scope 外(默认 -1) sij_full = pypto.matmul(qi, kj, ...) # (B) softmax vec 链:sg_set_scope=2(mul/amax/sub/exp/sum/cast 合为一子图) pypto.set_pass_options(sg_set_scope=2) sij = pypto.mul(sij_full, scale) ... pypto.set_pass_options(sg_set_scope=-1)

AT-21 明确约束:Cube 与 Vec 操作不得在同一 scope(混置会报F41007 OP_SCOPE_ERROR),不同阶段用不同正整数 ID,sg_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, keepdim=True) pij = pypto.exp(pypto.sub(scores_scaled, mij)) lij = pypto.sum(pij, dim=-1, keepdim=True)

这解释了为何 AT-12 的 scale 缩放必须在 matmul 内/紧邻完成:online softmax 的逐块行最大值、指数和都建立在已缩放的分数之上,scale 提前应用可避免重复乘法。

工程要点与扩展方向

Cube 性能配置

AT-12 作为 Attention 的两个 Cube 节点之一,其性能直接决定整体吞吐。SK-01 的实测配置表给出 FA 类算子的关键选项:

  • pass_options.cube_l1_reuse_setting:C1(QK^T) 与 C2(P@V) 的权重/激活 L1 复用,FA 的核心 cube 优化,{-1: 2~4}或分阶段{-1: 2, 0: 8}
  • pass_options.cube_nbuffer_setting:Cube 双缓冲,掩盖 K/V 加载延迟,{-1: 2}起步;
  • pass_options.vec_nbuffer_setting:softmax 阶段向量算子多,{-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 的 K=7168 拆半),但对 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 输出 dtype(SK-01 强制项);
  • 序列边界用view+valid_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, dtype=FP32, b_trans=True)完成转置矩阵乘,以mul(scores, scale)完成 1/√d 缩放,以dequant_dynamic支撑 FP8 量化输入的动态反量化。设计落地时,将其嵌入 SK-01 的 C1 位置、遵守 TileShape 分阶段配置与 Cube/Vec scope 隔离,即可与 online softmax、P@V 无缝衔接,写出结构清晰且可验证的注意力分数计算段。

【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询