ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

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

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 中提供的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::tupleaclTensor*, 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::tupleaclTensor*, 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::arrayOpArgList, 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),仅供参考
返回列表