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):
| # | 文件 | 角色 |
|---|---|---|
| 1 | repeat_aclnn_runner.cpp | ACLNN Runner |
| 2 | repeat_aclnn_runner.h | ACLNN Runner |
| 3 | repeat_operation.cpp | Operation 定义 |
| 4 | repeat_operation.h | Operation 定义 |
| 5 | repeat_ops_runner.cpp | Ops Runner |
| 6 | repeat_ops_runner.h | Ops 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 = 1、OUT_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),它做了三件事:
- 平台预加载检查:若当前平台为
ASCEND_950,会提前调用RepeatAclnnRunner::LoadMethod()确认 ACLNN 函数符号可加载,失败则提示检查 CANN 版本并返回ERROR_CANN_ERROR; multiples维度校验:multiples.size() > MAX_DIM(8)时报错;- 逐元素正数校验:任一
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 阶段各执行一次:
- 输入维数不得超过 multiples 长度,否则返回
ERROR_INVALID_TENSOR_DIM; - 统计真正发生放大的维度数
repeatDimNum:仅当multiples[i] > 1且in_dims[i] > 1时才计数(对长度为 1 的维做复制不视为“扩展维度”); - 总量约束:
repeatDimNum + inTensorDimNum <= MAX_DIM(8)且repeatDimNum + multiples.size() <= 8; - 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 张量,并记录tensorIdx与needUpdateTensorDataPtr(repeat_aclnn_runner.cpp#L28-L68);SetAclNNWorkspaceExecutor:先用aclCreateIntArray将param_.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。四个用例分别覆盖:
| 用例 | 输入 shape | multiples | 验证点 |
|---|---|---|---|
| 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 给出的顺序:
| 顺序 | 文件 | 重点关注 |
|---|---|---|
| 1 | repeat_operation.h | 了解输入输出数量、InferShape 签名 |
| 2 | repeat_operation.cpp | CreateRunner() 决策逻辑 |
| 3 | repeat_aclnn_runner.h | ACLNN API 封装接口 |
| 4 | repeat_aclnn_runner.cpp | Workspace 计算 + ACLNN API 调用 |
| 5 | repeat_ops_runner.h | 原生 Ops 执行接口 |
| 6 | repeat_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),仅供参考