news 2026/9/19 15:47:37

PyPTO-Gym 算子设计模式 AT-10:RMSNorm + Linear 融合(V→C 排布)的原理与 MLAProlog 实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PyPTO-Gym 算子设计模式 AT-10:RMSNorm + Linear 融合(V→C 排布)的原理与 MLAProlog 实战

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名称tagsflow_pattern
AT-09Linear Projection (Quantized MatMul)matmulC, V
AT-10RMSNorm + Linear (Fused)norm-linear-fusedV, C
AT-11RMSNorm + Linear + Quant (Fused)norm-linear-quant-fusedV, 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)

关键语义有三点:

  1. 先 V 后 C:RMSNorm 是逐 token 的向量归约,天然落在 Vector 单元;投影是矩阵乘,落在 Cube 单元。V→C 的顺序意味着 kernel 内部需要一次 Vector→Cube 的排布切换(PyPTO 中用set_vec_tile_shapes/set_cube_tile_shapes表达)。
  2. 中间精度收敛到 BF16:RMSNorm 内部在 FP32 下计算(保证精度),但喂给 matmul 之前必须cast(normed, BF16),让矩阵乘两侧都以 BF16 输入(dtype=BF16)。
  3. 权重按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是否加 biasGLMAttnFusion 使用
rsqrt_modersqrt vs sqrt+divrsqrt (Qwen3) / sqrt+div (GLM)
output_rstd是否输出 rstdInplaceAddRmsNorm 输出

对应到 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_biasout_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-rmsnormset_vec_tile_shapes(8, q_lora_rank)切回 Vector 做归一化,这是 PyPTO 中表达 V→C 排布的标准手法。
  • gamma 的 FP32 化gamma_cq_2d_fp32在循环外预先reshape + cast好(L350-L354),循环内只做一次mul,避免重复转换。
  • 中间 BF16 castrms_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_normmul(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_normcast BF16set_cube_tile_shapesmatmul(..., b_trans=True)assemble)。

七、AT-10 在整体骨架中的组织方式

AT-10 作为原子模式,可被不同骨架以不同粒度复用:

骨架组织方式对应算子
SK-03 Linear Projection单层 Loop 内一次 V→CMLAProlog(部分)、Qwen3PreAttn
SK-04 Multi-Stage Fused PrologC→V→C→V 多阶段串联,AT-10 反复出现MLAProlog、MLAPrologQuant
SK-05 Fused Pre-AttnPhase1 V-C-V + Phase2 Flash AttentionGLMAttnFusion、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 单 roott=4 → 4t=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/Dsum + epssqrtx_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 段)出现「归一化 + 投影」组合时:

  1. 先在 AT 索引 中按 tags(normmatmulnorm-linear-fused)命中 AT-10;
  2. 若后续还有量化,升级为 AT-11;若只是单次 V→C,直接采用 SK-03 结构;若存在 q_a→norm→q_b 多段串联,采用 SK-04;
  3. 用 AT-03 的rsqrt_mode/has_gamma参数和 AT-09 的quant_mode/out_dtype参数完成局部实例化,参考上文性能配置表落地 runtime/pass options;
  4. 最后按第八节的 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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/19 15:47:29

集团财务数字化规划:架构分层、数据主线与落地验证

简介:这份集团公司财务管理数字化规划方案(88页PPT)面向企业财务管理者、数字化转型规划人员及咨询顾问,系统展示了集团财务数字化转型的整体蓝图。内容涵盖业务流程体系设计、以用户体验为中心的全面需求调研、业务能力提升机会识…

作者头像 李华
网站建设 2026/9/19 15:45:18

C1科目一2025题库答案:用Python解析docx生成错题本与模拟卷

简介:2025年C1驾照科目一必考题库附含答案,面向正在备考C1驾驶证理论考试的学员,聚焦科目一高频考点与紧急驾驶应对策略。压缩包为docx格式,内含1个Word文档,整体大小仅49KB,便于下载后随时在电脑或手机上查…

作者头像 李华
网站建设 2026/9/19 15:44:57

EtherCAT星型拓扑断线故障解析:HotConnect配置与验证指南

简介:一份关于倍福EtherCAT HotConnect设置方法的PDF技术资料,面向工业自动化现场工程师与倍福控制器开发人员,解决设备热插拔或物理线路变动导致EtherCAT网络通讯中断、IO值停止刷新的常见问题。资源为单个PDF文件,压缩包大小119…

作者头像 李华
网站建设 2026/9/19 15:44:27

Emotion-LLaMA:面向情感识别的多模态LLaMA端到端构建实战

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/19 15:43:53

74LS芯片触发器实验:从RS到D触发器与乒乓球电路解析

简介:这份《数电》4.触发器及其应用文档,是一份面向数字电子技术课程的实验指导,适合高校电子类学生、实验课教师及自学者对照练习,重点解决触发器逻辑功能理解与基本时序电路设计问题。文档依次讲解基本RS、JK、D、T四种触发器的…

作者头像 李华
网站建设 2026/9/19 15:42:31

小天才电话手表root全攻略:解锁系统权限与刷机避坑指南

1. 项目缘起与整体思路拆解小天才电话手表在儿童智能穿戴市场里占有率很高,不少家长和数码爱好者手里都有闲置的旧款设备。这类手表本质上是一台跑着定制安卓系统的微型终端,出厂时厂商出于稳定性和安全考虑,把权限收得很紧,普通用…

作者头像 李华