ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

CANN ascend-transformer-boost 中的 NormRopeReshape 融合算子:RMSNorm→RoPE→ReshapeAndCache 三阶段流水解析

CANN ascend-transformer-boost 中的 NormRopeReshape 融合算子:RMSNorm→RoPE→ReshapeAndCache 三阶段流水解析 CANN ascend-transformer-boost 中的 NormRopeReshape 融合算子RMSNorm→RoPE→ReshapeAndCache 三阶段流水解析【免费下载链接】ascend-transformer-boost本项目是CANN提供的是一款高效、可靠的Transformer加速库基于华为Ascend AI处理器提供Transformer定制化场景的高性能融合算子。项目地址: https://gitcode.com/cann/ascend-transformer-boost导读NormRopeReshape 是 CANN ascend-transformer-boost 提供的一个multi_stage_fusion 型融合算子将 Transformer 解码阶段频繁出现的 RMSNorm 归一化、RoPE 旋转位置编码与 KV Cache 写入ReshapeAndCache三个阶段压缩为单 kernel 执行避免多次访存与 kernel 启动开销。本文以 .agent/knowledge/ops/norm/norm_rope_reshape/index.md 知识条目为骨架结合源码、参数定义与测试用例完整说明该算子的参数约束、7 输入 1 输出的张量布局、Shape 校验规则、Graph 构建链路及运行平台限制帮助读者在 Atlas 800I A2 推理场景中正确配置与使用该算子。一、算子定位与适用场景NormRopeReshape 是norm类别下的 S 级融合算子tier: Stype: fusion其核心价值在于减少中间张量的写回与重读单独的rms_norm、rope、kv_cache三个算子依次执行时RMSNorm 的输出需要落回 HBMRoPE 阶段再读入并写出ReshapeAndCache 阶段再读入写入 KV Cache存在多次全量访存融合后RMSNorm 输出直接留在片上寄存器/L1 中供 RoPE 消费RoPE 结果就地写入 KV Cache整条流水在单 kernel 内完成显著降低带宽压力与调度开销。该算子在项目中的典型应用是自回归解码inference阶段的 key 侧处理对归一化后的 key 向量施加旋转位置编码再按slotMapping指定的位置写入分页 KV Cache为后续 PagedAttention 等注意力算子准备缓存数据。从知识条目的source元数据可知该算子的核心实现位于类型仓库路径算子Operationsrc/ops/ops_infer/norm_rope_reshape/融合 Kernelsrc/kernels/mixkernels/rms_norm_and_rope_and_reshape_and_cache/对外参数头文件include/atb/infer_op_params.h说明知识条目中 Kernel 路径写作src/kernels/mixkernels/laser_attention但从 OpsRunner 源码 构建的图节点名RmsNormAndRopeAndReshapeAndCacheOperation及 mixkernels 目录 的实际实现来看融合 kernel 的落地路径为src/kernels/mixkernels/rms_norm_and_rope_and_reshape_and_cache/本文以源码实测为准。二、源码文件地图与阅读顺序知识条目给出了完整的源码阅读路径结合仓库实际文件各文件职责如下#文件角色1norm_rope_reshape_operation.hOperation 定义输入/输出数量、InferShape 与参数校验接口2norm_rope_reshape_operation.cppCreateOperation 入口、三阶段参数校验epsilon 校验 → 平台校验 → Shape 校验3norm_rope_reshape_ops_runner.hOpsRunner 声明4norm_rope_reshape_ops_runner.cppGraph 构建挂载RmsNormAndRopeAndReshapeAndCacheOperation融合节点推荐按知识条目的阅读顺序展开[1] operation.h → 融合算子接口 [2] operation.cpp → 三阶段参数校验 [3] ops_runner.cpp → Graph 构建RMSNorm→RoPE→Reshape三、参数结构与约束NormRopeReshapeParam对外参数定义在 include/atb/infer_op_params.h#L2696-L2716知识条目中的结构体与仓库实现一致//! \struct NormRopeReshapeParam //! \brief 融合rmsnorm、rope、reshapeAndCache。 //! \warning 仅Atlas 800I A2推理产品支持该算子 struct NormRopeReshapeParam { //! \brief precisionMode精度模式。 uint32_t precisionMode 0; //! \brief rotaryCoeff算子内Rope部分计算的旋转系数。 uint32_t rotaryCoeff 2; //! \brief epsilon归一化时加在分母上防止除零。 float epsilon 1e-5; //! \brief 预留参数 uint8_t rsv[16] {0}; };参数明细与约束参数类型默认值约束与含义precisionModeuint320精度模式透传至底层 Kernel 参数rotaryCoeffuint322RoPE 旋转系数控制旋转频率基数epsilonfloat1e-5RMSNorm 分母防除零项必须大于 0rsvuint8[16]全 0预留参数不可赋非零值两个关键校验逻辑在 CreateOperation 模板特化 中体现epsilon 合法性校验源码使用std::fabs(opParam.epsilon) THRESHOLD判定其中THRESHOLD std::numeric_limitsfloat::min()即 float 最小正规格化数 1.17549e-38。也就是说 epsilon 的绝对值不能小于该阈值否则报错Invalid epsilon, its recommended to init a nonzero value for eps.并返回ERROR_INVALID_PARAM平台校验通过GetSingletonConfig().Is910B()判断非 910B即非 Atlas 800I A2平台直接拒绝创建报错NormRopeReshapeOperation only supports Atlas 800I A2 inference products。平台限制是硬性约束其他平台一律返回ERROR_INVALID_PARAM。参数在底层通过SetNormRopeReshapeParam一对一透传为AtbOps::OpParam::RmsNormAndRopeAndReshapeAndCache的precisionMode / epsilon / rotaryCoeff字段见 norm_rope_reshape_ops_runner.cpp#L18-L25并在 Runner 构造时打印日志便于排查。四、7 输入 1 输出的张量布局从 ops_configs/atb_ops_info.ini#L5870-L5894 可确认算子的 I/O 规格nd格式与知识条目的输入描述完全对应序号名称dtype格式维度含义input0xfloat16nd3 维[batch, headDim, hiddenSize]RMSNorm 输入key 向量input1gammafloat16nd1 维[hiddenSize]RMSNorm 缩放权重input2keyRopefloat16nd2 维RoPE 旋转向量key 部分input3cosfloat16nd2 维RoPE 余弦表input4sinfloat16nd2 维RoPE 正弦表input5slotMappingint32nd1 维KV Cache 槽位映射input6keycacheinfloat16nd4 维[?, ?, 1, cacheSize]KV Cache 写入缓冲output0keycacheoutfloat16nd4 维写缓存后的输出GetInputNum()返回 7、GetOutputNum()返回 1见 norm_rope_reshape_operation.cpp#L66-L74。需要说明的是知识条目中列出的 dtype 为 fp16 主路径但底层 kernel 的CanSupport同时接受float16 与 bf16仅slotMapping第 5 个输入强制要求TENSOR_DTYPE_INT32见 rms_norm_and_rope_and_reshape_and_cache_kernel.cpp#L38-L50。五、Shape 校验规则910B 专用路径算子对输入 Shape 有严格约束全部集中在InferShapeCheckImpl→CheckDim910B/NormRopeReshapeCheckImpl910Bnorm_rope_reshape_operation.cpp#L83-L181以及底层 kernel 的CheckNormRopeCacherms_norm_and_rope_and_reshape_and_cache_operation.cpp#L63-L110。5.1 维度数量约束CheckDim910B输入要求xdimNum 3gammadimNum 1keyRope/cos/sindimNum 2slotMappingdimNum 1keycacheindimNum 45.2 维度一致性约束gamma与x的 dtype、format 必须一致且gamma最后一维等于x的 embed 维GammaBetaTensorCheckx.dims[0] keyRope.dims[0] cos.dims[0]batch 维一致keyRope.dims[1] cos.dims[1] sin.dims[1]RoPE 向量列维一致核心等式keycachein.dims[3] x.dims[2] keyRope.dims[1]即 Cache 的最后一维必须等于 RMSNorm 输入维与 RoPE 旋转维之和表示归一化后的 key 旋转增量拼接写入x.dims[2] % 16 0且gamma.dims[0] % 16 0末维必须是 16 的倍数满足向量化对齐底层 kernel 进一步要求x.dims[1] 1headDim 固定为 1且x.dims[2]、keyRope.dims[1]、keycachein.dims[3]均为 16 的倍数、任意输入维不得为 0。5.3 总缓存上限约束InferShapeCheckImpl开头还有一个容量守卫if (ELEVEN * inTensorDescs.at(0).shape.dims[DIM_TWO] * FLOAT16SIZE inTensorDescs.at(6).shape.dims[DIM_THREE] * FLOAT16SIZE MAXUBSIZE) { // MAXUBSIZE 196352即11 * x.dims[2] * 2 keycachein.dims[3] * 2 196352单位字节超出则返回ERROR_INVALID_TENSOR_DIM。这是融合 kernel 片上资源UB 空间的硬上限配置较大 hidden size 时需留意。5.4 输出 Shape 推断InferShapeImpl的实现非常简洁输出 desc 直接拷贝第 6 个输入keycachein的 descnorm_rope_reshape_operation.cpp#L76-L81底层 kernel 的InferShapeImpl同样执行outTensors[0].desc inTensor(6).desc。SetupCheckImpl中通过CheckOutTensorSame强制输出与keycachein的 dtype、format、dimNum、各维 size 完全一致。六、三阶段融合流水Computation Pipeline知识条目以 YAML 形式给出了计算流水定义pipeline_type: multi_stage_fusion stages: - stage: RMSNorm归一化 note: fp16 in → fp32 acc → fp16 out - stage: RoPE旋转位置编码 note: rotaryCoeff 控制旋转系数 - stage: ReshapeAndCache note: KV Cache 写入 note: 三阶段融合算子单 kernel 内完成。Golden 生成需匹配中间 dtype。三个阶段在单 kernel 内顺序执行RMSNorm对x按epsilon做归一化内部使用 fp32 累加以保证精度fp16 输入 → fp32 中间累加 → fp16 输出gamma作为逐元素缩放权重RoPE结合cos/sin旋转表与keyRope向量按rotaryCoeff决定的旋转系数对归一化后的 key 施加旋转位置编码ReshapeAndCache依据slotMapping将编码后的 key 写入keycachein指定的缓存槽位产出keycacheout。对 Golden 生成的关键提示知识条目明确标注由于中间阶段存在 fp16→fp32→fp16 的精度转换编写精度比对脚本时必须匹配中间 dtype否则容易产生 epsilon 级别的精度误差误报。仓库测试通过 CSV 方式驱动见下文第七节。七、执行路径与平台支持知识条目给出的执行路径为NormRopeReshapeOperation::CreateRunner() └── → NormRopeReshapeOpsRunner单一路径仅 A2 Kernel: RmsNormAndRopeAndReshapeAndCacheKernel融合 kernel源码印证了这条链路CreateRunner 直接构造NormRopeReshapeOpsRunner没有任何平台分支——因为平台限制在CreateOperation阶段就已拦截能走到 Runner 的必然是 A2NormRopeReshapeOpsRunner继承自OpsRunner构造时调用SetNormRopeReshapeParam透传参数再通过BuildNormRopeReshapeGraph构建包含 1 个节点RmsNormAndRopeAndReshapeAndCacheOperation的 MKI 图norm_rope_reshape_ops_runner.cpp#L27-L50并通过REG_RUNNER_TYPE/REG_OP_PARAM完成注册底层 MKI 侧RmsNormAndRopeAndReshapeAndCacheOperation::GetBestKernel返回RmsNormAndRopeAndReshapeAndCacheKernelrms_norm_and_rope_and_reshape_and_cache_operation.cpp#L43-L49Kernel 通过 tiling 目录下的RmsNormAndRopeAndReshapeAndCacheTiling完成分片计算GetTilingSize返回RmsNormAndRopeAndReshapeAndCacheTilingData大小InitImpl调用 Tiling 函数见 kernel 文件。平台Runner限制Atlas 800I A2910BOpsRunner唯一支持路径其他平台—CreateOperation 阶段直接拒绝八、测试用例与 Golden 约束仓库在 tests/apitest/opstest/csv/norm_rope_reshape.csv 中提供了可直接运行的 CSV 用例SocVersion 为Ascend910B是理解张量布局的最佳样例输入: 64,1,512; 512; 64,64; 64,64; 64,64; 64; 192,128,1,576 输出: 192,128,1,576 参数: {epsilon: 1e-8}用例要点x为[64, 1, 512]batch64, headDim1, hidden512gamma为[512]keyRope/cos/sin均为[64, 64]slotMapping为[64]keycachein为[192, 128, 1, 576]注意keycachein.dims[3] 576 512 64正好满足x.dims[2] keyRope.dims[1]的拼接等式且 512、64、576 均为 16 的倍数第 3 个用例将x的 dtype 改为 int32 后期望错误为ERROR_INVALID_TENSOR_INI_MATCH验证 dtype 不匹配会被拦截第 4 个用例将预留参数rsv置[1]且输入数为 0期望错误为ERROR_INVALID_PARAM说明预留参数必须保持全 0slotMapping的数据生成范围固定为0,0全 0 槽位DataGenType支持customize与random两种方式。测试侧实现位于 tests/apitest/opstest/python/operations/norm_rope_reshape/test_norm_rope_reshape.py高层冒烟用例见 tests/high_level_test/NormRopeReshapeOperation/Smoke/test_norm_rope_reshape_smoke.py。另外该算子的参数序列化由 src/atb/utils/param_to_json.cpp 中的OpParamToJson支持GetParamJson直接调用它完成 JSON 化输出。九、已知限制与关联算子已知问题#问题状态1仅 Atlas 800I A2 推理产品支持代码中Is910B()强校验限制lim2三阶段融合Golden 生成需匹配中间 dtypefp32 累加注意关联算子算子与本算子的关系rms_norm融合的第一阶段归一化rope融合的第二阶段旋转位置编码kv_cacheReshapeAndCache 阶段关联KV Cache 写入如果只需要其中单个阶段的能力可直接使用 rms_norm、rope、kv_cache 等独立算子当解码场景对延迟敏感且三者连续出现时NormRopeReshape是更优的融合选择。十、使用建议小结确认平台仅 Atlas 800I A2910B推理产品可用其他平台会在CreateOperation阶段报ERROR_INVALID_PARAM初始化参数epsilon必须为绝对值不小于 float 最小正规格化数1.17549e-38的非零值rotaryCoeff默认 2rsv保持全 0核对 Shape重点检查keycachein.dims[3] x.dims[2] keyRope.dims[1]、各维 16 字节对齐、x.dims[1] 1以及11 * x.dims[2] * 2 keycachein.dims[3] * 2 196352的 UB 容量上限输出规格keycacheout必须与keycachein的 dtype、format、Shape 完全一致底层 kernel 支持 fp16 / bf16slotMapping固定 int32精度比对生成 Golden 时匹配中间 dtypefp16 in → fp32 acc → fp16 out避免 epsilon 级误报。【免费下载链接】ascend-transformer-boost本项目是CANN提供的是一款高效、可靠的Transformer加速库基于华为Ascend AI处理器提供Transformer定制化场景的高性能融合算子。项目地址: https://gitcode.com/cann/ascend-transformer-boost创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进