news 2026/9/19 4:57:56

CANN opbase 算子开发:INFER_SHAPE 宏用法详解与输出 Shape 推导实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
CANN opbase 算子开发:INFER_SHAPE 宏用法详解与输出 Shape 推导实战

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 中提供的InferShape4BroadcastInferShape4ElewiseInferShape4Reduce等工具函数,就是面向广播、逐元素、归约三类典型算子形态的通用 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_INPUTOP_OUTPUTOP_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 = 0OP_OUTPUT_ARG = 1OP_ATTR_ARG = 2OP_WORKSPACE_ARG = 3OP_OUTSHAPE_ARG = 4OP_OPTION_ARG = 5OP_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:封装算子的输入aclTensoraclTensorList。注意:若算子存在非aclTensor/aclTensorList的输入,需先通过aclOpExecutor::ConvertToTensor转换为aclTensor后再传入。
  • OP_OUTPUT:封装算子的输出aclTensoraclTensorList。Shape 推导的结果会写回这些输出张量对象中。
  • OP_ATTR:封装算子的属性参数(即算子原型中声明的属性),例如adjX1adjX2等;属性的类型可覆盖布尔、整型、浮点、字符串、DataTypeOpImplModeaclScalar以及各类数组(详见 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; \ })

其核心执行步骤为:

  1. 构建参数上下文:调用GetOpArgContext(op_args)(内部通过MakeOpArgContext分配并初始化OpArgContext,见 op_arg_def.h)。OpArgContext内部以std::array<OpArgList, OP_ARG_TYPE_NUM>按参数类别存放输入、输出、属性等参数列表。
  2. 空指针保护:若上下文构建失败(例如内存分配失败),宏返回ACLNN_ERR_PARAM_NULLPTR,避免空指针解引用。
  3. 调用底层 InferShape:以KERNEL_NAME##OpTypeId()拼出算子类型 ID,并将上下文中的输入参数列表(OP_INPUT_ARG)、输出参数列表(OP_OUTPUT_ARG)、属性参数列表(OP_ATTR_ARG)分别取出,调用InferShape(optype, inputs, outputs, attrs)
  4. 资源释放:推导完成后调用DestroyOpArgContext(opArgCtx)释放上下文内存,防止泄漏。
  5. 返回状态码:宏整体以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():一次性为KernelRunContextAsyncAnyValue值数组分配内存;
  • EnsureContextCapacity():当输入/输出数量增长时,以 2 倍增长因子扩容上下文缓冲区(重新分配并memcpy旧数据,不使用realloc,以保证缓冲区尾部清零语义);
  • UpdateInferShapeContext():把内核上下文中的输入参数、输出参数与推导值挂接到推导上下文,正确设置input_sizeoutput_sizeoutput_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 个输入,其中x1x2是参与矩阵乘的两个张量,bias是偏置张量,nullptr表示该位置无输入占位。从 op_arg_def.h 可以看到,nullptr会被编码为OPARG_ACLTENSOR类型的空张量,从而保持参数位置与算子原型对齐,不影响 Shape 推导逻辑。
  • OP_OUTPUT(bmmOut):封装 1 个输出张量bmmOut,推导结果会写回该张量的 Shape 字段。
  • OP_ATTR(adjX1, adjX2, offsetX, opImplModeEnum):封装 4 个属性,其中adjX1adjX2是控制左右输入是否转置的布尔属性,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_INPUTOP_OUTPUTOP_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_SHAPEADD_TO_LAUNCHER_LIST_AICORE),是编写正确、健壮的 aclnn 算子接口代码的基础。相关配套文档可在 common_macros_and_classes.md 索引页中找到完整宏与类的说明。

【免费下载链接】opbase本项目是CANN算子库的基础框架库,为算子提供公共依赖文件和基础调度能力。项目地址: https://gitcode.com/cann/opbase

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/19 4:56:08

荧光标记胆酸衍生物CY5.5/CY7-甘氨胆酸的应用与特性

1. 荧光标记胆酸衍生物概述CY5.5-Glycocholic Acid&#xff08;CY5.5-甘氨胆酸&#xff09;和CY7-Glycocholic Acid&#xff08;CY7-甘氨胆酸&#xff09;是两种重要的荧光标记胆酸衍生物&#xff0c;在生物医学研究和分子影像领域具有广泛应用。这类化合物通过将近红外荧光染料…

作者头像 李华
网站建设 2026/9/19 4:54:27

psycopg2 参数化建表,把 AsIs 写法交给走 TaoToken 的 Codex 复查

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/19 4:53:54

MuJoCo 木块自由落体:Codex 走 TaoToken 生成并跑通 01_mujoco_helloworld.py

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/19 4:53:12

基于深度学习的鸟类识别检测系统:YOLO实战全解析

1. 项目概述与需求拆解1.1 为什么鸟类识别适合做毕设每年到了毕设选题季&#xff0c;总有学弟学妹来问我&#xff1a;"学长&#xff0c;深度学习方向的题目到底选什么好&#xff1f;"我的回答一直很明确&#xff1a;选一个数据好找、场景直观、算法成熟但又有优化空间…

作者头像 李华
网站建设 2026/9/19 4:52:55

Rust 打造 OpenObserve:替代 Elasticsearch 和 Prometheus 的可观测性实战

1. 为什么我又把日志和指标系统折腾了一遍如果你运维过中等规模的线上环境&#xff0c;大概率经历过这样的场景&#xff1a;Elasticsearch 集群的 JVM 堆内存三天两头告警&#xff0c;Prometheus 的 TSDB 在高峰期写入延迟飙升&#xff0c;Grafana 面板加载慢得让人想砸键盘。更…

作者头像 李华