ascend-transformer-boost 中 rope_grad 算子源码导读:RoPE 反向训练算子的完整实现链路
【免费下载链接】ascend-transformer-boost本项目是CANN提供的是一款高效、可靠的Transformer加速库,基于华为Ascend AI处理器,提供Transformer定制化场景的高性能融合算子。项目地址: https://gitcode.com/cann/ascend-transformer-boost
本文围绕 ATB(ascend-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、单节点的简单算子 |
| 文件数 | 4 | Op 定义 2 个文件 + Ops Runner 2 个文件 |
| Runner 类型 | OpsRunner, Operation | 走原生 Ops(MKI)执行路径,由 Operation 创建 OpsRunner |
| ACLNN | no | 不提供 ACLNN 接口 |
从算子的输入输出组织方式看(详见 rope_grad_ops_runner.cpp 中的SetupKernelGraph),它接收 4 个输入——Q 嵌入梯度qEmbeddedGrad、K 嵌入梯度kEmbeddedGrad、余弦表cos、正弦表sin——并输出 2 个结果qGrad、kGrad,这正是 RoPE 前向旋转的逆过程所需的数据形态。
二、文件清单:4 个源文件的角色分工
路由文档第一节的文件清单完整列出了该算子在 Host 侧的全部实现文件,本文将其扩展为带源码链接的完整清单:
| # | 文件 | 角色 | 关键内容 |
|---|---|---|---|
| 1 | rope_grad_operation.cpp | Operation 定义(实现) | InferShapeImpl、DimCheck/ParamCheck校验、CreateRunner()决策逻辑 |
| 2 | rope_grad_operation.h | Operation 定义(声明) | RopeGradOperation类声明,继承自OperationBase |
| 3 | rope_grad_ops_runner.cpp | Ops Runner(实现) | Kernel Graph 构建、参数变更处理、类型注册 |
| 4 | rope_grad_ops_runner.h | Ops 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 ¶m); uint32_t GetInputNum() const override; uint32_t GetOutputNum() const override; train::RopeGradParam GetParam() const; void SetParam(const train::RopeGradParam ¶m); protected: Status InferShapeImpl(const SVector<TensorDesc> &inTensorDescs, SVector<TensorDesc> &outTensorDescs) const override; std::shared_ptr<Runner> CreateRunner(Context &context) const override; Status InferShapeCheckImpl(const SVector<TensorDesc> &inTensorDescs) const override; Status SetupCheckImpl(const SVector<Tensor> &inTensors, const SVector<Tensor> &outTensors) const override; nlohmann::json GetParamJson() const override; private: Status ParamCheck(const SVector<TensorDesc> &inTensorDescs) const; Status DimCheck(const SVector<TensorDesc> &inTensorDescs) const; train::RopeGradParam param_; };从签名可以看出该算子的"契约":
- 输入输出数量:
GetInputNum()返回IN_TENSOR_NUM = 4,GetOutputNum()返回OUT_TENSOR_NUM = 2(见 rope_grad_operation.cpp); - InferShape 签名:
InferShapeImpl(const SVector<TensorDesc> &, SVector<TensorDesc> &),基于张量描述推导输出形状; - 两层校验钩子:
InferShapeCheckImpl(形状阶段)与SetupCheckImpl(Setup 阶段)都复用私有方法DimCheck+ParamCheck; - Runner 创建点:
CreateRunner(Context&)是该算子从"图描述"走向"执行体"的关键决策入口。
3.2 第二步:rope_grad_operation.cpp —— CreateRunner() 决策逻辑与参数校验
rope_grad_operation.cpp 承载了全部 Host 侧逻辑。按路由文档提示,重点是CreateRunner()的决策逻辑:
std::shared_ptr<Runner> RopeGradOperation::CreateRunner(Context &context) const { ContextBase *contextBase = dynamic_cast<ContextBase *>(&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.MallocRunner<RopeGradOpsRunner, train::RopeGradParam>(param_); if (!runner) { ATB_LOG(DEBUG) << "MallocRunner from pool failed!"; return std::make_shared<RopeGradOpsRunner>(param_); } return std::shared_ptr<Runner>(runner, &pool { pool.FreeRunner(runner); }); }(见 rope_grad_operation.cpp)决策逻辑分三步:
- 将 Context 向下转型为 ContextBase,获取其 Runner 资源池的访问入口;
- 通过名称查找 Runner 类型:
RunnerTypeRegister::GetRunnerTypeIdx("RopeGradOpsRunner")将字符串类型名解析为类型下标,再取出对应的RunnerPool; - 优先从池中复用 Runner:
pool.MallocRunner<RopeGradOpsRunner, train::RopeGradParam>(param_)尝试以参数模板分配(可复用)一个已构造的 Runner;池耗尽时退化为std::make_shared新建。返回的shared_ptr附带自定义删除器,析构时调用pool.FreeRunner归还对象——这是 ATB 中 Runner 对象池化的通用模式。
形状推导与校验。InferShapeImpl非常直接:两个输出分别继承第一、二个输入的描述(源码):
Status RopeGradOperation::InferShapeImpl(const SVector<TensorDesc> &inTensorDescs, SVector<TensorDesc> &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 = 2) | ERROR_INVALID_TENSOR_DIM |
| Q/K 梯度同形 | 输入 0 与输入 1 的dims[0]、dims[1]必须相等 | ERROR_INVALID_TENSOR_SIZE |
| cos/sin 同形 | 输入 2 与输入 3 的两个维度必须相等 | ERROR_INVALID_TENSOR_SIZE |
| hiddenSize 对齐 | 输入 0/1 的dims[1](hiddenSize)必须能被 128 整除 | ERROR_INVALID_TENSOR_DIM |
| headSize | cos(输入 2)的dims[1]必须等于HEAD_SIZE = 128 | ERROR_INVALID_TENSOR_DIM |
| qSeqLen 非空 | param_.qSeqLen.size() > 0 | ERROR_INVALID_PARAM |
| qSeqLen 合法区间 | 每个元素满足0 < qSeqLen[i] <= cos.dims[0](最大序列长度) | ERROR_INVALID_PARAM |
硬件平台限制。文件顶部还有一个匿名命名空间的全局ParamCheck(源码):
bool ParamCheck(const atb::train::RopeGradParam &opParam) { if (!atb::GetSingleton<atb::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_ = GetSingleton<AtbOperationIrCfg>().GetOperationIr("RopeGradOperation")从全局 IR 配置单例中取出该算子的 IR 描述;GetParamJson()则通过OpParamToJson(param_)将参数转为 JSON 用于日志与调试。
3.3 第三步:rope_grad_ops_runner.h —— 原生 Ops 执行接口
rope_grad_ops_runner.h 声明了执行侧的RopeGradOpsRunner:
class RopeGradOpsRunner : public OpsRunner { public: explicit RopeGradOpsRunner(const train::RopeGradParam ¶m); ~RopeGradOpsRunner() override; void SetParam(const Mki::Any ¶m) override; protected: Status SetupKernelGraph(const OpsTensorPack &opsTensorPack) override; private: train::RopeGradParam param_; };它继承自 OpsRunner(ATB 中"原生 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::vector<int32_t> qSeqLen),并各自定义了operator==用于参数变更检测; - 参数热更新:
SetParam中先做Mki::AnyCast<train::RopeGradParam>解包,newParam == param_不等时才更新并置isParamUpdated_ = true(源码),供上层决定是否重新执行 Tiling; - 类型注册:文件尾部的
REG_RUNNER_TYPE(RopeGradOpsRunner)把 Runner 类名注册进RunnerTypeRegister(供CreateRunner按名查找),REG_OP_PARAM(AtbOps::OpParam::RopeGrad)把参数类型注册进 MKI 的参数系统——这两行宏是该算子能被框架"按名发现"的关键。
四、参数详解:RopeGradParam
Host 侧参数结构定义在 include/atb/train_op_params.h:
//! \struct RopeGradParam //! \brief 旋转位置编码处理的反向。 struct RopeGradParam { //! \brief 存储unpad场景下每个batch实际qSseqlen的值。size不能为0 std::vector<int32_t> qSeqLen; uint8_t rsv[8] = {0}; // 预留参数 };参数语义与使用约束:
qSeqLen:unpad(变长)场景下每个 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::RopeGradOperation:MKI 算子入口(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.cpp(Kernel 侧) 定义了设备侧的AtbOps::RopeGradOperation(与 Host 侧类同名但不同命名空间):GetInputNum/GetOutputNum固定返回 4 与 2;CheckRopeGrad用MKI_CHECK宏复核了与 Host 侧一致的维度约束(输入 0/1 同形且 hiddenSize 被 128 整除、cos/sin 同形且dims[1] == 128);InferShapeImpl同样令输出继承输入 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 维,输出为 FLOAT16;GetTilingSize按batch = qSeqLen.size()计算 Tiling 缓冲大小:
uint32_t batch = AnyCast<OpParam::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 前向旋转的伴随(转置)运算;
- 三级流水:
CopyIn(DataCopy从 GM 拷入 UB 队列)→Compute(VEC 指令: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对齐、headSize=128、qSeqLen 合法性 └─ CreateRunner RunnerPool 池化分配 RopeGradOpsRunner └─ RopeGradOpsRunner::SetupKernelGraph └─ 单节点 "RopeGradOperation" + OpParam::RopeGrad{qSeqLen} Kernel 侧(src/kernels/mixkernels/rope_grad/) AtbOps::RopeGradOperation(InferShape / GetBestKernel) └─ AtbOps::RopeGradKernel(CanSupport:FP16/ND/2D;Tiling 大小按 batch 计) └─ AscendC RopeGrad:head 级核间并行,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),仅供参考