
CANN opbase 算子开发aclTensor::SetData 接口详解与源码实现【免费下载链接】opbase本项目是CANN算子库的基础框架库为算子提供公共依赖文件和基础调度能力。项目地址: https://gitcode.com/cann/opbase本指南围绕 CANN 算子库基础框架 opbase 中aclTensor::SetData接口展开讲解如何向由AllocHostTensor申请的 host 侧张量写入数据覆盖两种函数原型、参数语义、数据类型转换规则、底层实现与配套专用接口帮助算子开发者快速掌握在自定义算子实现中构造和填充 host 侧aclTensor的完整方法。一、SetData 的功能定位在 opbase 框架中aclTensor是算子实现层与上层调度框架之间的核心张量载体类定义位于 include/nnopbase/opdev/common_types.h它同时携带 shape、format、stride、数据地址等元信息与数据内容。而SetData正是aclTensor对外提供的数据写入接口针对通过AllocHostTensor申请得到的 host 侧 tensor设置指定位置的数据。也就是说SetData并不是一个独立创建张量的接口而是与AllocHostTensor配套使用的数据填充环节。AllocHostTensor负责在 host 侧分配张量内存参见 AllocHostTensor 接口文档SetData则负责把业务数据按目标数据类型写入这块内存。两者配合即可在算子 host 侧构造出携带真实数据的aclTensor用于后续的 shape 推导校验或经aclOpExecutor提交执行。二、函数原型SetData以模板成员函数的形式声明于aclTensor类中声明见 common_types.h提供两种重载设置指定索引处的值针对 tensor 中的第index个元素写入单个值。void SetData(int64_t index, const T value, op::DataType dataType)用一块已有内存初始化 tensor 数据将指针value指向的一段内存整体写入 tensor。void SetData(const T *value, uint64_t size, op::DataType dataType)其中T为模板参数op::DataType即ge::DataType在 common_types.h 中通过using DataType ge::DataType;定义。这意味着调用方可以传入int64_t、float、bool等任意源类型由接口内部完成到目标dataType的转换后落盘。三、参数说明1. 设置指定索引处值的接口参数输入/输出说明index输入需要修改aclTensor的第几个元素从 0 开始计数的元素下标。value输入目标值将aclTensor的指定元素修改为该值。dataType输入数据类型为op::DataType即ge::DataType。value会被转换为指定的dataType后再写入aclTensor。2. 用已有内存初始化 tensor 数据的接口参数输入/输出说明value输入指向需要写入aclTensor的数据内存的指针。size输入需要写入的元素个数注意是元素个数而非字节数。dataType输入数据类型为op::DataType即ge::DataType。数据会被转换为指定的dataType后再写入aclTensor。3. 返回值说明无返回值void。四、底层实现原理SetData的实现位于 src/nnopbase/common/utils/common_types.cpp从源码可以梳理出以下关键机制。1. 仅对 host 侧张量生效两个重载在函数体入口都会先做 placement 检查if (this-GetPlacement() op::TensorPlacement::kOnHost) { ... }GetPlacement()返回tensor_-GetPlacement()见 common_types.cpp即只有张量驻留在 hostkOnHost时SetData才会真正执行写入。这正与文档中针对通过AllocHostTensor申请得到的 host 侧 tensor的定位一致——该接口面向的是 host 端数据准备场景而非设备侧内存。2. 按目标数据类型进行强制类型转换单元素重载的核心是一个基于dataType的分发 switch将value强制转换为目标类型后写入GetStorageAddr()返回的存储地址void* dataAddr this-GetStorageAddr(); switch (dataType) { case op::DataType::DT_FLOAT: SetDataByDataTypeT, float(index, dataAddr, value); break; case op::DataType::DT_FLOAT16: SetDataByDataTypeT, op::fp16_t(index, dataAddr, value); break; case op::DataType::DT_BF16: SetDataByDataTypeT, op::bfloat16(index, dataAddr, value); break; case op::DataType::DT_INT8: SetDataByDataTypeT, int8_t(index, dataAddr, value); break; ... }SetDataByDataTypecommon_types.cpp内部通过static_castdataType(value)完成转换。这里有一个实现细节值得注意对于自定义浮点类型如op::fp16_t、op::bfloat16、op::Float8E5M2等通过op::internal::IsCustomFloat判定为避免模板推导歧义会先经double中转再转换if constexpr (op::internal::IsCustomFloattypename std::decayT::type::value) { *(tmpDataAddr index) static_castdataType(static_castdouble(value)); } else { *(tmpDataAddr index) static_castdataType(value); }3. bool 目标的特殊语义当dataType为DT_BOOL时走的是独立的SetDataByBool分支common_types.cpp。对于浮点类源类型含各类自定义浮点bool 判定规则为*(tmpDataAddr index) std::abs(static_castfloat(value)) std::numeric_limitsfloat::epsilon();即非零绝对值不小于 float epsilon即置 true其余类型则直接static_castbool(value)。这意味着 NaN、Inf 等特殊浮点值转换为 bool 时会得到true与常规static_castbool语义有所区分。4. 批量重载是单元素版本的循环封装指针批量版本并没有单独的内存拷贝逻辑而是逐元素调用单元素重载common_types.cppfor (uint64_t i 0; i size; i) { SetData(i, value[i], dataType); }因此批量版本天然继承了单元素版本的全部行为同样受kOnHost限制、同样按dataType逐元素转换。相应地元素个数size必须不超过张量容量否则将越界写入。5. 不支持的 dataType 处理当dataType不在 switch 支持列表内时会记录一条不支持数据类型的错误日志OP_LOGE_FOR_NOT_SUPPORTED_DATA_TYPE并在日志中列出当前支持的枚举范围[DT_FLOAT(0), DT_FLOAT16(1), DT_INT8(2), DT_INT32(3), DT_UINT8(4), DT_INT16(6), DT_UINT16(7), DT_UINT32(8), DT_INT64(9), DT_UINT64(10), DT_DOUBLE(11), DT_BOOL(12), DT_BF16(27)]由此可从源码确认当前支持的数据类型为DT_FLOAT、DT_FLOAT16、DT_BF16、DT_INT8、DT_INT16、DT_UINT8、DT_UINT16、DT_INT32、DT_UINT32、DT_INT64、DT_UINT64、DT_DOUBLE、DT_BOOL。五、约束说明入参指针不能为空批量重载的value指针不得为nullptr。仅 host 侧张量可用接口对TensorPlacement非kOnHost的张量不产生写入效果源码层面直接跳过。index/size需在张量元素范围内源码未做边界检查越界访问属于未定义行为调用方需自行保证。dataType须为支持列表内类型否则仅记录错误日志不执行写入。六、调用示例以下示例完整展示了SetData两种重载的用法先用一块int64_t内存初始化input的前 10 个元素再把myArray的第一个值写入input的第 11 个元素下标 10。// 初始化一块 int64_t 内存分别将 input 的前 10 个数字置为该内存的内容 // 并将 input 的第 11 个数字置为 myArray 的第一个数字。 void Func(const aclTensor *input) { int64_t myArray[10]; input-SetData(myArray, 10, DT_INT64); input-SetData(10, myArray[0], DT_INT64); }结合AllocHostTensor的完整使用链路可参见 AllocHostTensor 接口文档典型组合是先AllocHostTensor(shape, dataType, format)拿到 host 张量再通过SetData填充内容后参与算子执行。七、与专用类型接口的关系SetData是一个泛型模板入口aclTensor还在 common_types.h 中提供了一系列针对固定源类型的专用重载内部均直接复用SetData(value, size, dataType)实现见 common_types.cppSetBoolData(const bool* value, ...)SetIntData(const int64_t* value, ...)SetFloatData(const float* value, ...)SetFp16Data(const op::fp16_t* value, ...)SetBf16Data(const op::bfloat16* value, ...)SetFloat8E5M2Data/SetFloat8E4M3FNData/SetFloat8E8M0DataSetFloat6E3M2Data/SetFloat6E2M3DataSetFloat4E2M1Data/SetFloat4E1M2DataSetHiFloat4Data/SetHiFloat8Data这些接口的文档位于 common_types 目录如 SetIntData、SetFloatData、SetFp16Data、SetBf16Data、SetBoolData。当源数据就是标准 C 类型或框架自定义浮点类型时使用对应专用接口可以免去模板推导代码意图也更明确当需要把一种源类型统一转换为多种目标dataType时直接使用泛型SetData更为灵活。八、测试验证opbase 在单测中覆盖了SetData的基本行为见 tests/nnopbase/ut/composite_op/test_common_types.cpp 的aclTensorSetData用例TEST_F(CommonTypesTest, aclTensorSetData) { float fpValue 3.2; uint64_t size 1; aclTensor* floatTensor new aclTensor(fpValue, size, op::DataType::DT_FLOAT); float intArr[5] {1., 2., 3., 4., 5.}; floatTensor-SetData(intArr, 5, op::DataType::DT_QINT16); }该用例通过aclTensor的指针构造版本创建 host 张量再以SetData批量写入印证了先构造、后填充的典型使用模式。需要说明的是用例中传入的DT_QINT16并不在SetData的 switch 支持列表内因此该场景实际会落入 default 分支打印不支持日志测试目的更侧重于验证调用路径不崩溃实际业务使用时应选择第四节列出的支持类型。九、总结aclTensor::SetData是 CANN opbase 框架中 host 侧张量数据写入的标准入口核心要点可归纳为定位配合AllocHostTensor使用负责将业务数据按目标数据类型填充进 host 张量两种重载单元素按索引写入、指针批量写入批量版本内部逐元素复用单元素逻辑类型转换按dataType分发并static_cast转换自定义浮点类型经double中转bool 目标采用非零即 true含 NaN/Inf的判定规则生效前提仅对kOnHost张量生效入参指针不可为空index/size需在张量元素范围内dataType需在支持列表内。掌握该接口后即可在算子 host 侧自由构造携带实际数据的aclTensor为后续 shape 校验、图构建与算子执行提供正确的数据输入。【免费下载链接】opbase本项目是CANN算子库的基础框架库为算子提供公共依赖文件和基础调度能力。项目地址: https://gitcode.com/cann/opbase创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考