ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

CANN PyPTO 算子指南:pypto.rms_norm 根均方层归一化(RMSNorm)的用法、原理与性能实践

CANN PyPTO 算子指南:pypto.rms_norm 根均方层归一化(RMSNorm)的用法、原理与性能实践 CANN PyPTO 算子指南pypto.rms_norm 根均方层归一化RMSNorm的用法、原理与性能实践【免费下载链接】pyptoPyPTO发音: pai p-t-oParallel Tensor/Tile Operation编程范式。项目地址: https://gitcode.com/cann/pypto导读本文聚焦 CANN PyPTO 提供的归一化算子pypto.rms_norm讲解其函数原型、参数语义、调用示例与数学原理并结合仓库源码剖析其 FP32 高精度提升路径、轻量化 NPU 分支以及 gamma 广播机制最后通过单元测试与编程指南中的性能建议说明如何在 LLM 场景含 QKV 融合中正确、高效地使用该算子。功能说明与数学原理pypto.rms_norm沿输入 Tensor 的最后一个维度执行根均方层归一化Root Mean Square LayerNormRMSNorm。与经典 LayerNorm 不同RMSNorm 不计算均值、不做均值中心化而是仅基于均方根RMS进行缩放因此省去了均值归约与偏置项计算量更小、更适合长序列场景。数学定义如下设最后一个维度大小为 C即归一化轴长度为n Crms sqrt( sum( x_i^2 * (1/n), i in [0, C) ) epsilon ) y_i x_i / rms 若提供 gamma y_i y_i * gamma_i即先对最后一个维度的元素求平方和并乘以1/n得到均方值加epsilon后开平方得到 RMS再用x / rms完成归一化如果提供了gamma则在最后一个维度上对归一化结果做逐元素缩放。从源码看该公式被精确实现于 python/pypto/operator.pyn x.shape[-1] y pypto.sqrt(pypto.sum(x * x * (1.0 / n), -1, keepdimTrue) epsilon)其中epsilon是数值稳定性常数用于避免分母为 0默认值为1e-6。产品支持情况pypto.rms_norm当前支持的硬件产品如下Ascend 950PR / Ascend 950DT支持Atlas A3 训练系列产品 / Atlas A3 推理系列产品支持Atlas A2 训练系列产品 / Atlas A2 推理系列产品支持此外从源码中的is_lite_npu()判断见 python/pypto/operator.py可以推断算子内部还针对 Lite 形态的 SoCsoc_version为Kirin9030、KirinX90提供了专门的实现分支说明该算子在上述两类芯片路径上均有对应实现。函数原型与参数说明函数原型rms_norm(input: Tensor, gamma: Tensor None, epsilon: float 1e-6) - Tensor参数说明参数名输入/输出说明input输入源操作数。支持 PyPTO 支持的数据类型可以是任意 Shape 的Tensor[..., C]最后一个维度 C 通常表示通道数或特征数。gamma输入可选的缩放参数Shape 应为[C]与 input 最后一个维度大小一致。默认值为None即不做缩放。epsilon输入数值稳定性常数默认值为1e-6。其中参数gamma是可选的从源码 python/pypto/operator.py 可以看到当gamma is not None时内部会将[C]形状的 gamma 通过reshape广播为[1, ..., 1, C]保持输入 rank再与归一化结果逐元素相乘因此支持任意维度的输入 Tensor。返回值说明返回归一化后的 TensorShape 与输入 Tensorinput相同若输入数据类型非 FP32输出 Tensor 会被转换回输入 Tensor 的原始数据类型见 python/pypto/operator.py。调用示例与结果验证基本调用示例x pypto.tensor([2, 4], pypto.DT_FP32) gamma pypto.tensor([4], pypto.DT_FP32) y pypto.rms_norm(x, gamma)结果示例如下输入数据x: [[1, 2, 3, 4], [5, 6, 7, 8]] 输入数据gamma: [1, 1, 1, 1] 输出数据y: [[0.3651, 0.7302, 1.0954, 1.4605], [0.7580, 0.9097, 1.0613, 1.2129]]这里gamma取全 1 向量因此输出等价于纯归一化结果。读者可自行验算第一行均方值为(14916)/4 7.5RMS 为sqrt(7.51e-6) ≈ 2.73861/2.7386 ≈ 0.36512/2.7386 ≈ 0.7302与示例输出完全一致。在 JIT Kernel 中的用法结合单元测试 python/tests/ut/kirin/common_rmsnorm.py可以看出一段完整的可运行 Kernel 写法pypto.frontend.jit( codegen_options{soc_version: soc_version}, runtime_options{run_mode: pypto.RunMode.SIM}, ) def kernel(a: pypto.Tensor([...], dtype), out: pypto.Tensor([...], dtype), gamma: pypto.Tensor([...], dtype)): pypto.set_vec_tile_shapes(*tile_shapes) out[:] pypto.rms_norm(a, gamma)测试用例覆盖了 FP16/FP32、有无 gamma、以及从一维到四维的多种 Shape如(2, 32)、(2, 4, 160)、(5, 2, 4, 176)、(1, 2, 64, 2048)等可作为该算子在实际 Kernel 中的参考用法。源码级实现两条计算路径查看 python/pypto/operator.py 可以看到rms_norm对外是统一入口内部根据目标 SoC 自动分派到两条实现路径1. FP32 提升计算路径rms_norm_fp32_castdef rms_norm_fp32_cast(input: Tensor, gamma: Tensor None, epsilon: float 1e-6) - Tensor: in_dtype input.dtype x pypto.cast(input, pypto.DT_FP32) n x.shape[-1] y pypto.sqrt(pypto.sum(x * x * (1.0 / n), -1, keepdimTrue) epsilon) ones pypto.full(y.shape, 1.0, pypto.DT_FP32) y pypto.div(x * ones, y, pypto.PrecisionType.INTRINSIC) if gamma is not None: rank input.dim shape [1] * rank shape[-1] gamma.shape[0] g pypto.cast(pypto.reshape(gamma, shape), pypto.DT_FP32) y * g if in_dtype ! pypto.DT_FP32: y pypto.cast(y, in_dtype) return y要点先把输入cast到DT_FP32全程以 FP32 计算 RMS 与除法最后再cast回输入原始类型保证 FP16/BF16 等低精度输入下的归一化精度除法使用pypto.PrecisionType.INTRINSIC内建精度可在满足精度要求的前提下获得更优性能gamma 先 reshape 成可广播形状再参与逐元素乘法。2. 轻量化 NPU 无转换路径rms_norm_no_castdef rms_norm_no_cast(input: Tensor, gamma: Tensor None, epsilon: float 1e-6) - Tensor: in_dtype input.dtype n input.shape[-1] y pypto.sqrt(pypto.sum(input * input * (1.0 / n), -1, keepdimTrue) epsilon) ones pypto.full(y.shape, 1.0, in_dtype) y pypto.div(input * ones, y, pypto.PrecisionType.INTRINSIC) if gamma is not None: rank input.dim shape [1] * rank shape[-1] gamma.shape[0] g pypto.reshape(gamma, shape) y * g return y该路径不做 FP32 提升直接在输入数据类型上完成计算避免 cast 带来的额外开销适用于对精度敏感度较低、追求极致性能的 Lite NPU 场景。两条路径的选择由is_lite_npu()根据codegen_options中的soc_version自动完成使用者无需关心。典型应用场景LLM 归一化与 QKV 融合RMSNorm 是当前大语言模型LLM中替代 LayerNorm 的主流归一化方式。在仓库的融合算子测试 python/tests/ut/kirin/common_qkv_rmsnorm_rope_scatternd.py 中可以看到pypto.rms_norm被直接用于 Q、K 向量的归一化并与 RoPE 旋转位置编码、KV Cache 写入融合进同一个 Kerneldef apply_rmsnorm(input_tensor: pypto.Tensor, gamma: pypto.Tensor, epsilon: float 1e-6) - pypto.Tensor: # gamma 入参 shape 为 (1, 128, 1, 1)flatten 到 (128,) 给 rms_norm 用 normalized pypto.rms_norm(input_tensor, pypto.reshape(gamma, [128]), epsilonepsilon) return normalized该场景给出了两个值得借鉴的工程细节gamma 形状适配权重在模型中以(1, 128, 1, 1)形式存储使用前通过pypto.reshape展平为[128]正好对应rms_norm要求的[C]形状融合收益将 RMSNorm 与后续 transpose、RoPE、index_put_KV Cache 写入放在同一 Kernel 内避免中间结果在 GM 与片上存储之间反复搬运。性能调优建议关于该算子的性能可参考 PyPTO 编程指南 性能调优 中的明确建议归约类计算Reduce 运算如 sum、max、min 等尽可能不要在归约轴上进行切分。例如输入 Shape 为 (56, 1024) 的 RMSNorm它的最后一维 TileShape 应当设为 1024。由于rms_norm本质上是沿最后一个维度的归约求和对 reduce 轴切分会导致多个子图的输出需要在同一子图中再次 reduce引入 GM 搬运与调度开销而将最后一维的 TileShape 设为整个 C则上下游子图可以合并消除 GM 搬运与调度开销。在编写 Kernel 时应通过pypto.set_vec_tile_shapes将最后一维保持为完整宽度这一点在单元测试 python/tests/ut/kirin/common_rmsnorm.py 的用例中也有体现如输入(4, 128)、tile shape(1, 128)归约轴保持 128 不切分。正确性验证方式仓库单元测试 python/tests/ut/kirin/common_rmsnorm.py 给出了该算子的标准 golden 校验流程def _compute_golden_rmsnorm(a, gammaNone, eps1e-6): n a.shape[-1] a_f a.float() rms torch.sqrt(torch.sum(a_f * a_f * (1.0 / n), -1, keepdimTrue) eps) y a_f / rms if gamma is not None: y y * gamma.float() return y.to(a.dtype)校验维度包括check_nan检查输出无 NaN以及输出与 golden 的余弦相似度不低于0.9999。若要在自己的开发环境中验证pypto.rms_norm可以参照该流程随机生成输入FP16/FP32、有无 gamma、不同 Shape用 PyTorch 按上述公式计算 golden再与 PyPTO 输出做余弦相似度对比。小结pypto.rms_norm(input, gammaNone, epsilon1e-6)沿最后一个维度执行 RMSNorm支持任意 Shape 的Tensor[..., C]与可选 gamma 缩放算子内部自动分派非 Lite NPU 走 FP32 提升路径保证精度Lite NPUKirin9030/KirinX90走无转换路径追求性能输出统一恢复为输入原始数据类型在 LLM 场景中可无缝嵌入 QKV 融合 KernelRMSNorm RoPE KV Cache 写入注意将 gamma 适配为[C]形状性能上牢记归约轴不切分原则将最后一维 TileShape 保持为完整 C可消除 GM 搬运与调度开销。【免费下载链接】pyptoPyPTO发音: pai p-t-oParallel Tensor/Tile Operation编程范式。项目地址: https://gitcode.com/cann/pypto创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表