CANN opbase 算子开发:INFER_SHAPE 宏用法详解与输出 Shape 推导实战
【免费下载链接】opbase本项目是CANN算子库的基础框架库,为算子提供公共依赖文件和基础调度能力。项目地址: https://gitcode.com/cann/opbase
导读:本文围绕 CANN opbase 算子库中用于推导算子输出 Shape 的核心宏INFER_SHAPE,系统讲解其宏功能、宏原型、参数语义、与ADD_TO_LAUNCHER_LIST_AICORE等宏的调用顺序约束,并结合仓库源码(make_op_executor.h、op_arg_def.h、op_executor.cpp)与单元测试用例,还原宏的展开与底层调用链。读者阅读后可掌握在 aclnn 二阶段算子接口中正确使用INFER_SHAPE完成输出 Shape 推导的完整方法。
宏功能:在算子执行前推导输出 Shape
在 CANN opbase 的 aclnn 算子开发体系中,算子对外暴露的接口通常是"一阶段(aclnn_Xxx)+ 二阶段(aclnn_Xxx_)"的组合。其中一阶段接口负责参数校验、Shape 推导与执行器的准备,二阶段接口负责真正把算子任务下发执行。INFER_SHAPE宏正是位于一阶段接口中、用来触发 Shape 推导的关键工具:
针对指定算子,运行其 InferShape 函数,推导输出 shape。
也就是说,只要某个算子在算子原型中注册了 InferShape 实现,开发者在接口实现中调用一次INFER_SHAPE(KERNEL_NAME, ...),即可借助算子的 InferShape 逻辑,根据传入的输入张量与属性参数自动推导出输出张量的 Shape,并将推导结果回填到输出aclTensor中。这一过程不启动内核、不分配设备内存,属于"轻量级"的形状推断阶段。
在 opbase 中,InferShape 能力不止存在于 aclnn 侧。算子侧(host 侧)的gert::InferShapeContext同样被大量复用,例如 infershape_broadcast_util.h、infershape_elewise_util.h、infershape_reduce_util.h 中提供的InferShape4Broadcast、InferShape4Elewise、InferShape4Reduce等工具函数,就是面向广播、逐元素、归约三类典型算子形态的通用 Shape 推导实现。INFER_SHAPE宏与这些实现配合,构成了 opbase 中从"算子原型注册"到"接口侧形状推导"的完整闭环。
宏原型与参数说明
INFER_SHAPE的宏原型如下:
INFER_SHAPE(KERNEL_NAME, op_args...)| 参数 | 输入/输出 | 说明 |
|---|---|---|
| KERNEL_NAME | 输入 | 算子名,例如Add。宏内部会基于它拼出算子类型 ID 符号(KERNEL_NAME##OpTypeId()),因此必须与OP_TYPE_REGISTER注册的算子名保持一致。 |
| op_args... | 输入 | 算子的参数,包括输入 OP_INPUT、输出 OP_OUTPUT、属性 OP_ATTR 等参数。 |
值得强调的是,op_args...是一个可变参数包,其中每一项都必须是以OP_INPUT、OP_OUTPUT、OP_ATTR等"参数封装宏"为单位的参数组,而不是裸的aclTensor*指针。这些封装宏统一把实参打包成带有类型标签的元组对象,INFER_SHAPE宏内部才能据此对输入、输出、属性进行分类处理。opbase 中这类参数封装宏的定义集中在 op_arg_def.h:
#define OP_INPUT(x...) op::OpInput(std::make_tuple(x)) #define OP_OUTPUT(x...) op::OpOutput(std::make_tuple(x)) #define OP_ATTR(x...) op::OpAttr(std::make_tuple(x)) #define OP_WORKSPACE(x...) op::OpWorkspace(std::make_tuple(x)) #define OP_OUTSHAPE(x...) op::OpOutshape(std::tuple<aclTensor*, uint64_t>(x)) #define OP_OPTION(x...) op::OpOption(std::make_tuple(x)) #define OP_EMPTY_ARG op::EMPTY_OP_ARG #define OP_MODE(x...) op::OpMode(std::make_tuple(x))对应的参数分类依据是 op_arg_def.h 中定义的OpArgDef枚举:OP_INPUT_ARG = 0、OP_OUTPUT_ARG = 1、OP_ATTR_ARG = 2、OP_WORKSPACE_ARG = 3、OP_OUTSHAPE_ARG = 4、OP_OPTION_ARG = 5、OP_EXEC_MODE_ARG = 6等。INFER_SHAPE宏正是按照这一分类,从上下文对象中取出输入、输出、属性三类参数列表,再交给底层的InferShape函数。
关联接口一览
原文档中明确指出,以下接口是INFER_SHAPE宏定义内部会调用到的关联接口:
OP_INPUT(x...) OP_OUTPUT(x...) OP_ATTR(x...) OP_WORKSPACE(x...) OP_OUTSHAPE(x...) OP_OPTION(x...) OP_EMPTY_ARG OP_MODE(x...)各关联接口的职责如下:
- OP_INPUT:封装算子的输入
aclTensor与aclTensorList。注意:若算子存在非aclTensor/aclTensorList的输入,需先通过aclOpExecutor::ConvertToTensor转换为aclTensor后再传入。 - OP_OUTPUT:封装算子的输出
aclTensor与aclTensorList。Shape 推导的结果会写回这些输出张量对象中。 - OP_ATTR:封装算子的属性参数(即算子原型中声明的属性),例如
adjX1、adjX2等;属性的类型可覆盖布尔、整型、浮点、字符串、DataType、OpImplMode、aclScalar以及各类数组(详见 op_arg_def.h 的OpArgType枚举)。 - OP_WORKSPACE:封装算子执行所需的 workspace 张量列表,用于在内核执行前申请临时工作空间。
- OP_OUTSHAPE:以
std::tuple<aclTensor*, uint64_t>形式封装输出 Shape 相关信息,用于携带"带确定 Shape 的输出"场景。 - OP_OPTION:封装算子执行选项(如实现模式
OpImplMode)。 - OP_EMPTY_ARG:表示空参数占位(
op::EMPTY_OP_ARG),用于参数位置对齐。 - OP_MODE:封装算子执行模式(对应
OpExecMode)。
这些宏虽各有用途,但在INFER_SHAPE的上下文里,真正参与 Shape 推导的核心是OP_INPUT(推导的输入依据)、OP_OUTPUT(推导结果的承载对象)与OP_ATTR(影响推导结果的属性参数),其余接口更多是作为同一套参数描述体系中的可选成员存在。
源码级原理:宏展开与底层调用链
宏定义的展开过程
INFER_SHAPE宏的实际定义位于 make_op_executor.h,展开后的逻辑可以用如下伪代码概括:
#define INFER_SHAPE(KERNEL_NAME, op_args...) \ ({ \ aclnnStatus inferShapeRet; \ do { \ op::OpArgContext* opArgCtx = GetOpArgContext(op_args); \ if (opArgCtx == nullptr) { \ inferShapeRet = ACLNN_ERR_PARAM_NULLPTR; \ } else { \ inferShapeRet = InferShape(KERNEL_NAME##OpTypeId(), \ *opArgCtx->GetOpArg(op::OP_INPUT_ARG), \ *opArgCtx->GetOpArg(op::OP_OUTPUT_ARG), \ *opArgCtx->GetOpArg(op::OP_ATTR_ARG)); \ op::DestroyOpArgContext(opArgCtx); \ } \ } while (0); \ inferShapeRet; \ })其核心执行步骤为:
- 构建参数上下文:调用
GetOpArgContext(op_args)(内部通过MakeOpArgContext分配并初始化OpArgContext,见 op_arg_def.h)。OpArgContext内部以std::array<OpArgList, OP_ARG_TYPE_NUM>按参数类别存放输入、输出、属性等参数列表。 - 空指针保护:若上下文构建失败(例如内存分配失败),宏返回
ACLNN_ERR_PARAM_NULLPTR,避免空指针解引用。 - 调用底层 InferShape:以
KERNEL_NAME##OpTypeId()拼出算子类型 ID,并将上下文中的输入参数列表(OP_INPUT_ARG)、输出参数列表(OP_OUTPUT_ARG)、属性参数列表(OP_ATTR_ARG)分别取出,调用InferShape(optype, inputs, outputs, attrs)。 - 资源释放:推导完成后调用
DestroyOpArgContext(opArgCtx)释放上下文内存,防止泄漏。 - 返回状态码:宏整体以
aclnnStatus返回,可作为一阶段接口的返回值。
这里有一个值得注意的细节:KERNEL_NAME##OpTypeId()这个符号名要求算子已通过OP_TYPE_REGISTER或类似机制注册过算子类型 ID。以测试文件 test_infer_shape.cpp 为例,使用INFER_SHAPE前必须先写OP_TYPE_REGISTER(Add);、OP_TYPE_REGISTER(ReduceSum);等注册语句,宏才能通过AddOpTypeId()拿到对应的类型 ID。
底层 InferShape 接口
INFER_SHAPE宏最终调用的是 op_executor.h 中声明的接口:
aclnnStatus InferShape(uint32_t optype, op::OpArgList& inputs, op::OpArgList& outputs, op::OpArgList& attrs);该接口在 op_executor.cpp 中实现,直接转发到内部实现op::internal::InferShape(optype, inputs, outputs, attrs)。内部实现会基于算子类型 ID 查找到该算子的 InferShape 函数指针,加载输入 Shape 与属性,执行推导并把结果写回输出参数。
在推导上下文的准备上,opbase 提供了InferShapeContextHolder(见 infershape_context_holder.cpp),它负责:
BuildInferShapeContext():一次性为KernelRunContext与AsyncAnyValue值数组分配内存;EnsureContextCapacity():当输入/输出数量增长时,以 2 倍增长因子扩容上下文缓冲区(重新分配并memcpy旧数据,不使用realloc,以保证缓冲区尾部清零语义);UpdateInferShapeContext():把内核上下文中的输入参数、输出参数与推导值挂接到推导上下文,正确设置input_size、output_size、output_start等字段,为 InferShape 函数执行做好准备。
由此可见,INFER_SHAPE宏虽然在接口代码中只是一行调用,但其背后串联了参数上下文构建、算子类型 ID 解析、InferShape 函数查表、推导上下文构建与资源生命周期管理等多个环节,是 aclnn 一阶段接口中"形状推导"子系统的统一入口。
约束说明:与 ADD_TO_LAUNCHER_LIST_AICORE 的调用顺序
使用INFER_SHAPE有一条关键约束:
如果算子需要
INFER_SHAPE,那么此宏需要在ADD_TO_LAUNCHER_LIST_AICORE之前调用。
也就是说,在 aclnn 一阶段接口中,代码顺序应为:
// 1. 先推导输出 Shape auto ret = INFER_SHAPE(KERNEL_NAME, op_args...); // 2. 再创建 AI Core 算子执行任务并加入执行队列 ADD_TO_LAUNCHER_LIST_AICORE(KERNEL_NAME, op_args...);这一顺序要求与两个宏的职责划分是一致的:
INFER_SHAPE只做"形状推导":它根据输入与属性计算输出 Shape,不涉及内核二进制、tiling、workspace 等执行期信息;ADD_TO_LAUNCHER_LIST_AICORE负责"创建执行任务":它会构建AiCoreKernelLauncher并调用BuildGraph组装算子执行图(见 op_executor.cpp),这个过程依赖已经推导完成的输出 Shape 来构造内核参数与图结构。
反过来看,ADD_TO_LAUNCHER_LIST_AICORE 文档中也明确写着"如果算子需要 INFER_SHAPE,那么此宏需要在 INFER_SHAPE 之后调用",两份文档互相印证。若顺序颠倒,执行任务创建时拿到的输出 Shape 还是未推导的初始状态,将导致后续 tiling 与内核下发阶段得到错误的形状信息。
调用示例:BatchMatMulV3 输出 Shape 推导
原文档给出的示例以 BatchMatMulV3 算子为对象,推导其输出 Shape:
// 调用INFER_SHAPE推导batchmatmul算子的输出shape,其中BatchMatMulV3是算子的名字, // OP_INPUT是算子输入参数,OP_OUTPUT是算子输出参数,OP_ATTR是算子的属性参数 INFER_SHAPE(BatchMatMulV3, OP_INPUT(x1, x2, bias, nullptr), OP_OUTPUT(bmmOut), OP_ATTR(adjX1, adjX2, offsetX, opImplModeEnum));逐段拆解该示例:
BatchMatMulV3:算子名。宏内部将其拼成BatchMatMulV3OpTypeId(),因此该算子必须已完成类型注册。OP_INPUT(x1, x2, bias, nullptr):封装 4 个输入,其中x1、x2是参与矩阵乘的两个张量,bias是偏置张量,nullptr表示该位置无输入占位。从 op_arg_def.h 可以看到,nullptr会被编码为OPARG_ACLTENSOR类型的空张量,从而保持参数位置与算子原型对齐,不影响 Shape 推导逻辑。OP_OUTPUT(bmmOut):封装 1 个输出张量bmmOut,推导结果会写回该张量的 Shape 字段。OP_ATTR(adjX1, adjX2, offsetX, opImplModeEnum):封装 4 个属性,其中adjX1、adjX2是控制左右输入是否转置的布尔属性,offsetX是矩阵乘偏移量属性,opImplModeEnum是实现模式枚举。这些属性会作为 InferShape 的输入参与推导:例如adjX1为 true 时,输出 Shape 的对应维度的计算依据会从x1的转置形状得出。
调用后返回的aclnnStatus可直接作为一阶段接口的返回值。若返回ACLNN_SUCCESS,说明推导成功,bmmOut中已填充推导出的 Shape;若返回ACLNN_ERR_PARAM_NULLPTR,说明参数上下文构建失败(如空指针入参)。
更多形态示例
为便于对照,仓库测试 test_infer_shape.cpp 中给出了不同算子形态的典型用法(其中算子通过IMPL_OP(...).InferShape(...)注册推导函数,接口侧再以INFER_SHAPE触发):
// 逐元素算子:输出 Shape 直接复制输入 Shape // IMPL_OP(Add).InferShape(...); // 推导逻辑:*output = *input_shape auto ret = INFER_SHAPE(Add, OP_INPUT(self.get(), other.get()), OP_OUTPUT(out.get())); EXPECT_EQ(ret, ACL_SUCCESS); EXPECT_EQ(out->GetOriginalShape(), otherShape); EXPECT_EQ(out->GetStorageShape(), otherShape); EXPECT_EQ(out->GetViewShape(), otherShape);// 归约算子:携带属性参与推导 // IMPL_OP(ReduceSum).InferShape(...); auto ret = INFER_SHAPE(ReduceSum, OP_INPUT(x.get(), rAxesTensor.get()), OP_OUTPUT(out.get()), OP_ATTR(false));这些用例同时给出了INFER_SHAPE的验证方式:推导完成后,输出张量的GetOriginalShape()(原始 Shape)、GetStorageShape()(存储 Shape)、GetViewShape()(视图 Shape)应与预期一致。这也从侧面说明,INFER_SHAPE完成的是包括原始形状、存储形状、视图形状在内的完整 Shape 三元组推导,而非仅仅填充一个维度列表。
实战建议
- 先注册后推导:使用
INFER_SHAPE前,务必确保算子已通过OP_TYPE_REGISTER(KERNEL_NAME)完成类型 ID 注册,否则KERNEL_NAME##OpTypeId()符号无法解析。 - 顺序不可颠倒:在 aclnn 一阶段接口中,
INFER_SHAPE必须位于ADD_TO_LAUNCHER_LIST_AICORE之前。 - 参数分类必须规范:输入、输出、属性必须分别通过
OP_INPUT、OP_OUTPUT、OP_ATTR封装,不要混用或裸传指针;缺参位置可用nullptr(输入/输出)占位。 - 检查返回值:
INFER_SHAPE的返回值是一阶段接口状态的一部分,建议与后续ADD_TO_LAUNCHER_LIST_AICORE的返回值一并检查,任一失败都应提前返回,避免进入执行阶段。 - 非张量参数先行转换:若算子的输入或输出中包含非
aclTensor/aclTensorList类型,需先调用aclOpExecutor::ConvertToTensor完成转换,再交给OP_INPUT/OP_OUTPUT封装。
总结
INFER_SHAPE是 CANN opbase 中连接"算子 InferShape 实现"与"aclnn 接口调用"的桥梁宏:它统一了算子参数的描述方式(OP_INPUT/OP_OUTPUT/OP_ATTR等),把推导动作封装为一行可读性极高的调用,并妥善处理了空指针防护与资源释放。理解它的宏展开过程(make_op_executor.h)、底层接口(op_executor.h)与调用顺序约束(先INFER_SHAPE后ADD_TO_LAUNCHER_LIST_AICORE),是编写正确、健壮的 aclnn 算子接口代码的基础。相关配套文档可在 common_macros_and_classes.md 索引页中找到完整宏与类的说明。
【免费下载链接】opbase本项目是CANN算子库的基础框架库,为算子提供公共依赖文件和基础调度能力。项目地址: https://gitcode.com/cann/opbase
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考