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),即只有张量驻留在 host(kOnHost)时,SetData才会真正执行写入。这正与文档中"针对通过AllocHostTensor申请得到的 host 侧 tensor"的定位一致——该接口面向的是 host 端数据准备场景,而非设备侧内存。
2. 按目标数据类型进行强制类型转换
单元素重载的核心是一个基于dataType的分发 switch,将value强制转换为目标类型后写入GetStorageAddr()返回的存储地址:
void* dataAddr = this->GetStorageAddr(); switch (dataType) { case op::DataType::DT_FLOAT: SetDataByDataType<T, float>(index, dataAddr, value); break; case op::DataType::DT_FLOAT16: SetDataByDataType<T, op::fp16_t>(index, dataAddr, value); break; case op::DataType::DT_BF16: SetDataByDataType<T, op::bfloat16>(index, dataAddr, value); break; case op::DataType::DT_INT8: SetDataByDataType<T, int8_t>(index, dataAddr, value); break; ... }SetDataByDataType(common_types.cpp)内部通过static_cast<dataType>(value)完成转换。这里有一个实现细节值得注意:对于自定义浮点类型(如op::fp16_t、op::bfloat16、op::Float8E5M2等,通过op::internal::IsCustomFloat判定),为避免模板推导歧义,会先经double中转再转换:
if constexpr (op::internal::IsCustomFloat<typename std::decay<T>::type>::value) { *(tmpDataAddr + index) = static_cast<dataType>(static_cast<double>(value)); } else { *(tmpDataAddr + index) = static_cast<dataType>(value); }3. bool 目标的特殊语义
当dataType为DT_BOOL时,走的是独立的SetDataByBool分支(common_types.cpp)。对于浮点类源类型(含各类自定义浮点),bool 判定规则为:
*(tmpDataAddr + index) = std::abs(static_cast<float>(value)) >= std::numeric_limits<float>::epsilon();即非零(绝对值不小于 float epsilon)即置 true;其余类型则直接static_cast<bool>(value)。这意味着 NaN、Inf 等特殊浮点值转换为 bool 时会得到true,与常规static_cast<bool>语义有所区分。
4. 批量重载是单元素版本的循环封装
指针批量版本并没有单独的内存拷贝逻辑,而是逐元素调用单元素重载(common_types.cpp):
for (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.cpp):
SetBoolData(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),仅供参考