ARTICLE DETAIL

资讯详情

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

CANN opbase 算子开发指南:aclTensor::SetData 接口原理与实操

CANN opbase 算子开发指南:aclTensor::SetData 接口原理与实操 CANN opbase 算子开发指南aclTensor::SetData 接口原理与实操【免费下载链接】opbase本项目是CANN算子库的基础框架库为算子提供公共依赖文件和基础调度能力。项目地址: https://gitcode.com/cann/opbase导读SetData是 CANN opbase 算子库中aclTensor类提供的 host 侧数据写入接口用于向通过AllocHostTensor申请得到的 host 侧 tensor 写入数据支持“按索引写单个元素”和“用一段已有内存批量初始化”两种形态。本文围绕 SetData 接口文档 展开先讲清接口定位与两种重载的原型、参数与调用示例再深入 common_types.cpp 源码剖析其底层实现原理数据类型转换、布尔特判、placement 校验、支持的 DT 枚举并给出边界约束与工程实践建议。读完本文你将掌握如何在算子开发中正确、高效地使用SetData填充 host 侧 tensor 数据并能基于源码证据判断其适用场景与限制。接口定位谁在调用 SetDataSetData是aclTensor的成员方法。在 CANN opbase 中aclTensor是算子开发侧表示张量的核心对象声明于 common_types.h与GetStorageShape/SetStorageShape/SetDataType等接口并列属于 opdev 的 common_types 一组。整个接口族的使用概览可参见 common_types.md。其使用前提是“通过AllocHostTensor申请得到的 host 侧 tensor”即 tensor 的内存位于 hostCPU 侧而非 device昇腾设备侧。AllocHostTensor系列接口声明于 op_executor.h由aclOpExecutor提供支持int64_t、uint64_t、bool、char、int32_t、uint32_t、int16_t、uint16_t、int8_t、uint8_t、double、float、fp16_t、bfloat16等多种基础类型指针输入其实现位于 op_executor.cpp每个重载内部都通过new aclTensor(...)创建对象并把 tensor 登记进allocatedObjList_与allocatedTensorList_统一管理生命周期。从源码看AllocHostTensor创建的 host tensor 有两种来源形态仅指定 shape 与数据类型对应 common_types.cpp 第 164-190 行 的构造函数内部按storageShape的元素个数乘以TypeSize(dataType)分配内存并将aclStorage的地址指向该内存TensorPlacement标记为kOnHost直接传入已有数据指针value, size, dataType形态将指针地址作为存储地址。SetData针对的正是这类placement 为kOnHost的 tensor——这一前提在实现中被显式检查详见下文源码剖析。函数原型与参数说明SetData提供两个模板重载原型如下// 重载一针对 tensor设置指定索引处的值 template typename T void SetData(int64_t index, const T value, op::DataType dataType); // 重载二针对 tensor用一块已有内存初始化 tensor 数据 template typename T void SetData(const T *value, uint64_t size, op::DataType dataType);两个接口均为模板函数模板参数T由调用方传入的value实参推导因此value可以是任意数值基础类型含op::fp16_t、op::bfloat16等 opbase 自定义类型。重载一按索引写入单个元素参数输入/输出说明index输入需要修改 aclTensor 的第几个元素从 0 开始的元素序号。value输入将 aclTensor 的指定元素修改为 value 的值。dataType输入目标数据类型op::DataType即ge::DataType。value 会被转换为该指定的 dataType 后再写入 aclTensor。重载二用已有内存初始化 tensor参数输入/输出说明value输入指向需要写入 aclTensor 的数据内存指针。size输入需要写入的元素个数注意是元素个数不是字节数。dataType输入目标数据类型op::DataType即ge::DataType。数据会被转换为该指定的 dataType 后再写入 aclTensor。说明op::DataType即ge::DataType的别名二者是同一枚举。ge::DataType的完整取值说明参见《基础数据结构和接口参考》中“ge 命名空间 DataType”一节。返回值与约束返回值无void。写入失败不会通过返回值上报而是依赖日志OP_LOGE输出错误信息。约束入参指针不能为空。对于重载二value指针必须指向有效的内存对于重载一虽然传入的是值而非指针但同样要求 tensor 本身是有效的 host 侧 tensor。调用示例官方文档给出的完整示例// 初始化一块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); }代码解读input-SetData(myArray, 10, DT_INT64)把myArray指向的 10 个int64_t元素逐个写入input的第 09 号元素input-SetData(10, myArray[0], DT_INT64)把myArray[0]的值写入input的第 10 号元素。注意示例中const aclTensor *input是 const 指针而SetData是 const 成员函数声明与定义中均不带 mutable 限定之外的特殊标记因此 const 对象也可调用。一个更贴近实际开发流程的完整用法——先申请 host tensor再填充数据// 假设 executor 为已创建的 aclOpExecutor 对象 op::Shape shape({5}); aclTensor* hostTensor executor.AllocHostTensor(shape, op::DataType::DT_INT64, op::Format::FORMAT_ND); int64_t values[5] {10, 20, 30, 40, 50}; hostTensor-SetData(values, 5, op::DataType::DT_INT64); // 批量写入 hostTensor-SetData(3, 999, op::DataType::DT_INT64); // 再单独改写第 3 个元素源码级原理剖析SetData的实现在 common_types.cpp 中。以下按实现层次逐步拆解。1. Placement 前置检查两个重载的入口处都先检查this-GetPlacement() op::TensorPlacement::kOnHosttemplate typename T void aclTensor::SetData(int64_t index, const T value, op::DataType dataType) { if (this-GetPlacement() op::TensorPlacement::kOnHost) { void* dataAddr this-GetStorageAddr(); ... } } template typename T void aclTensor::SetData(const T* value, uint64_t size, op::DataType dataType) { if (this-GetPlacement() op::TensorPlacement::kOnHost) { for (uint64_t i 0; i size; i) { SetData(i, value[i], dataType); } } }这说明SetData只对 host 侧 tensor 生效与文档“针对通过AllocHostTensor申请得到的 host 侧 tensor”的定位完全一致。若对 device 侧 tensorplacement 为kOnDeviceHbm调用方法体不会执行任何写入操作也不会报错——这是使用时必须注意的静默行为。2. 重载二基于重载一逐元素实现重载二指针批量版没有做批量内存拷贝而是简单循环逐元素调用重载一for (uint64_t i 0; i size; i) { SetData(i, value[i], dataType); }因此两种调用的数据转换逻辑完全一致每个元素都会独立经历一次类型转换后再写入。3. 按 dataType 分发到具体写入分支重载一内部根据dataType枚举做 switch 分发将value以模板参数dataType目标类型写入dataAddr index偏移处。源码支持的目标类型包括浮点类DT_FLOATfloat、DT_FLOAT16op::fp16_t、DT_BF16op::bfloat16、DT_DOUBLEdouble整数类DT_INT8/DT_INT16/DT_INT32/DT_INT64、DT_UINT8/DT_UINT16/DT_UINT32/DT_UINT64布尔类DT_BOOL走专门的SetDataByBool逻辑。对于DT_BOOL之外的类型统一调用SetDataByDataTypeT, dataType完成写入template typename T, typename dataType static void SetDataByDataType(int64_t index, void* dataAddr, const T value) { dataType* tmpDataAddr static_castdataType*(dataAddr); if constexpr (op::internal::IsCustomFloattypename std::decayT::type::value) { // For custom float types, convert through double to avoid ambiguity *(tmpDataAddr index) static_castdataType(static_castdouble(value)); } else { *(tmpDataAddr index) static_castdataType(value); } }关键点采用static_castdataType(value)做C 显式类型转换即窄化转换同样被允许因此float - int、int64_t - int8_t这类可能损失精度的转换不会在编译期被拒绝当模板实参T是 opbase 的自定义浮点类型op::fp16_t、bfloat16、Float8E5M2等通过IsCustomFloat判断时会先转为double再转目标类型以避免多重重载/构造歧义指针以dataAddr存储区起始地址为基址写入位置是dataAddr index * sizeof(dataType)即按元素序号索引。4. 布尔类型的特殊语义当目标类型为DT_BOOL时走的是独立的SetDataByBool其语义与普通static_castbool不同template typename T static void SetDataByBool(int64_t index, void* dataAddr, const T value) { auto tmpDataAddr static_castbool*(dataAddr); if constexpr (std::is_floating_pointT::value || 是自定义浮点类型...) { *(tmpDataAddr index) std::abs(static_castfloat(value)) std::numeric_limitsfloat::epsilon(); } else { *(tmpDataAddr index) static_castbool(value); } }即当写入源是浮点类型含fp16_t、Float8*、Float6*、Float4*、HiFloat4/HiFloat8等自定义浮点时取绝对值是否不小于 float 最小非零值epsilon作为布尔判定而当写入源是整数/其他类型时走普通static_castbool。因此浮点 0 写入 bool 得到false任何非零值得到true。5. 不支持数据类型的兜底switch 的default分支对应源码第 645-652 行会通过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)]。从该列表可见SetData目前不支持复数类型DT_COMPLEX64/DT_COMPLEX128与 Float8/Float6/Float4/HiFloat 等自定义低比特浮点作为目标 dataType——即使源值类型T是自定义浮点目标dataType也只接受上述 13 种。6. 相关便捷接口在SetData之上aclTensor还封装了一批类型化便捷接口common_types.h它们内部均委托给SetDataSetBoolData(const bool* value, uint64_t size, op::DataType dataType)实现见 common_types.cpp 第 676 行SetIntData(const int64_t* value, uint64_t size, op::DataType dataType)SetFloatData(const float* value, uint64_t size, op::DataType dataType)SetFp16Data(const op::fp16_t* value, uint64_t size, op::DataType dataType)SetBf16Data(const op::bfloat16* value, uint64_t size, op::DataType dataType)以及SetFloat8E5M2Data/SetFloat8E4M3FNData/SetFloat8E8M0Data/SetFloat6E3M2Data/SetFloat6E2M3Data/SetFloat4E2M1Data/SetFloat4E1M2Data/SetHiFloat4Data/SetHiFloat8Data等。这些接口只是把T类型固定为对应基础类型后转发到SetData便于调用方写出类型更明确、更不易出错的代码。测试用例验证仓库单元测试 test_common_types.cpp 中提供了SetData的直接用例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); }该用例展示了SetData的典型调用形态先构造一个 host 侧aclTensor以float数据初始化再调用批量重载写入 5 个float元素。同文件中还有大量围绕aclTensor的 Get/Set 接口族测试可作为阅读 common_types 模块整体行为的入口。边界约束与工程实践建议综合文档与源码SetData的使用应遵循以下约束与建议仅限 host 侧 tensor只对AllocHostTensor申请的 tensorplacement 为kOnHost生效对 device 侧 tensor 调用会被静默忽略。index/size 越界风险实现不做边界检查index超出Numel()、size超过 tensor 容量时会产生越界写。调用前应自行用Numel()view 元素个数确认。指针非空批量重载的value指针不能为空size为 0 时循环不会执行属于安全的空操作。类型转换语义写入过程使用static_cast窄化转换float - int、double - float等可能损失精度DT_BOOL目标对浮点源有基于epsilon的特殊判定。若需要精确控制转换行为建议在调用前完成转换。目标 dataType 范围仅支持DT_FLOAT/DT_FLOAT16/DT_BF16/DT_DOUBLE/DT_INT8/DT_INT16/DT_INT32/DT_INT64/DT_UINT8/DT_UINT16/DT_UINT32/DT_UINT64/DT_BOOL13 种复数与自定义低比特浮点类型不被支持使用前应核对枚举取值。优先使用类型化便捷接口当写入数据的基础类型明确如int64_t、float、bool时优先使用SetIntData/SetFloatData/SetBoolData等接口代码意图更清晰。总结SetData是 CANN opbase 中向 host 侧aclTensor写入数据的基础接口提供按索引单元素写入与按内存批量写入两种形态核心能力在于“写入时按目标dataType自动做类型转换”。通过 common_types.cpp 的源码可以看到其实质是一个带 placement 检查、按枚举分发、逐元素转换写入的模板方法且批量版内部就是单元素版的循环复用。结合AllocHostTensor申请、SetData填充、再交给aclOpExecutor下发的完整链路即可在算子开发中灵活构造与初始化 host 侧输入数据。【免费下载链接】opbase本项目是CANN算子库的基础框架库为算子提供公共依赖文件和基础调度能力。项目地址: https://gitcode.com/cann/opbase创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表