ascend-transformer-boost 中 rope_grad 算子源码导读RoPE 反向训练算子的完整实现链路【免费下载链接】ascend-transformer-boost本项目是CANN提供的是一款高效、可靠的Transformer加速库基于华为Ascend AI处理器提供Transformer定制化场景的高性能融合算子。项目地址: https://gitcode.com/cann/ascend-transformer-boost本文围绕 ATBascend-transformer-boost仓库中rope_grad训练侧算子的路由文档展开以文件清单 推荐阅读顺序 源码路径为骨架带你沿着src/ops/ops_train/rope_grad/下 4 个核心源文件逐层深入从RopeGradOperation的输入输出定义、InferShape签名与CreateRunner()决策逻辑到RopeGradOpsRunner的原生 Ops 执行接口与 Kernel Graph 构建直至src/kernels/mixkernels/rope_grad/下的 AscendC Kernel 实现。读完后你能够独立掌握 ATB 训练算子Operation → Runner → Kernel的三层结构与代码阅读方法并能复述 rope_grad 的张量约束、参数校验规则与设备侧三级流水计算流程。一、rope_grad 是什么训练侧的 RoPE 反向算子rope_grad是 ATB 面向 Transformer 训练场景提供的旋转位置编码RoPE反向算子。根据路由文档 .agent/knowledge/routing/rope_grad.md 的元信息它的分类与定位如下属性取值说明分类train训练侧算子位于src/ops/ops_train/目录复杂度S单 Kernel、单节点的简单算子文件数4Op 定义 2 个文件 Ops Runner 2 个文件Runner 类型OpsRunner, Operation走原生 OpsMKI执行路径由 Operation 创建 OpsRunnerACLNNno不提供 ACLNN 接口从算子的输入输出组织方式看详见 rope_grad_ops_runner.cpp 中的SetupKernelGraph它接收 4 个输入——Q 嵌入梯度qEmbeddedGrad、K 嵌入梯度kEmbeddedGrad、余弦表cos、正弦表sin——并输出 2 个结果qGrad、kGrad这正是 RoPE 前向旋转的逆过程所需的数据形态。二、文件清单4 个源文件的角色分工路由文档第一节的文件清单完整列出了该算子在 Host 侧的全部实现文件本文将其扩展为带源码链接的完整清单#文件角色关键内容1rope_grad_operation.cppOperation 定义实现InferShapeImpl、DimCheck/ParamCheck校验、CreateRunner()决策逻辑2rope_grad_operation.hOperation 定义声明RopeGradOperation类声明继承自OperationBase3rope_grad_ops_runner.cppOps Runner实现Kernel Graph 构建、参数变更处理、类型注册4rope_grad_ops_runner.hOps Runner声明RopeGradOpsRunner类声明继承自OpsRunner此外算子在 Kernel 侧还有配套的独立目录src/kernels/mixkernels/rope_grad/包含 Kernel 描述、Tiling 与 AscendC 算子实现属于 MKI 算子注册体系的一部分见第五节。三、推荐阅读顺序按路由文档的 4 步走读路由文档给出了明确的阅读顺序与各文件的关注点以下按此顺序逐一展开并给出对应的源码级依据。3.1 第一步rope_grad_operation.h —— 了解输入输出数量、InferShape 签名rope_grad_operation.h 中声明了 Host 侧的核心类RopeGradOperation它继承自atb::OperationBase持有train::RopeGradParam参数对象class RopeGradOperation : public OperationBase { public: explicit RopeGradOperation(const train::RopeGradParam param); uint32_t GetInputNum() const override; uint32_t GetOutputNum() const override; train::RopeGradParam GetParam() const; void SetParam(const train::RopeGradParam param); protected: Status InferShapeImpl(const SVectorTensorDesc inTensorDescs, SVectorTensorDesc outTensorDescs) const override; std::shared_ptrRunner CreateRunner(Context context) const override; Status InferShapeCheckImpl(const SVectorTensorDesc inTensorDescs) const override; Status SetupCheckImpl(const SVectorTensor inTensors, const SVectorTensor outTensors) const override; nlohmann::json GetParamJson() const override; private: Status ParamCheck(const SVectorTensorDesc inTensorDescs) const; Status DimCheck(const SVectorTensorDesc inTensorDescs) const; train::RopeGradParam param_; };从签名可以看出该算子的契约输入输出数量GetInputNum()返回IN_TENSOR_NUM 4GetOutputNum()返回OUT_TENSOR_NUM 2见 rope_grad_operation.cppInferShape 签名InferShapeImpl(const SVectorTensorDesc , SVectorTensorDesc )基于张量描述推导输出形状两层校验钩子InferShapeCheckImpl形状阶段与SetupCheckImplSetup 阶段都复用私有方法DimCheckParamCheckRunner 创建点CreateRunner(Context)是该算子从图描述走向执行体的关键决策入口。3.2 第二步rope_grad_operation.cpp —— CreateRunner() 决策逻辑与参数校验rope_grad_operation.cpp 承载了全部 Host 侧逻辑。按路由文档提示重点是CreateRunner()的决策逻辑std::shared_ptrRunner RopeGradOperation::CreateRunner(Context context) const { ContextBase *contextBase dynamic_castContextBase *(context); if (!contextBase) { ATB_LOG(DEBUG) context cast to contextBase failed!; return nullptr; } int64_t runnerTypeIdx RunnerTypeRegister::GetRunnerTypeIdx(RopeGradOpsRunner); RunnerPool pool contextBase-GetRunnerPool(runnerTypeIdx); Runner *runner pool.MallocRunnerRopeGradOpsRunner, train::RopeGradParam(param_); if (!runner) { ATB_LOG(DEBUG) MallocRunner from pool failed!; return std::make_sharedRopeGradOpsRunner(param_); } return std::shared_ptrRunner(runner, pool { pool.FreeRunner(runner); }); }见 rope_grad_operation.cpp决策逻辑分三步将 Context 向下转型为 ContextBase获取其 Runner 资源池的访问入口通过名称查找 Runner 类型RunnerTypeRegister::GetRunnerTypeIdx(RopeGradOpsRunner)将字符串类型名解析为类型下标再取出对应的RunnerPool优先从池中复用 Runnerpool.MallocRunnerRopeGradOpsRunner, train::RopeGradParam(param_)尝试以参数模板分配可复用一个已构造的 Runner池耗尽时退化为std::make_shared新建。返回的shared_ptr附带自定义删除器析构时调用pool.FreeRunner归还对象——这是 ATB 中 Runner 对象池化的通用模式。形状推导与校验。InferShapeImpl非常直接两个输出分别继承第一、二个输入的描述源码Status RopeGradOperation::InferShapeImpl(const SVectorTensorDesc inTensorDescs, SVectorTensorDesc outTensorDescs) const { outTensorDescs.at(0) inTensorDescs.at(0); outTensorDescs.at(1) inTensorDescs.at(1); return NO_ERROR; }即qGrad与qEmbeddedGrad同形、kGrad与kEmbeddedGrad同形——RoPE 反向只是逐元素变换不改变张量形状。校验逻辑集中在DimCheck与ParamCheck两个私有方法中源码完整规则如下校验项规则错误码输入维度4 个输入全部为 2 维INPUT_SHAPE_DIM 2ERROR_INVALID_TENSOR_DIMQ/K 梯度同形输入 0 与输入 1 的dims[0]、dims[1]必须相等ERROR_INVALID_TENSOR_SIZEcos/sin 同形输入 2 与输入 3 的两个维度必须相等ERROR_INVALID_TENSOR_SIZEhiddenSize 对齐输入 0/1 的dims[1]hiddenSize必须能被 128 整除ERROR_INVALID_TENSOR_DIMheadSizecos输入 2的dims[1]必须等于HEAD_SIZE 128ERROR_INVALID_TENSOR_DIMqSeqLen 非空param_.qSeqLen.size() 0ERROR_INVALID_PARAMqSeqLen 合法区间每个元素满足0 qSeqLen[i] cos.dims[0]最大序列长度ERROR_INVALID_PARAM硬件平台限制。文件顶部还有一个匿名命名空间的全局ParamCheck源码bool ParamCheck(const atb::train::RopeGradParam opParam) { if (!atb::GetSingletonatb::Config().Is910B()) { ATB_LOG(ERROR) RopeGradOperation is not supported in Atlas 800I A2 inference product.; return false; } return atb::OperationUtil::QSeqLenCheck(opParam.qSeqLen); }从源码看RopeGradOperation构造时会经OPERATION_PARAM_FUNCS宏路径触发该校验仅当设备通过Config::Is910B()判断时才允许创建算子实例否则输出not supported in Atlas 800I A2 inference product的错误日志并拒绝执行同时OperationUtil::QSeqLenCheck对qSeqLen再做一次通用约束检查如非空与长度上限。这意味着使用该算子的前提是部署在 910B 训练设备上。参数序列化。构造函数中operationIr_ GetSingletonAtbOperationIrCfg().GetOperationIr(RopeGradOperation)从全局 IR 配置单例中取出该算子的 IR 描述GetParamJson()则通过OpParamToJson(param_)将参数转为 JSON 用于日志与调试。3.3 第三步rope_grad_ops_runner.h —— 原生 Ops 执行接口rope_grad_ops_runner.h 声明了执行侧的RopeGradOpsRunnerclass RopeGradOpsRunner : public OpsRunner { public: explicit RopeGradOpsRunner(const train::RopeGradParam param); ~RopeGradOpsRunner() override; void SetParam(const Mki::Any param) override; protected: Status SetupKernelGraph(const OpsTensorPack opsTensorPack) override; private: train::RopeGradParam param_; };它继承自 OpsRunnerATB 中原生 Ops执行器的基类只需实现两个关键虚函数SetupKernelGraph将 ATB 的张量包翻译成 MKI 的 Kernel Graph节点 张量引用SetParam当上层图动态更新算子参数时把Mki::Any中的新参数解包并与旧参数比较。3.4 第四步rope_grad_ops_runner.cpp —— 原生 Ops 调用链 平台适配rope_grad_ops_runner.cpp 完成了ATB 参数 → MKI Kernel Graph的最后一跳。SetupKernelGraph的完整流程源码Status RopeGradOpsRunner::SetupKernelGraph(const OpsTensorPack opsTensorPack) { (void)opsTensorPack; kernelGraph_.inTensors.resize(IN_TENSOR_COUNT); // 4 kernelGraph_.outTensors.resize(OUT_TENSOR_COUNT); // 2 Mki::Tensor qEmbeddedGrad kernelGraph_.inTensors.at(0); Mki::Tensor kEmbeddedGrad kernelGraph_.inTensors.at(1); Mki::Tensor cos kernelGraph_.inTensors.at(2); Mki::Tensor sin kernelGraph_.inTensors.at(3); Mki::Tensor qGrad kernelGraph_.outTensors.at(0); Mki::Tensor kGrad kernelGraph_.outTensors.at(1); kernelGraph_.nodes.resize(1); auto ropeGradNode kernelGraph_.nodes.at(0); AtbOps::OpParam::RopeGrad ropeGradParam; ropeGradParam.qSeqLen param_.qSeqLen; ropeGradNode.opDesc {0, RopeGradOperation, ropeGradParam}; ropeGradNode.inTensors {qEmbeddedGrad, kEmbeddedGrad, cos, sin}; ropeGradNode.outTensors {qGrad, kGrad}; return NO_ERROR; }要点解读单节点图rope_grad 的 Kernel Graph 只有一个节点节点名为RopeGradOperation——这个名字与 Kernel 侧REG_OPERATION注册的类名对应见第五节参数桥接Host 侧的train::RopeGradParam定义于 include/atb/train_op_params.h中的qSeqLen被复制到 Kernel 侧的AtbOps::OpParam::RopeGrad定义于 src/kernels/include/atbops/params/rope_grad.h。两侧结构体字段一致均为std::vectorint32_t qSeqLen并各自定义了operator用于参数变更检测参数热更新SetParam中先做Mki::AnyCasttrain::RopeGradParam解包newParam param_不等时才更新并置isParamUpdated_ true源码供上层决定是否重新执行 Tiling类型注册文件尾部的REG_RUNNER_TYPE(RopeGradOpsRunner)把 Runner 类名注册进RunnerTypeRegister供CreateRunner按名查找REG_OP_PARAM(AtbOps::OpParam::RopeGrad)把参数类型注册进 MKI 的参数系统——这两行宏是该算子能被框架按名发现的关键。四、参数详解RopeGradParamHost 侧参数结构定义在 include/atb/train_op_params.h//! \struct RopeGradParam //! \brief 旋转位置编码处理的反向。 struct RopeGradParam { //! \brief 存储unpad场景下每个batch实际qSseqlen的值。size不能为0 std::vectorint32_t qSeqLen; uint8_t rsv[8] {0}; // 预留参数 };参数语义与使用约束qSeqLenunpad变长场景下每个 batch 的实际 Q 序列长度元素个数为 batch size。文档注释明确size 不能为 0结合 Host 侧ParamCheck的逐元素校验0 qSeqLen[i] maxSeqLen与 Kernel 侧的批次上限检查batch 0 batch 100000见 rope_grad_kernel.cpp取值区间为每个元素(0, maxSeqLen]、数组长度[1, 100000)rsv[8]8 字节预留区默认全零用于未来字段扩展时保持参数结构兼容Kernel 侧AtbOps::OpParam::RopeGrad只保留qSeqLen一个有效字段是 Tiling 与 Kernel 计算分块的核心依据。五、源码路径与 Kernel 侧实现路由文档第三节给出了三处源码入口均已在仓库中核实存在Op 目录src/ops/ops_train/rope_grad/——上节走读的 4 个文件Kernel 目录src/kernels/mixkernels/rope_grad/——MKI 算子注册 AscendC 实现参数头文件include/atb/train_op_params.h——对外 API 参数。Kernel 目录内部结构与执行链如下src/kernels/mixkernels/rope_grad/ ├── rope_grad_operation.cpp # AtbOps::RopeGradOperationMKI 算子入口InferShape、选核 ├── rope_grad_kernel.cpp # AtbOps::RopeGradKernel能力检查、Tiling 大小计算 ├── op_kernel/rope_grad.cpp # AscendC Kernel设备侧 CopyIn/Compute/CopyOut 三级流水 ├── tiling/ │ ├── rope_grad_tiling.cpp / .h # Tiling 计算 │ └── tiling_data.h # RopeGradTilingData / RopeGradSampleTilingData └── CMakeLists.txtMKI 算子入口。rope_grad_operation.cppKernel 侧 定义了设备侧的AtbOps::RopeGradOperation与 Host 侧类同名但不同命名空间GetInputNum/GetOutputNum固定返回 4 与 2CheckRopeGrad用MKI_CHECK宏复核了与 Host 侧一致的维度约束输入 0/1 同形且 hiddenSize 被 128 整除、cos/sin 同形且dims[1] 128InferShapeImpl同样令输出继承输入 0/1 的描述GetBestKernel按名返回RopeGradKernel末尾REG_OPERATION(RopeGradOperation)完成注册——这正是 Runner 侧节点名RopeGradOperation的落点两端以此衔接。Kernel 能力检查与 Tiling 大小。rope_grad_kernel.cpp 中RopeGradKernel::CanSupport要求4 入 2 出、参数类型为OpParam::RopeGrad、4 个输入全部为 FLOAT16、ND 格式、2 维输出为 FLOAT16GetTilingSize按batch qSeqLen.size()计算 Tiling 缓冲大小uint32_t batch AnyCastOpParam::RopeGrad(launchParam.GetParam()).qSeqLen.size(); MKI_CHECK(batch 0 batch BTACH_LIMIT, OpParam is invalid, return 0); return sizeof(RopeGradTilingData) sizeof(RopeGradSampleTilingData) * batch;即 Tiling 数据 1 份全局数据RopeGradTilingData 每 batch 1 份样本数据RopeGradSampleTilingData与 tiling_data.h 中的结构定义对应。AscendC 设备实现。op_kernel/rope_grad.cpp 实现了真正的反向计算值得逐段理解按 head 维度做核间并行Init中每个 AI Core 以GetBlockIdx() * headSize为偏移绑定一段 GM 缓冲源码即不同核并行处理不同的 head分块尺寸MAX_PROCESS_NUM 192 * 1024 / sizeof(half) / 8单次循环最大处理元素数rowsPerLoop MAX_PROCESS_NUM / headSize按rowsPerLoop行一轮地循环处理每个 batch 的qSeqLen行计算主体Compute 函数Muls(workLocal[mask], sinLocal, scalars, mask, currentloopRows, {1, 1, 8, 8}); // -sin后半 head Add(workLocal, cosLocal, workLocal, ...); // work cos - sin按半 head 掩码叠加 Mul(qgradLocal, qembedgradLocal, workLocal, ...); // qGrad qEmbeddedGrad * work Mul(kgradLocal, kembedgradLocal, workLocal, ...); // kGrad kEmbeddedGrad * work即先构造work cos - sin在 head 的后半部分施加-sin项再与嵌入梯度逐元素相乘得到 Q、K 梯度——这是 RoPE 前向旋转的伴随转置运算三级流水CopyInDataCopy从 GM 拷入 UB 队列→ComputeVEC 指令Muls/Add/Mul配合pipe_barrier(PIPE_V)流水→CopyOut写回 GM每个 batch 处理完执行pipe_barrier(PIPE_ALL)同步后累加cursumseqlen进入下一个 batch。六、测试与验证仓库为该算子提供了高层测试用例目录 tests/high_level_test/RopeGradOperation/按测试维度组织为三类Smoke/——冒烟用例验证基本路径可跑通Dtype_dataFormat/——数据类型与数据格式边界对应 Kernel 侧 FLOAT16 / ND 的强制约束Boundary_value/——边界值用例对应 hiddenSize 128 整除、headSize 128、qSeqLen 区间等校验规则。这些 CSV 驱动的用例与第五节列出的校验规则一一呼应可以作为验证算子约束是否被正确拦截的回归依据。七、小结从路由文档到源码的完整链路回看路由文档 .agent/knowledge/routing/rope_grad.md 的知识条目指引rope_grad 的更完整知识归档位于 .agent/knowledge/ops/train/rope_grad/index.md可作为延伸阅读。将本文走读内容收敛为一张链路图Host 侧src/ops/ops_train/rope_grad/ RopeGradOperation ├─ InferShapeImpl out[0]in[0], out[1]in[1] ├─ DimCheck/ParamCheck 2维、同形、128对齐、headSize128、qSeqLen 合法性 └─ CreateRunner RunnerPool 池化分配 RopeGradOpsRunner └─ RopeGradOpsRunner::SetupKernelGraph └─ 单节点 RopeGradOperation OpParam::RopeGrad{qSeqLen} Kernel 侧src/kernels/mixkernels/rope_grad/ AtbOps::RopeGradOperationInferShape / GetBestKernel └─ AtbOps::RopeGradKernelCanSupportFP16/ND/2DTiling 大小按 batch 计 └─ AscendC RopeGradhead 级核间并行CopyIn→Compute(cos−sin 乘法)→CopyOut掌握这套Operation 定义 → Runner 池化创建 → Kernel Graph 构建 → MKI 注册选核 → AscendC 三级流水的走读方法后你可以将其平移到 ATB 仓库中其他训练侧算子如fast_soft_max、strided_batch_matmul的源码阅读上针对 rope_grad 本身则应牢记其使用前提910B 设备、FP16/ND、headSize 128、unpad 场景的qSeqLen参数不可为空。【免费下载链接】ascend-transformer-boost本项目是CANN提供的是一款高效、可靠的Transformer加速库基于华为Ascend AI处理器提供Transformer定制化场景的高性能融合算子。项目地址: https://gitcode.com/cann/ascend-transformer-boost创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
