已关闭
[Requirement|需求建议]: 新增 cast_v3 算子 #2289
YOLO-MIC创建于  7月22日关闭于  29 天前
YOLO-MIC
YOLO-MIC
7月22日 创建

Thanks for sending an requirement! Please fill in the following template to help quickly solve your problem.

一、背景信息(必填)

在 Atlas 推理系列产品上使用 AscendC 实现 Cast 算子(命名 cast_v3),并新增支持 BF16 数据类型输入。

现有 Cast 算子(TBE 实现)在 310P AICore 上仅支持 40 组 dtype 组合,BF16 及其他未列出组合走 AICPU 实现,性能较差。在包含 bfloat16 计算的 LLM 微调、混合精度训练等场景中,bf16↔fp32、bf16↔int8/int16、int64→int32、fp16→bool 等类型转换路径会报 unsupported dtype combination,导致模型无法在 Atlas 推理系列产品上跑通。

本次任务使用 AscendC 新增独立的 cast_v3 算子,扩展支持的 dtype 组合至 53 组,并将原 AICPU 路径下沉到 AICore,性能显著提升。

二、价值/作用(必填)

cast_v3 算子的主要功能是执行张量类型转换,将输入 tensor 从一种 dtype 转换为另一种 dtype,计算公式 out = cast(x, dst_type)。在深度学习训练与推理中,类型转换是基础且高频的操作。

现有 Cast 算子部分 dtype 组合走 AICPU 实现,涉及 host→device 数据搬运与 CPU 计算,性能较差,成为端到端推理的瓶颈。cast_v3 算子通过 AscendC 实现,将这些组合下沉到 AICore 执行,利用 Vector/Core 侧硬件并行计算能力,性能显著优于 AICPU 方案。

三、设计方案(必填)

3.1 使能方式

Aclnn 直调(aclnnCastV3 两段式接口)

3.2 总体设计
3.2.1 算子支持的数据类型

输入:BOOL、FLOAT16、FLOAT、INT8、UINT8、INT16、INT32、INT64、BF16(共 9 种)
输出:BOOL、FLOAT16、FLOAT、INT8、UINT8、INT16、INT32(共 7 种,比输入少 BF16 和 INT64)
共 53 组合法 dtype 组合
属性:dst_type(int,指定目标 dtype)

支持形状:ND 格式,输入输出 shape 相同。

3.2.2 host 侧设计

Host 侧将数据视为一维向量,仅考虑数据个数,不考虑维度信息。

Tiling 数据结构:

字段 类型 含义
batchSize int64_t 输入元素总数
ubProcessNum int32_t 单次 UB 处理元素数
formerBatchSize uint64_t 首核处理的数据量
tailBatchSize uint64_t 尾核处理的数据量
formerCoreNum uint64_t 首核数量(多分配数据的核数量)

Tiling Key 规划:

Tiling Key 输入类型 输出类型 Kernel 类 说明
1 SINGLE_ROW(通用模式,单行处理) 任意 CastGeneric 通用数据类型转换路径
2 SEGMENTED(数据量较大,分块处理) 任意 CastGeneric 分块处理路径
3 PACKED(打包模式) 任意 CastGeneric 打包处理路径
4 BF16_SPECIAL(bf16 专用) 任意 CastBf16 BF16 输入路径,通过 float32 中转位重建

分核策略: 优先满核,核间不能均分时,前几个核多处理一块。
内存优化: 充分使用 UB 空间,根据 UB 大小、输入/输出 dtype 字节数、double buffer 及各 tiling key 路径所需临时 buffer 大小,计算单次 UB 处理元素数。

3.2.3 kernel 侧设计

Kernel 侧进行 Init 和 Process 两个阶段,其中 Process 包括 CopyIn → Compute → CopyOut 三个阶段。

  • cast_base.h:基类 CastBase,管理 GM/UB 地址、tiling 参数、RunProcess 模板循环
  • cast_generic.h:通用 kernel,按输入 T 分派 ComputeFromFP32 / ComputeFromFP16 / ComputeFromInt32 / ComputeFromInt64
  • cast_bf16.h:bf16 专用 kernel,通过位重建算法(拆高低字节 → 数值运算 → CAST_RINT 还原位模式)实现 bf16→float 转换
  • cast_ops.h:辅助函数 PackInt32ToInt16 / PackInt32ToByte / CastToBool
  • cast_tiling_key.h:tilingKey 常量定义

关键语义对齐:

  • 浮点→整型:NaN 转换为 0;截断小数部分(CAST_TRUNC)
  • 整型→窄整型:取低位字节(modulo-wrap),对齐 PyTorch 行为
  • →bool:非零即 True(bool(x) = (x != 0)
  • int64→float/half:通过 GatherMask 位重建
3.3 支持硬件
支持的芯片版本 涉及勾选
Atlas 推理系列产品

3.4 算子约束限制

  • 输入 shape 与输出 shape 必须相同
  • 输出 dtype 由 dst_type 属性指定,必须为支持的 8 种输出 dtype 之一
  • 不支持广播(shape 必须一致)
  • 仅支持 ND format
  • 浮点→整型时,输入数据中存在 nan,则将 nan 转换为 0
  • INT32→INT8 场景:只能保证输入数据在 (-2048, 1920) 范围内精度无误差
  • INT64→FLOAT32 场景:只能保证输入数据在 (-2147483648, 2147483647) 范围内精度无误差
  • 当前仅支持 ascend310p

关联 PR:
https://gitcode.com/cann/ops-math/merge_requests/4159

likedislike
YOLO-MICYOLO-MIC
7月22日 修改了issue 的描述
YOLO-MICYOLO-MIC
7月22日 修改了issue 的描述
陈思
陈思成员
7月23日 评论:

/assign @YOLO-MIC

likedislike
CANN-robotCANN-robot成员
7月23日 将 YOLO-MIC 设为负责人
CANN-robotCANN-robot成员
29 天前 关闭了 issue
CANN-robotCANN-robot成员
28 天前 添加了label:resolved