Ascend Transformer Boost 中 Repeat 算子实现解析:参数校验、形状推导与双 Runner 执行路径
2026/9/18 7:06:38 网站建设 项目流程

Ascend Transformer Boost 中 Repeat 算子实现解析:参数校验、形状推导与双 Runner 执行路径

【免费下载链接】ascend-transformer-boost本项目是CANN提供的是一款高效、可靠的Transformer加速库,基于华为Ascend AI处理器,提供Transformer定制化场景的高性能融合算子。项目地址: https://gitcode.com/cann/ascend-transformer-boost

本文以 CANN / ascend-transformer-boost 仓库中的 Repeat 算子路由知识文件为核心,完整梳理 Repeat 算子的 6 个源文件组织、推荐阅读顺序与执行链路;并结合 infer_op_params.h 的参数定义与 src/ops/ops_infer/repeat/ 下的真实实现,讲清该算子如何完成参数校验、输出形状推导,并按平台在 ACLNN 路径与原生 Ops 路径之间选择执行方式。读完后你可以独立定位 Repeat 算子的每个实现环节,并理解其在 Ascend 平台上双执行路径的切换机制。

1. 算子定位与文件组织

Repeat 算子在 ATB(Ascend Transformer Boost)中被归类为infer类别的推理侧算子,复杂度为 M 级,由 6 个源文件组成,支持三种 Runner 类型(OpsRunner、ACLNNRunner、Operation),并支持 ACLNN 接口。官方文档语义定义(见 infer_op_params.h)为:将输入 Tensor 的 Shape 按指定轴扩展指定倍数,即等价于torch.Tensor.repeat的广播复制语义——输出 y 的维度和 multiples 维度一致,每个维度大小为输入 x 广播到 multiples 维度后与 multiples 对应维度的乘积。

完整的文件清单如下(源自路由文件 repeat.md):

#文件角色
1repeat_aclnn_runner.cppACLNN Runner
2repeat_aclnn_runner.hACLNN Runner
3repeat_operation.cppOperation 定义
4repeat_operation.hOperation 定义
5repeat_ops_runner.cppOps Runner
6repeat_ops_runner.hOps Runner

这些文件均位于Op 目录src/ops/ops_infer/repeat/下,参数头文件为 include/atb/infer_op_params.h。需要说明的是,从源码结构看,Repeat 算子最终执行依赖的是 CANN 原生能力:ACLNN 路径调用aclnnRepeat两段式接口,原生 Ops 路径映射到系统ExpandOperation,而非常规的 ATB 自研 kernel 目录(路由文件中登记的 Kernel 目录src/kernels/mixkernels/laser_attention为路由文件中的关联记录,实际执行链路以上述两个 Runner 源码为准)。

配套的详细知识条目见 ops/other/repeat/index.md,主索引见 knowledge/README.md。

2. RepeatParam 参数定义与约束

Repeat 算子只有 1 个输入 Tensor、1 个输出 Tensor(IN_TENSOR_NUM = 1OUT_TENSOR_NUM = 1,见 repeat_operation.cpp),全部算子配置收敛在RepeatParam结构体中:

// include/atb/infer_op_params.h struct RepeatParam { //! 每一维度上扩展的倍数。 //! \warning //! - 支持在不超过两个维度上进行扩展 //! - multiples的维度小于等于8且需大于或等于输入x的维度,每一个元素要求大于0。 SVector<int64_t> multiples; //! 预留参数 uint8_t rsv[8] = {0}; };

(代码见 infer_op_params.h#L2017-L2030)

关键约束逐条说明:

  • multiples是逐维扩展倍数,其长度必须大于等于输入 Tensor 的维度数——不足的高位维度按 1 补齐后再相乘;
  • 维度上限multiples.size() <= MAX_DIM(即 8),超过会在构造算子时直接报ERROR_INVALID_PARAM
  • 取值范围:每个元素必须大于 0,否则报错退出;
  • 扩展维度数限制:官方约束为“支持在不超过两个维度上进行扩展”(即真正发生放大的维度个数不超过 2),这一点在源码的 shape 检查中通过repeatDimNum统计落地,下文详解;
  • rsv[8]为预留参数,调用方应保持为 0(CreateOperation入口会执行OP_PARAM_RSV_CHECK校验,见 repeat_operation.cpp#L29)。

3. Operation 层:算子创建与入口校验

RepeatOperation继承自OperationBase,对外暴露标准接口:输入/输出数量、形状推导、执行前检查与 Runner 创建,类定义见 repeat_operation.h。

工厂函数CreateOperation是算子创建入口(repeat_operation.cpp#L24-L56),它做了三件事:

  1. 平台预加载检查:若当前平台为ASCEND_950,会提前调用RepeatAclnnRunner::LoadMethod()确认 ACLNN 函数符号可加载,失败则提示检查 CANN 版本并返回ERROR_CANN_ERROR
  2. multiples维度校验multiples.size() > MAX_DIM(8)时报错;
  3. 逐元素正数校验:任一multiples[i] <= 0时报错。
if (Mki::PlatformInfo::Instance().GetPlatformType() == Mki::PlatformType::ASCEND_950) { if (RepeatAclnnRunner::LoadMethod() != NO_ERROR) { ATB_LOG(ERROR) << "Load aclnn function failed, please check your CANN version."; return ERROR_CANN_ERROR; } }

构造算子时还会通过单例AtbOperationIrCfg获取"RepeatOperation"的 IR 配置(repeat_operation.cpp#L58-L61),用于统一的算子 IR 描述。

4. 形状推导 InferShapeImpl:广播语义与溢出保护

InferShapeImpl是理解 Repeat 语义的核心(repeat_operation.cpp#L75-L96)。推导规则:

  • 输出 dtype 与 format 直接继承输入;
  • 输出维数 =param_.multiples.size()
  • 从输入的高位维度向低位对齐遍历(先取输入最后几维与 multiples 尾部对齐):
    • 若输入还有剩余维度(idx > 0):out_dims[i] = in_dims[idx-1] * multiples[i]
    • 若输入维度已用尽:out_dims[i] = multiples[i](即输入未覆盖的高位维直接取倍数值,相当于在维度 1 的张量上扩展);
  • 每一步乘法前先做int64_t上溢检查max / in_dim < multiple即判溢出),溢出返回ERROR_INVALID_PARAM

以测试用例in: [2,3,5] + multiples: [1,4,4]为例,输出为[2,12,20];若输入是[1,2,32,1]multiples = [1,1,1,2],则输出为[1,2,32,2]

5. 执行前检查:InferShapeCheckImpl 与 SetupCheckImpl

两处检查逻辑一致(repeat_operation.cpp#L98-L138),在形状推断阶段和 Setup 阶段各执行一次:

  1. 输入维数不得超过 multiples 长度,否则返回ERROR_INVALID_TENSOR_DIM
  2. 统计真正发生放大的维度数repeatDimNum:仅当multiples[i] > 1in_dims[i] > 1时才计数(对长度为 1 的维做复制不视为“扩展维度”);
  3. 总量约束repeatDimNum + inTensorDimNum <= MAX_DIM(8)repeatDimNum + multiples.size() <= 8
  4. Setup 阶段额外要求输出维数等于multiples.size()

这套检查把“不超过两个维度扩展”的产品约束与 8 维上限统一收敛在维度计数中,是调用前最容易踩坑的地方。

6. CreateRunner:双执行路径的平台决策

CreateRunner是整个算子分叉的关键(repeat_operation.cpp#L140-L147):

std::shared_ptr<Runner> RepeatOperation::CreateRunner(Context &context) const { (void)context; if (Mki::PlatformInfo::Instance().GetPlatformType() == Mki::PlatformType::ASCEND_950) { return std::make_shared<RepeatAclnnRunner>(param_); } return std::make_shared<RepeatOpsRunner>(param_); }

即:ASCEND_950 平台走RepeatAclnnRunner(ACLNN 两段式接口),其余平台走RepeatOpsRunner(原生 Ops 内核图路径)。这一决策在CreateOperation阶段已有预检,保证 950 平台上函数符号必然可用。

7. ACLNN Runner 路径:Workspace 计算与两段式调用

RepeatAclnnRunner继承自AclnnRunner(repeat_aclnn_runner.h),封装 CANNaclnnop/aclnn_repeat.h中的两段式接口,两个函数指针均通过运行时动态加载获得:

Status RepeatAclnnRunner::LoadMethod() { ... Status status = LoadFromSharedObjectFile("aclnnRepeatGetWorkspaceSize", "aclnnRepeat", RepeatAclnnRunner::aclnnGetWorkspaceSizeFunc_, RepeatAclnnRunner::aclnnExecuteFunc_); return status; }

(repeat_aclnn_runner.cpp#L114-L125)动态加载方式使 ATB 对 CANN 版本的 ACLNN 库不产生编译期硬依赖。三个核心钩子:

  • BuildAclnnVariantPack:把 ATB 的输入/输出 Tensor 逐个包装为AclNNTensor,通过CallAclCreateTensor创建 ACL 张量,并记录tensorIdxneedUpdateTensorDataPtr(repeat_aclnn_runner.cpp#L28-L68);
  • SetAclNNWorkspaceExecutor:先用aclCreateIntArrayparam_.multiples转成aclIntArray *size_,再调用aclnnRepeatGetWorkspaceSize(self, size_, out, &workspaceBufferSize, &rawExecutorPtr)完成第一段:算 workspace 大小并生成aclOpExecutor,随后包装为atbAclOpExecutor并查询是否可复用(IsRepeatable())(repeat_aclnn_runner.cpp#L70-L93);
  • LaunchAclnnKernel:从 Context 取执行 stream,调用aclnnRepeat(workspace, workspaceSize, executor, stream)下发内核(repeat_aclnn_runner.cpp#L95-L112)。

8. 原生 Ops Runner 路径:映射到 ExpandOperation 的视图展开

非 950 平台走RepeatOpsRunner(repeat_ops_runner.h),其核心思路是把 repeat 语义改写为“输入视图变形 + 一次 Expand”的原生算子调用

构造函数中搭建单节点内核图,opDesc 直接指向系统ExpandOperation(repeat_ops_runner.cpp#L29-L42):

repeatNode.opDesc = {0, "ExpandOperation", {}}; ... repeatNode.inferShapePreFunc = this { launchParam.SetParam(AsdOps::OpParam::Expand({repeatParam_})); };

真正的技巧在InTensorViewFunc(repeat_ops_runner.cpp#L44-L76)。它在不拷贝数据的前提下重排输入维度,为每个需要放大的维度插入1维再放原维,例如输入[a, b]且某维倍数为k > 1时,视图变为[a, 1, k, b];随后逐维累积repeatParam_[i] *= newDims[i]得到 Expand 的目标形状,并在累积前做int64_t上溢检查。这样原生 Expand 一次视图扩展就完成了整段 repeat,避免了显式的分块复制逻辑。

文件末尾通过REG_RUNNER_TYPE(RepeatOpsRunner)REG_OP_PARAM(AsdOps::OpParam::Expand)完成 Runner 类型与参数类型的静态注册(repeat_ops_runner.cpp#L80-L81)。

9. 测试验证:四组用例覆盖典型形态

算子级测试位于 tests/apitest/opstest/python/operations/repeat/test_repeat.py,参数矩阵数据在 tests/apitest/opstest/csv/repeat.csv。四个用例分别覆盖:

用例输入 shapemultiples验证点
TestRepeatOperation1[2,3,5]fp16[1,4,4]同维数下两维扩展
TestRepeatOperation2[1,2,32,1]fp16[1,1,1,2]单位数维上的复制
TestRepeatOperation3[256,2,32,1]fp16[1,1,1,2]大批量输入
TestRepeatOperation4[256,2,1,1,128]fp16[1,1,1,16,1]5 维输入且含 format cast

每个用例的 golden 直接以in_tensors[0].repeat(multiples)作为参考实现,即 Repeat 语义与 PyTorchtorch.repeat严格对齐;TestRepeatOperation4还通过torch_npu.npu_format_cast覆盖了非 ND 格式输入的场景。

10. 推荐阅读顺序(来自路由文件)

若要复现本文的源码阅读路径,可直接沿用路由文件 routing/repeat.md 给出的顺序:

顺序文件重点关注
1repeat_operation.h了解输入输出数量、InferShape 签名
2repeat_operation.cppCreateRunner() 决策逻辑
3repeat_aclnn_runner.hACLNN API 封装接口
4repeat_aclnn_runner.cppWorkspace 计算 + ACLNN API 调用
5repeat_ops_runner.h原生 Ops 执行接口
6repeat_ops_runner.cpp原生 Ops 调用链 + 平台适配

11. 小结

Repeat 算子在 ascend-transformer-boost 中的实现呈现出典型的 ATB 单算子结构:RepeatOperation负责参数校验、广播形状推导与 8 维上限检查;CreateRunner按平台在两条执行路径间切换——950 平台经运行时加载的aclnnRepeat/aclnnRepeatGetWorkspaceSize两段式接口执行,其余平台则将 repeat 改写为“输入视图变形 + 原生 ExpandOperation”的单节点内核图。两条路径都内置int64_t溢出保护,且语义与torch.repeat完全对齐,测试用例可直接复用 PyTorch 作为 golden。

【免费下载链接】ascend-transformer-boost本项目是CANN提供的是一款高效、可靠的Transformer加速库,基于华为Ascend AI处理器,提供Transformer定制化场景的高性能融合算子。项目地址: https://gitcode.com/cann/ascend-transformer-boost

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询