我要提问
ARTICLE DETAIL

资讯详情

前沿编程新知与开发实战干货的深度解读。

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

PyPTO-Gym 算子设计模式 AT-10:RMSNorm + Linear 融合(V→C 排布)的原理与 MLAProlog 实战 PyPTO-Gym 算子设计模式 AT-10RMSNorm Linear 融合V→C 排布的原理与 MLAProlog 实战【免费下载链接】pypto-gymPyPTO-Gym 是基于 PyPTO 编程框架构建的算子与模型样例仓库项目地址: https://gitcode.com/cann/pypto-gym导读AT-10RMSNorm 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 索引再读取候选卡片」SKSkeleton描述 kernel 整体结构ATAtom描述局部计算。AT-10 属于 atoms/index.md 中编号第 16 条的局部计算模式ID名称tagsflow_patternAT-09Linear Projection (Quantized MatMul)matmulC, VAT-10RMSNorm Linear (Fused)norm-linear-fusedV, CAT-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 ProjectionC 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, dtypeBF16, b_transTrue)关键语义有三点先 V 后 CRMSNorm 是逐 token 的向量归约天然落在 Vector 单元投影是矩阵乘落在 Cube 单元。V→C 的顺序意味着 kernel 内部需要一次 Vector→Cube 的排布切换PyPTO 中用set_vec_tile_shapes/set_cube_tile_shapes表达。中间精度收敛到 BF16RMSNorm 内部在 FP32 下计算保证精度但喂给 matmul 之前必须cast(normed, BF16)让矩阵乘两侧都以 BF16 输入dtypeBF16。权重按b_transTrue传入matmul 的 B 侧权重以[N, K]布局传入并做转置这是 PyPTO matmul 对「激活 × 权重」的标准调用形态。使用算子MLAPrologq_a_proj → norm → q_b_proj、Qwen3PreAttninput_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, keepdimTrue), 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, keepdimTrue) # 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 sqrtdivrsqrt (Qwen3) / sqrtdiv (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 计算流中的那行 matmuly matmul(x, w, dtypeBF16, b_transTrue)AT-09 另外定义了两种量化模式AT-10 的可选扩展方向INT8 W8A8 模式x_int8, x_scale AT-05(x) # 量化激活 y_int32 matmul(x_int8, w_int8, dtypeINT32) # 整数矩阵乘 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_modenone / int8_w8a8 / mxfp8、has_bias、out_dtypeBF16/FP32。当选择量化模式时AT-10 升级为 AT-11 RMSNorm Linear Quant排布变为 V → C → V。五、实战案例一MLAPrologDeepSeek-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, nameMLA_BS_LOOP, idx_namebs_offset, unroll_listunroll_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 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_transTrue的[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 BF16L390-L396是同一 AT-10 模式在 KV 分支上的平行实例。DeepSeek-V2-Lite 的混合实现 mla_prolog.py 则展示了另一种组织用loop_unroll循环包裹kv_b_projmatmul并把权重预转置后直接 matmul省去每任务重复的b_trans跨步寻址。六、实战案例二Qwen3PreAttnPre-Attention 的 input_norm → QKV_projAT-10 在 Pre-Attention 算子中的形态是input_norm → QKV_proj对输入通常先做残差相加执行 RMSNorm再把归一化结果一次性投影为 Q/K/V 拼接张量。仓库中的骨架文档 SK-05 Fused Pre-Attn (Two-Phase) 以 Qwen3PreAttnFused 为典型算子给出了两阶段排布Phase 1Pre-ProcessingV-C-VV(Norm) → C(Quant Linear) → V(DequantSplitNormRoPECache)——其中 Norm 与 Linear 正是 AT-10或量化版 AT-11示例骨架片段为normed rms_norm(x_add, gamma, eps) # V: NormAT-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 2Flash Attention复用 SK-01 的 C1-V1-C2 在线 softmax 结构。在 Qwen3 这类非量化实现中Phase 1 直接退化为标准 AT-10rms_norm(x_add, gamma, eps)→cast(BF16)→matmul(normed_bf16, w_qkv, BF16, b_transTrue)输出经 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_transTrue)→assemble。七、AT-10 在整体骨架中的组织方式AT-10 作为原子模式可被不同骨架以不同粒度复用骨架组织方式对应算子SK-03 Linear Projection单层 Loop 内一次 V→CMLAProlog部分、Qwen3PreAttnSK-04 Multi-Stage Fused PrologC→V→C→V 多阶段串联AT-10 反复出现MLAProlog、MLAPrologQuantSK-05 Fused Pre-AttnPhase1 V-C-V Phase2 Flash AttentionGLMAttnFusion、Qwen3PreAttnFusedSK-04 特别强调loop_unroll必须使用双变量解包for bs_offset, tile_bs in pypto.loop_unroll(...)每个 unroll 因子生成独立子循环路径路径内tile_bs特化为编译期整数从而满足view的静态List[int]约束t 整除法则为t8 → tile_bs8 单 root、t4 → 4、t16 → 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 阶段nbuffer2 即可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_axisTrue必配jit 首行尾轴 broadcast 内联 brcbpypto.reshape(..., inplaceTrue)推荐tile 内 reshape避免临时张量分配token 循环展开按需、单值SK-03 候选 128/64/32/16/8/1SK-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 按 M16 出多个任务vec tile 用 D 轴大 tile权重 host 侧预转置省去b_transcube 用[16,16],[512,1024],[128,128]enable_split_kTrue。该变体的有效性依赖 loop_unroll 结构前提不可拆分套用到平铺 BT-loop 上。九、验证方式golden 对照与测试用例AT-10 的正确性验证在 pypto-gym 中采用「golden 参考 数值对比」方式。test_mla_prolog_v4.py 提供了与 kernel 等价的 torch 参考实现rms_norm_new无 gammax_f32 * x_f32→* 1/D→sum eps→sqrt→x_f32 / reduce_sqrt与 kernel 变体 B 完全同构rms_norm带 gamma额外执行res_div * gamma对应has_gammaTrue形态golden 计算链L125-L160q_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使用t16, num_heads64, h4096, q_lora_rank1024, head_dim512, qk_rope_head_dim64L410-L424unroll_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 索引 中按 tagsnorm、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_transTrue与量化扩展路径。以 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),仅供参考
返回列表