多彩编程 多彩编程MZPH · CODE BLOG
ARTICLE DETAIL

文章详情

深耕前端与后端开发技术的一线实战笔记与踩坑复盘。

CANN ops-transformer 数据类型互推导机制详解:Tensor-Tensor 与 Tensor-Scalar 自动类型提升规则

CANN ops-transformer 数据类型互推导机制详解:Tensor-Tensor 与 Tensor-Scalar 自动类型提升规则 CANN ops-transformer 数据类型互推导机制详解Tensor-Tensor 与 Tensor-Scalar 自动类型提升规则【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer本文以 CANN/ops-transformer 算子库的 互推导关系 为核心系统讲解 aclnn 两段式接口在输入数据类型不一致时的自动类型推导Type Promotion机制包括 Tensor 与 Tensor 之间、Tensor 与 Scalar 之间两套完整推导规则表、规律解读与调用示例并结合仓库源码与算子调用流程说明其实际作用。读完本文读者将能够准确预判aclnnAdd、aclnnAdds等算子 API 在混合类型输入下内部统一使用的计算类型理解×不支持组合的边界并在实际算子调用与样例验证中加以运用。一、什么是互推导两段式接口内部的类型统一机制在 CANN/ops-transformer 中基于单算子 API 执行方式调用算子时通常采用两段式接口样式形如aclnnStatus aclxxXxxGetWorkspaceSize(const aclTensor *src, ..., aclTensor *out, ..., uint64_t *workspaceSize, aclOpExecutor **executor); aclnnStatus aclxxXxx(void *workspace, uint64_t workspaceSize, aclOpExecutor *executor, aclrtStream stream);其中aclxx表示算子接口前缀如aclnnXxx表示算子类型如 Add 算子。完整的调用约定参见 两段式接口。当 API输入的aclTensor数据类型不一致例如一个为ACL_FLOAT16、一个为ACL_FLOAT时API 内部会推导出一个数据类型将输入数据统一转换成该类型后再进行计算——这个过程在项目文档中被称为互推导deduction。推导出的统一类型可以理解为本次计算的计算类型compute type它由参与运算的所有输入共同决定这正是互字的含义不是单方面以某个输入为准而是根据组合规则共同推导。互推导解决的是输入之间类型不一致的问题它与另外两个概念共同构成算子 API 的完整类型/形状语义闭环概念解决的问题对应文档broadcast 关系输入形状不一致时如何广播broadcast关系互推导关系输入数据类型不一致时如何统一计算类型互推导关系本文主题互转换关系输出数据类型与计算类型不一致时如何转换结果互转换关系三者合在一起定义了 aclnn 算子 API 在输入形状、输入类型、输出类型三个维度上的自动处理行为。互推导的推导原理与 PyTorch 的 Type Promotion 机制类似但针对 NPU 算子库的数据类型体系做了适配规则细节以本文两张推导表为准。二、前置基础aclTensor 数据类型与简写约定通过aclCreateTensor接口创建aclTensor时支持的全量数据类型定义于《Runtime 运行时 API》的aclDataType中。项目文档为了方便描述将数据类型统一采用简写形式完整简写对照表参见 数据类型。参与互推导规则的 16 种数据类型及其简写如下原始数据类型简写原始数据类型简写ACL_FLOATf32ACL_UINT32u32ACL_FLOAT16f16ACL_INT64s64ACL_DOUBLEf64ACL_UINT64u64ACL_BF16bf16ACL_BOOLboolACL_INT8s8ACL_COMPLEX32c32ACL_UINT8u8ACL_COMPLEX64c64ACL_INT16s16ACL_COMPLEX128c128ACL_UINT16u16——需要特别说明的是推导规则只覆盖上表中的 16 种类型。aclDataType中还有ACL_STRING、ACL_INT4、ACL_UINT1、ACL_HIFLOAT8、ACL_FLOAT8_E5M2、ACL_FLOAT8_E4M3FN、ACL_FLOAT8_E8M0、ACL_FLOAT6_E3M2、ACL_FLOAT6_E2M3、ACL_FLOAT4_E2M1、ACL_FLOAT4_E1M2等类型完整清单见 数据类型这些类型不适用本文两张推导表具体算子 API 是否支持它们以该 API 参数说明为准。简写不区分大小写例如ACL_FLOAT可写作FLOAT或FLOAT32ACL_BF16可写作BF16或BFLOAT16。三、Tensor-Tensor 互推导规则当一个 API如aclnnAdd、aclnnMul等输入的 aclTensor 数据类型不一致时API 内部会推导出一个数据类型将输入数据转换成该数据类型进行计算。类型推导的规则如下表所示说明为方便描述表格中使用的数据类型是简写形式代表的含义ACL_FLOAT(f32)、ACL_FLOAT16(f16)、ACL_DOUBLE(f64)、ACL_BF16(bf16)、ACL_INT8(s8)、ACL_UINT8(u8)、ACL_INT16(s16)、ACL_UINT16(u16)、ACL_INT32(s32)、ACL_UINT32(u32)、ACL_INT64(s64)、ACL_UINT64(u64)、ACL_BOOL(bool)、ACL_COMPLEX32(c32)、ACL_COMPLEX64(c64)、ACL_COMPLEX128(c128)。表格里表头和最左侧一列分别表示待推导的两个输入数据类型表格中对应位置表示推导出的数据类型。表中叉号×表示这两种类型不能进行推导计算。表1Tensor-Tensor 数据类型推导关系数据类型f32f16f64bf16s8u8s16u16s32u32s64u64boolc32c64c128f32f32f32f64f32f32f32f32×f32×f32×f32c64c64c128f16f32f16f64f32f16f16f16×f16×f16×f16c32c64c128f64f64f64f64f64f64f64f64×f64×f64×f64c128c128c128bf16f32f32f64bf16bf16bf16bf16×bf16×bf16×bf16c32c64c128s8f32f16f64bf16s8s16s16×s32×s64×s8c32c64c128u8f32f16f64bf16s16u8s16×s32×s64×u8c32c64c128s16f32f16f64bf16s16s16s16×s32×s64×s16c32c64c128u16×××××××u16××××××××s32f32f16f64bf16s32s32s32×s32×s64×s32c32c64c128u32×××××××××u32××××××s64f32f16f64bf16s64s64s64×s64×s64×s64c32c64c128u64×××××××××××u64××××boolf32f16f64bf16s8u8s16×s32×s64×boolc32c64c128c32c64c32c128c32c32c32c32×c32×c32×c32c32c64c128c64c64c64c128c64c64c64c64×c64×c64×c64c64c64c128c128c128c128c128c128c128c128c128×c128×c128×c128c128c128c128四、Tensor-Tensor 推导规律解读从表1可以归纳出以下几条规律均为对表格的观察结论可作为调用时预判计算类型的依据浮点类型之间按精度提升wider wins两个浮点输入推导时取精度更高者如f16 f32 → f32、bf16 f32 → f32、f16 f64 → f64。特别的两个半精度f16 bf16 → f32即半精度互相组合时会提升到单精度 f32 参与计算。浮点与整数混合时浮点类型胜出整数输入会被提升为浮点输入的类型如s8 f32 → f32、s64 f16 → f16。有符号与无符号整数混合时向有符号的更宽类型扩展例如s8 u8 → s16、s8 u16不可推导、s32 u8 → s32。u16、u32、u64是孤立类型它们只允许与自身同类型组合u16 u16 → u16与任何其他类型组合均标记为×无法进行推导计算。bool 参与推导且向另一侧靠拢bool f32 → f32、bool f16 → f16、bool s8 → s8、bool u8 → u8、bool bool → bool。bool 不会压过任何数值类型。复数类型具有支配地位任意复数与实数组合结果一定是复数。实数一侧的精度决定复数提升的档位c32 f16 → c32、c32 f32 → c64、c32 f64 → c128整数与 c32 组合得到 c32与 c64 组合得到 c64与 c128 组合得到 c128。×的含义表示这两种类型不能进行推导计算。从调用角度看这类组合属于 API 不支持的输入类型组合无法完成自动类型统一具体报错行为以 API 实际返回为准。五、Tensor-Tensor 推导示例原文档给出的两个标准示例调用aclnnAdd接口时如果输入参数的数据类型不一致一个为 float16一个为 float32那么 API 内部就会将 float16 的数据类型转换成 float32 的数据类型然后进行计算对应表1中f16 f32 → f32。调用aclnnAdd接口时如果输入参数的数据类型不一致一个为 float32一个为 bool那么 API 内部就会将 bool 的数据类型转换成 float32 的数据类型然后进行计算对应表1中bool f32 → f32。六、Tensor-Scalar 互推导规则当一个 API如aclnnAdds、aclnnMuls等输入的 Tensor 数据类型和输入的 Scalar 数据类型不一致时API 内部会推导出一个数据类型将输入数据转换成该数据类型进行计算。类型推导的规则如下表所示说明为方便描述表格中使用的数据类型是简写形式代表的含义同表1。表格里表头表示待推导的输入Tensor数据类型最左侧一列表示待推导的输入Scalar数据类型表格中对应位置表示推导出的数据类型。表中叉号x表示这两种类型不能进行推导计算。表2Tensor-Scalar 数据类型推导关系表数据类型f32f16f64bf16s8u8s16u16s32u32s64u64boolc32c64c128f32f32f16f64bf16f32f32f32xf32xf32xf32c32c64c128f16f32f16f64bf16f32f32f32xf32xf32xf32c32c64c128f64f32f16f64bf16f32f32f32xf32xf32xf32c128c128c128bf16f32f16f64bf16f32f32f32xf32xf32xf32c32c64c128s8f32f16f64bf16s8u8s16u16s32u32s64u64s8c32c64c128u8f32f16f64bf16s8u8s16u16s32u32s64u64u8c32c64c128s16f32f16f64bf16s8u8s16u16s32u32s64u64s16c32c64c128u16f32f16f64bf16s8u8s16u16s32u32s64u64xc32c64c128s32f32f16f64bf16s8u8s16u16s32u32s64u64s32c32c64c128u32f32f16f64bf16s8u8s16u16s32u32s64u64xc32c64c128s64f32f16f64bf16s8u8s16u16s32u32s64u64s64c32c64c128u64f32f16f64bf16s8u8s16u16s32u32s64u64xc32c64c128boolf32f16f64bf16s8u8s16u16s32u32s64u64boolc32c64c128c32c64c32c128c64c64c64c64c64c64c64c64c64c64c32c64c128c64c64c32c128c64c64c64c64c64c64c64c64c64c64c32c64c128c128c64c32c128c64c64c64c64c64c64c64c64c64c64c32c64c128七、Tensor-Scalar 推导规律解读从表2可以归纳出以下规律特别注意它与表1Tensor-Tensor的行为有明显差异Tensor 的实数类型通常主导推导结果当 Scalar 为浮点、整数或 bool 时结果往往跟随 Tensor 的类型。例如Tensor(f16) Scalar(f32) → f16、Tensor(f64) Scalar(f32) → f64、Tensor(u8) Scalar(s8) → u8。这与表1中精度更高的浮点胜出形成鲜明对比在 Tensor-Scalar 组合中即使 Scalar 精度更高也会被转换到 Tensor 的类型。例外Tensor 为整数、Scalar 为浮点时结果为浮点。例如Tensor(s8) Scalar(f32) → f32、Tensor(u64) Scalar(f16) → f16。此时 Scalar 的浮点类型成为计算类型整数 Tensor 被提升。整数 Scalar 与整数 Tensor 组合时结果为 Tensor 的类型Scalar(s8) Tensor(u16) → u16、Scalar(s64) Tensor(s32) → s32。注意在表2中整数 Scalar 与u16/u32/u64Tensor 的组合是允许的结果为对应的无符号类型这与表1中u16/u32/u64的孤立行为不同。bool Scalar 转换为 Tensor 的类型Scalar(bool) Tensor(f32) → f32、Scalar(bool) Tensor(s16) → s16但当 Tensor 为u16/u32/u64时组合为x不支持。复数 Scalar 与实数 Tensor 组合时结果由 Tensor 的实数精度决定复数档位Scalar(c32) Tensor(f16) → c32、Scalar(c32) Tensor(f32) → c64、Scalar(c32) Tensor(f64) → c128整数/布尔 Tensor 与复数 Scalar 组合统一得到 c64。值得注意表2中c64、c128两个 Scalar 行与c32行结果完全一致即复数 Scalar 自身的精度不参与抬升结果只取决于 Tensor 一侧。复数 Tensor 与任意 Scalar 组合时结果为 Tensor 的复数类型Scalar(c32) Tensor(c64) → c64、Scalar(f32) Tensor(c32) → c32复数 Tensor 类型保持主导。八、Tensor-Scalar 推导示例原文档给出的两个标准示例如果输入 Tensor 的数据类型为 float16输入 Scalar 的数据类型为 float32那么 API 内部就会将输入 Scalar 的 float32 数据类型转换成 float16 数据类型然后进行计算对应表2中Tensor(f16) Scalar(f32) → f16。如果输入 Tensor 的数据类型为 bool输入 Scalar 的数据类型为 float32那么 API 内部就会将输入 Tensor 的 bool 数据类型转换成 float32 数据类型然后进行计算对应表2中Tensor(bool) Scalar(f32) → f32。两个示例恰好展示了对偶的两种行为前者是Scalar 迁就 Tensor后者是Tensor 迁就 Scalar具体走哪条路完全由表2决定。九、实战在算子调用中观察互推导9.1 两段式调用中的推导发生位置互推导发生在算子 API 内部调用方无需也无法显式指定计算类型。以aclnnAdd类接口为例完整的两段式调用流程为调用第一段接口aclnnAddGetWorkspaceSize(...)API 在此阶段完成输入校验与类型推导并返回计算所需的workspaceSize与executor按workspaceSize申请 NPU 上的 workspace 内存调用第二段接口aclnnAdd(...)真正执行计算。因此只要输入 Tensor或 Scalar的数据类型不同推导动作就会在第一段接口内部自动触发推导出的计算类型决定了输入被转换为何种类型参与运算。9.2 结合 add 样例工程观察调用形态仓库的 add_example 样例 给出了一个完整的 aclnn 两段式调用实现其关键步骤为通过aclCreateTensor(shape.data(), shape.size(), dataType, strides.data(), 0, aclFormat::ACL_FORMAT_ND, ...)创建aclTensor代码位置其中dataType传入aclDataType枚举如ACL_FLOAT调用第一段接口aclnnAddExampleGetWorkspaceSize(selfX, selfY, out, workspaceSize, executor)代码位置调用第二段接口aclnnAddExample(workspaceAddr, workspaceSize, executor, stream)代码位置。在该样例中两个输入selfX、selfY均以ACL_FLOAT创建代码位置类型一致不会触发互推导。读者若将其中一个输入的dataType改为ACL_FLOAT16保持形状一致即可按表1验证f16 f32 → f32的推导行为API 内部会将 f16 输入提升为 f32 后计算输出仍以 f32 类型写入。同样在aclnnAdds类含 Scalar 参数的接口中可对照表2验证 Tensor 与 Scalar 类型不一致时的推导结果。9.3 如何运行算子样例验证推导行为项目提供了无需搭建调用工程的快速验证方式完整说明见 算子调用。基于自定义算子包执行算子样例的通用命令为bash build.sh --run_example ${op} ${mode} ${pkg_mode} [--example_name${example_name}] [--vendor_name${vendor_name}] [--soc${soc_version}] [--simulator${simulator}] [--experimental${experimental}]其中${op}待执行算子名小写下划线形式如flash_attention_score${mode}调用方式支持eageraclnn 调用与graph图模式调用${pkg_mode}包模式目前仅支持cust自定义算子包${example_name}可选examples 目录下样例文件去掉test_aclnn_前缀和.cpp后缀的名称${vendor_name}可选与构建的自定义算子包设置一致默认custom${soc_version}可选NPU 型号默认ascend910b${simulator}可选仿真模式仅eager场景可用${experimental}可选执行experimental贡献目录下的算子。基于 ops-transformer 包执行样例时命令简化为bash build.sh --run_example ${op} ${mode} [--soc${soc_version}]通过修改样例中创建aclTensor时的dataType参数即可构造不同类型的输入直观验证本文两张推导表中各组合的推导结果。十、互推导与互转换、广播的关系一个完整的语义闭环在 aclnn 算子 API 内部一次调用输入到输出的完整数据处理链条可以概括为形状维度若输入形状不一致先按 broadcast关系 的规则进行广播维度数不足的在左侧补 1然后按维度 1 拉伸使形状兼容类型维度输入若输入数据类型不一致按本文的互推导规则推导出统一的计算类型并将各输入转换为该类型计算以统一后的形状与计算类型在 NPU 上执行算子运算类型维度输出若输出aclTensor声明的数据类型与推导出的计算类型不一致按 互转换关系 将计算结果转换为输出类型。其中互转换关系对哪些类型之间可以转换有独立约束整数类型间可以转换也支持往浮点、复数类型转换浮点类型间可以转换也支持往复数类型转换复数类型间可以转换BOOL 支持往整数、浮点、复数类型转换除此之外的转换均不支持。因此即使某组输入能通过互推导得到计算类型最终输出能否按声明类型写出还需要满足互转换关系的约束——三个文档broadcast、互推导、互转换共同决定了 aclnn 接口对混合形状、混合类型输入的完整支持边界。十一、小结与使用建议本文围绕 CANN/ops-transformer 的 互推导关系 文档完整呈现了 Tensor-Tensor、Tensor-Scalar 两套数据类型推导规则表并结合规律解读与调用样例给出了可落地的预判方法。核心要点速查如下两套规则体系并存且行为不同Tensor-Tensor 组合遵循精度提升/复数支配规则表1Tensor-Scalar 组合通常由 Tensor 类型主导表2需分别记忆不能混用。u16、u32、u64需要特别小心表1中它们只能与自身组合表2中整数 Scalar 可与之组合但浮点/布尔 Scalar 与之组合为x。bool 是可塑性最强的类型它会向对侧数值类型转换不会压制任何数值类型但反过来也意味着与 bool 混算时计算结果可能发生精度变化。复数类型不可逆地支配结果只要任一侧是复数计算类型即为复数且档位c32/c64/c128由实数侧的精度决定。推导规则仅覆盖 16 种类型ACL_INT4、ACL_UINT1、各类 Float8/Float6/Float4 以及ACL_STRING等类型不适用上述推导表是否可用取决于具体 API 的参数支持范围。验证手段现成可通过 add_example 样例 构造不同dataType的输入结合 算子调用 的build.sh --run_example快速运行验证。编写调用代码时建议为每个算子 API 的输入提前核对两张推导表避免出现×组合导致的调用失败同时留意 Tensor-Scalar 场景下Scalar 被转换到 Tensor 类型带来的潜在精度变化如 f64 Scalar 遇到 f32 Tensor 时会降精度参与计算。【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表