PyPTO-Gym 算子设计模式 AT-10:RMSNorm + Linear 融合(V→C 排布)的原理与 MLAProlog 实战
【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym
导读
AT-10(RMSNorm + Linear Fused)是 PyPTO 算子设计模式库中描述「先 RMSNorm 归一化、再 Linear 投影」这一高频子结构的原子模式(Atom),它正是 MLAProlog、Qwen3PreAttn 等 Prolog / Pre-Attention 算子的标准骨架。本文以该模式卡片为纲,结合 pypto-gym 仓库中 DeepSeek-V4 MLA Prolog 实现 与 对应测试用例 的源码级证据,拆解 V→C 两阶段的每一条计算指令、TileShape 切换与性能配置要点,帮助你直接复刻该模式完成自己的算子设计。
一、AT-10 在模式体系中的定位
pypto-gym 的算子设计工作流(pypto-op-design/SKILL.md)要求设计者「先读 SK 索引 和 AT 索引,再读取候选卡片」:SK(Skeleton)描述 kernel 整体结构,AT(Atom)描述局部计算。AT-10 属于 atoms/index.md 中编号第 16 条的局部计算模式:
| ID | 名称 | tags | flow_pattern |
|---|---|---|---|
| AT-09 | Linear Projection (Quantized MatMul) | matmul | C, V |
| AT-10 | RMSNorm + Linear (Fused) | norm-linear-fused | V, C |
| AT-11 | RMSNorm + Linear + Quant (Fused) | norm-linear-quant-fused | V, C, V |
其中 C 表示 Cube(矩阵乘单元)、V 表示 Vector(向量单元),flow_pattern 仅示意主要计算顺序。AT-10 直接依赖两个更底层的原子:AT-03 RMSNorm(纯 V)与 AT-09 Linear Projection(C + V),并向上支撑 SK-03 / SK-04 / SK-05 三个骨架。
二、CV 排布与标准计算流
模式卡片给出的完整定义如下:
描述:先做 RMSNorm 归一化,再做 Linear 投影。这是 Prolog 算子和 Pre-Attention 算子的标准子结构。
CV 排布:V → C
计算流:
# V 阶段: RMSNorm normed = AT-03(x, gamma, eps) normed_bf16 = cast(normed, BF16) # C 阶段: Linear projected = matmul(normed_bf16, weight, dtype=BF16, b_trans=True)关键语义有三点:
- 先 V 后 C:RMSNorm 是逐 token 的向量归约,天然落在 Vector 单元;投影是矩阵乘,落在 Cube 单元。V→C 的顺序意味着 kernel 内部需要一次 Vector→Cube 的排布切换(PyPTO 中用
set_vec_tile_shapes/set_cube_tile_shapes表达)。 - 中间精度收敛到 BF16:RMSNorm 内部在 FP32 下计算(保证精度),但喂给 matmul 之前必须
cast(normed, BF16),让矩阵乘两侧都以 BF16 输入(dtype=BF16)。 - 权重按
b_trans=True传入:matmul 的 B 侧权重以[N, K]布局传入并做转置,这是 PyPTO matmul 对「激活 × 权重」的标准调用形态。
使用算子:MLAProlog(q_a_proj → norm → q_b_proj)、Qwen3PreAttn(input_norm → QKV_proj)。前者在 MLA 架构中把低秩 Query 投影拆成「压缩投影 + 归一化 + 展开投影」两段,后者在 Pre-Attention 阶段先对输入归一化再做 QKV 联合投影。
三、V 阶段深入:RMSNorm 的两种变体与源码级拆解
AT-10 的 V 阶段完整继承 AT-03 RMSNorm 的计算语义。AT-03 的输入为x: Tensor[*, D](BF16/FP16)、gamma: Tensor[D](BF16,可选)、eps: float;输出为y: Tensor[*, D],可选的rstd: Tensor[*, 1]。
变体 A — rsqrt(推荐,硬件融合指令)
x_fp32 = cast(x, FP32) x_sq = mul(x_fp32, x_fp32) mean_sq = mul(sum(x_sq, dim=-1, keepdim=True), 1.0/D) var = add(mean_sq, eps) rstd = rsqrt(var) y = mul(x_fp32, rstd) [可选] y = mul(y, cast(gamma, FP32)) y_out = cast(y, BF16)变体 B — sqrt + div
...同上到 mean_sq... var = add(mean_sq, eps) std = sqrt(var) rstd = div(ones, std) ...pypto-gym 中 DeepSeek-V4 MLA Prolog 实现的rms_norm函数 采用的就是变体 B的逐指令写法,可直接对照:
def rms_norm(input_tensor: pypto.Tensor, epsilon: float) -> pypto.Tensor: input_fp32 = pypto.cast(input_tensor, pypto.DT_FP32) dim = len(input_tensor.shape) y = pypto.mul(input_fp32, input_fp32) # x^2 y = pypto.mul(y, 1.0 / input_tensor.shape[dim - 1]) # * 1/D y = pypto.sum(y, -1, keepdim=True) # mean_sq y = pypto.add(y, epsilon) # + eps y = pypto.sqrt(y) # sqrt ones_vector = pypto.full(y.shape, 1.0, pypto.DT_FP32) y = pypto.div(ones_vector, y) # 1/std y = pypto.mul(input_fp32, y) # x * rstd return y可见实现顺序与变体 B 完全一致:FP32 计算全程、sum沿最后一维 keepdim、sqrt + div求倒数。gamma缩放由调用方在函数外完成(见下文 MLAProlog 的pypto.mul(qr, gamma_cq_2d_fp32)),这对应 AT-03 实例化参数表中has_gamma的有/无两种形态。
AT-03 实例化参数
| 参数 | 说明 | 变体 |
|---|---|---|
has_gamma | 是否乘 gamma | 有 gamma (MLAProlog) / 无 gamma (mhc_pre) |
has_bias | 是否加 bias | GLMAttnFusion 使用 |
rsqrt_mode | rsqrt vs sqrt+div | rsqrt (Qwen3) / sqrt+div (GLM) |
output_rstd | 是否输出 rstd | InplaceAddRmsNorm 输出 |
对应到 PyPTO 约束,API 约束文档 中 C-API-05 明确「精度敏感的归约和跨循环累加优先使用 FP32」,这正是 RMSNorm 内部全程 FP32 的原因;C-API-02 则要求 matmul 两侧输入满足 dtype 配对要求,这解释了为什么 V→C 之间必须有cast(normed, BF16)。
四、C 阶段深入:Linear 投影的三种模式
AT-10 的 C 阶段对应 AT-09 Linear Projection,输入x: Tensor[M, K]、权重w: Tensor[K, N]或[N, K],可选 bias 与量化参数。标准 BF16 模式即 AT-10 计算流中的那行 matmul:
y = matmul(x, w, dtype=BF16, b_trans=True)AT-09 另外定义了两种量化模式(AT-10 的可选扩展方向):
INT8 W8A8 模式:
x_int8, x_scale = AT-05(x) # 量化激活 y_int32 = matmul(x_int8, w_int8, dtype=INT32) # 整数矩阵乘 y = AT-06(y_int32, x_scale, w_scale) # 反量化MXFP8 模式:
y = scaled_mm(x_fp8, w_fp8, FP32, x_scale, w_scale)其中 AT-05 是逐 token 对称量化(amax求 max →127.0/max得 scale → 三次 cast 完成舍入与饱和),AT-06 负责反量化。AT-09 的实例化参数为quant_mode(none / int8_w8a8 / mxfp8)、has_bias、out_dtype(BF16/FP32)。当选择量化模式时,AT-10 升级为 AT-11 RMSNorm + Linear + Quant,排布变为 V → C → V。
五、实战案例一:MLAProlog(DeepSeek-V4 源码逐段对照)
AT-10 在 MLA Prolog 中表现为q_a_proj → norm → q_b_proj的两段式低秩投影。仓库中的 mla_prolog_v4_impl.py 提供了精确的工程实现:
unroll_list = configs.unroll_list for tIdx, unrollLength in pypto.loop_unroll(0, t, 1, name="MLA_BS_LOOP", idx_name="bs_offset", unroll_list=unroll_list): t_tile = unrollLength x_tile = pypto.view(x, [t_tile, h], [tIdx, 0], valid_shape=[t_tile, h]) # ===== AT-10 第一次出现: wq_a 投影 + RMSNorm ===== pypto.set_semantic_label("wqa-linear") pypto.set_cube_tile_shapes([32, 32], [512, 512], [64, 64]) q = pypto.matmul(x_tile, wq_a, pypto.DataType.DT_BF16) # C: Linear (q_a_proj) pypto.set_semantic_label("q-rmsnorm with weight") pypto.set_vec_tile_shapes(8, q_lora_rank) qr = rms_norm(q, attrs.eps) # V: RMSNorm qr = pypto.mul(qr, gamma_cq_2d_fp32) # V: * gamma (has_gamma) qr = pypto.cast(qr, pypto.DataType.DT_BF16) # V: cast 回 BF16 pypto.assemble(qr, [tIdx, 0], qr_out) # ===== AT-10 第二次出现: q_b 展开投影 ===== pypto.set_semantic_label("wqb-linear") pypto.set_cube_tile_shapes([32, 32], [128, 128], [256, 256]) q = pypto.matmul(qr, wq_b, pypto.DataType.DT_BF16) # C: Linear (q_b_proj) ...对照 AT-10 计算流可以逐行印证:
- V→C 切换:
wqa-linear段先set_cube_tile_shapes做 matmul,紧接着q-rmsnorm段set_vec_tile_shapes(8, q_lora_rank)切回 Vector 做归一化,这是 PyPTO 中表达 V→C 排布的标准手法。 - gamma 的 FP32 化:
gamma_cq_2d_fp32在循环外预先reshape + cast好(L350-L354),循环内只做一次mul,避免重复转换。 - 中间 BF16 cast:
rms_norm返回 FP32 结果,pypto.cast(qr, DT_BF16)后作为下一个 matmul 的 A 侧输入,完全符合 AT-10「normed_bf16 = cast(normed, BF16)」的约定。 - 权重 B 侧布局:
wq_b以静态[STATIC, STATIC]BF16 张量传入(L413),对应b_trans=True的[N, K]布局语义。 - jit 入口配置:kernel 外层
@pypto.frontend.jit(runtime_options={"stitch_function_max_num": 128})(L406-L408),用于多阶段融合调度。
此外 KV 路径(kv = matmul(x_tile, wkv, BF16)→rms_norm→mul(gamma_ckv)→cast BF16,L390-L396)是同一 AT-10 模式在 KV 分支上的平行实例。DeepSeek-V2-Lite 的混合实现 mla_prolog.py 则展示了另一种组织:用loop_unroll循环包裹kv_b_projmatmul,并把权重预转置后直接 matmul(省去每任务重复的b_trans跨步寻址)。
六、实战案例二:Qwen3PreAttn(Pre-Attention 的 input_norm → QKV_proj)
AT-10 在 Pre-Attention 算子中的形态是input_norm → QKV_proj:对输入(通常先做残差相加)执行 RMSNorm,再把归一化结果一次性投影为 Q/K/V 拼接张量。仓库中的骨架文档 SK-05 Fused Pre-Attn (Two-Phase) 以 Qwen3PreAttnFused 为典型算子,给出了两阶段排布:
- Phase 1(Pre-Processing,V-C-V):
V(Norm) → C(Quant Linear) → V(Dequant+Split+Norm+RoPE+Cache)——其中 Norm 与 Linear 正是 AT-10(或量化版 AT-11),示例骨架片段为:
normed = rms_norm(x_add, gamma, eps) # V: Norm(AT-10 的 V 阶段) # C: Quant Linear (INT8 W8A8) x_int8, x_scale = quantize(normed) y_int32 = matmul(x_int8, w_int8, INT32) y = dequant(y_int32, x_scale, w_scale) # C 阶段(AT-11 形态) # V: QKV Split + RoPE q = rms_norm_per_head(y[:, :q_dim], q_gamma) ...- Phase 2(Flash Attention):复用 SK-01 的 C1-V1-C2 在线 softmax 结构。
在 Qwen3 这类非量化实现中,Phase 1 直接退化为标准 AT-10:rms_norm(x_add, gamma, eps)→cast(BF16)→matmul(normed_bf16, w_qkv, BF16, b_trans=True),输出经 split 切成 Q、K、V 再分别做 per-head Norm 与 RoPE。同时 SK-03 Linear Projection (Norm→MatMul) 也把 Qwen3PreAttn 列为典型算子,并给出了单阶段 V→C 的参考结构(rms_norm→cast BF16→set_cube_tile_shapes→matmul(..., b_trans=True)→assemble)。
七、AT-10 在整体骨架中的组织方式
AT-10 作为原子模式,可被不同骨架以不同粒度复用:
| 骨架 | 组织方式 | 对应算子 |
|---|---|---|
| SK-03 Linear Projection | 单层 Loop 内一次 V→C | MLAProlog(部分)、Qwen3PreAttn |
| SK-04 Multi-Stage Fused Prolog | C→V→C→V 多阶段串联,AT-10 反复出现 | MLAProlog、MLAPrologQuant |
| SK-05 Fused Pre-Attn | Phase1 V-C-V + Phase2 Flash Attention | GLMAttnFusion、Qwen3PreAttnFused |
SK-04 特别强调loop_unroll必须使用双变量解包(for bs_offset, tile_bs in pypto.loop_unroll(...)):每个 unroll 因子生成独立子循环路径,路径内tile_bs特化为编译期整数,从而满足view的静态List[int]约束;t 整除法则为t=8 → tile_bs=8 单 root、t=4 → 4、t=16 → 8×2。这一规则在 mla_prolog_v4_impl.py 中即为for tIdx, unrollLength in pypto.loop_unroll(...)的实践。
八、AT-10 性能调优要点(开箱配置清单)
综合 SK-03 / SK-04 / SK-05 三个骨架文档,AT-10 形态 kernel 的推荐配置如下:
| 维度 | 推荐配置 | 取值经验 | 作用 |
|---|---|---|---|
runtime_options.stitch_function_max_num | 必配 | 128 | 多阶段融合;部分平台会回退,需按平台验证 |
pass_options.cube_l1_reuse_setting | 必配 | SK-03:{-1: 2, 1: 1}分轴;SK-04:{-1: 8, 0: 1, 1: 1} | 权重轴用 1(不复用),激活轴双缓冲,匹配「权重静态、激活动态」 |
pass_options.vec_nbuffer_setting | 推荐 | {0: 2}(Norm 阶段) | RMSNorm 在 V 阶段,nbuffer=2 即可 |
pypto.set_cache_policy(NONE_CACHEABLE, True) | 条件配 | 仅当权重在 loop 内被单次消费 | 权重只读一次不占 L2,避免与激活竞争 cache;⚠️ 若权重跨迭代复用且可驻留 L2(如 mhc_pre 的 phi 被 8 次 unroll 复用),标记 NONE_CACHEABLE 反而每迭代回 HBM 重读——勿用 |
pypto.set_semantic_label(...) | 推荐/必配 | 每阶段一个标签(如 "wqa-linear"、"q-rmsnorm") | 编译器靠语义标签做阶段隔离调度,遗漏会导致跨阶段错误融合 |
combine_axis=True | 必配 | jit 首行 | 尾轴 broadcast 内联 brcb |
pypto.reshape(..., inplace=True) | 推荐 | tile 内 reshape | 避免临时张量分配 |
| token 循环展开 | 按需、单值 | SK-03 候选 128/64/32/16/8/1;SK-04 候选 8/4/2/1 | 初始设计只选一个值并验证余数处理,其余留作调优候选 |
一个特殊形态值得注意:当 AT-10 退化为「单 matmul + norm」且形状落在输出宽 ≤ 64、归约维 N·D ≥ 2^14、FP32 计算(典型:mhc_pre 的matmul [B,28672]×[28672,24])时,SK-03 提供了专用变体——loop_unroll(0, BS, 1, unroll_list=[16])让 Vector 连续处理整 D 归约、Cube 按 M=16 出多个任务;vec tile 用 D 轴大 tile;权重 host 侧预转置省去b_trans;cube 用[16,16],[512,1024],[128,128]+enable_split_k=True。该变体的有效性依赖 loop_unroll 结构前提,不可拆分套用到平铺 BT-loop 上。
九、验证方式:golden 对照与测试用例
AT-10 的正确性验证在 pypto-gym 中采用「golden 参考 + 数值对比」方式。test_mla_prolog_v4.py 提供了与 kernel 等价的 torch 参考实现:
- rms_norm_new(无 gamma):
x_f32 * x_f32→* 1/D→sum + eps→sqrt→x_f32 / reduce_sqrt,与 kernel 变体 B 完全同构; - rms_norm(带 gamma):额外执行
res_div * gamma,对应has_gamma=True形态; - golden 计算链(L125-L160):
q_a_proj = torch.matmul(x, wq_a)→rms_norm(q_a_proj, gamma_cq)→q_b_proj = torch.matmul(q_a_layernorm, wq_b)→reshape(num_heads, head_dim)→ 再次 RMSNorm,逐段复现 AT-10 两次出现的位置; - 对比容差:
compare(output, golden, name, 0.0001, 0.0078125, 0.005)(L337-L341); - 测试输入规模:如
test_t16_pa_nd_bf16使用t=16, num_heads=64, h=4096, q_lora_rank=1024, head_dim=512, qk_rope_head_dim=64(L410-L424),unroll_list=[128, 64, 32, 16, 1]、cube_l1_reuse_setting={2: 4}(L440-L447),并标注为 large test case 默认 skip。
十、设计工作流中的应用建议
按 pypto-op-design/SKILL.md 的流程,当新算子(如某个模型的 prolog / pre-attention 段)出现「归一化 + 投影」组合时:
- 先在 AT 索引 中按 tags(
norm、matmul、norm-linear-fused)命中 AT-10; - 若后续还有量化,升级为 AT-11;若只是单次 V→C,直接采用 SK-03 结构;若存在 q_a→norm→q_b 多段串联,采用 SK-04;
- 用 AT-03 的
rsqrt_mode/has_gamma参数和 AT-09 的quant_mode/out_dtype参数完成局部实例化,参考上文性能配置表落地 runtime/pass options; - 最后按第八节的 golden 对照模式编写测试,验证数值容差与尾块处理。
结语
AT-10 看似只有两行计算流,但它锚定了 PyPTO 算力编排中最关键的一次 V→C 排布切换,并串联起 RMSNorm 的精度策略(FP32 内部计算 + BF16 输出)、matmul 的权重布局约定(b_trans=True)与量化扩展路径。以 mla_prolog_v4_impl.py 为参照实现、test_mla_prolog_v4.py 为验证基线,你可以在自己的 Prolog / Pre-Attention 算子设计中快速落地这一模式,并沿着 SK-03/04/05 的性能方向继续调优。
【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考