PyPTO vf 寄存器类型转换全解析:astype / bit_cast / truncate 的底层原理与实战
【免费下载链接】pyptoPyPTO(发音: pai p-t-o):Parallel Tensor/Tile Operation编程范式。项目地址: https://gitcode.com/cann/pypto
PyPTO(Parallel Tensor/Tile Operation)的 SIMD 寄存器计算(reg_computation)体系中,vf.astype、vf.bit_cast、vf.truncate三个接口覆盖了向量寄存器(RegTensor)上全部的类型转换需求:数值语义转换、按位重解释与向零取整。本文以 docs/zh/api/pro_api/SIMD-API/reg_computation/type_conversion/index.md 为主线,结合仓库中的接口声明与示例代码,系统讲解三个接口的语义差异、参数行为、支持的数据类型矩阵以及可直接运行的调用示例,帮助你写出符合硬件约束、行为可预期的高效向量内核。
类型转换接口总览:三个接口,三种语义
在 PyPTO 的向量函数(@pl.vector_function)中,RegTensor 是向量指令的操作对象,而数据类型转换是浮点量化、精度对齐、位运算前的必经环节。type_conversion目录下三个接口分别解决三类问题:
| 接口 | 语义 | 典型场景 |
|---|---|---|
| vf.astype | 数值类型转换(vcvt 指令) | FP32→FP16、BF16→FP4、FP32→INT32 等真实数值换算 |
| vf.bit_cast | 按位重解释(仅改类型标签) | 对浮点寄存器执行按位与/或/异或 |
| vf.truncate | 向零取整(vtrc + ROUND_Z) | 丢弃小数部分、保留浮点数据类型 |
三者最关键的区别在于:astype会真实改变元素的值(如 257.0 转为 uint8 后变成 1),bit_cast保持比特模式完全不变,而truncate只做小数截断、数据类型不变。从源码声明看,三个接口均定义于 python/pypto_pro/language/_vf_api.py(astype位于 L709、bit_cast位于 L1466、truncate位于 L1864),并且在 python/pypto_pro/language/parser/_call_parser.py 的解析表中分别注册(L506-L541),是 vf API 的一等公民。
产品支持情况
根据三个接口文档中一致的「产品支持情况」声明:
- Ascend 950PR / Ascend 950DT:支持(三个接口全部支持,含 FP8/FP4 低精度转换);
- Atlas A3 训练/推理系列:不支持;
- Atlas A2 训练/推理系列:不支持。
即类型转换接口是 Ascend 950 系列特有的向量能力,编写代码前需先确认目标设备。仓库中 python/pypto_pro/language/_vf_api.py L722 也明确标注 "FP8/FP4 conversion support (Ascend 950PR/950DT only)"。
vf.astype:通用数值类型转换
vf.astype是三者中功能最丰富、参数最多的接口,用于将源操作数的数据类型转换为目标数据类型,覆盖浮点转整数、浮点转浮点、整数转浮点、整数转整数四类转换。
函数原型与参数
astype(src, preg, dtype: DType, layout: Optional[CastLayout] = None, round_mode: Optional[VFRoundMode] = None, saturate: Optional[SaturateMode] = None, mode: Optional[MergeMode] = None) -> dst| 参数 | 输入/输出 | 说明 |
|---|---|---|
src | 输入 | 源操作数,reg_tensor。 |
preg | 输入 | mask_reg,按源操作数筛选有效元素。 |
dtype | 输入 | 必选,目标寄存器数据类型(如pl.DT_FP16、pl.DT_INT32)。目标类型与源类型不同,必须显式指定。 |
layout | 输入 | 可选,CastLayout。源与目的位宽不同时,控制位宽小的元素在寄存器中的排布:CastLayout.ZERO放偶数索引(默认)、CastLayout.ONE放奇数索引;FP4 类型额外支持CastLayout.TWO/CastLayout.THREE(4 倍扩展/缩窄时使用)。 |
round_mode | 输入 | 可选,VFRoundMode,浮点舍入模式,默认VFRoundMode.CAST_RINT。 |
saturate | 输入 | 可选,SaturateMode,OFF(默认,非饱和)或ON(饱和)。 |
mode | 输入 | 可选,MergeMode。当前设备仅支持MergeMode.ZEROING(preg 未筛选元素在 dst 中置 0),MERGING不支持。 |
layout、mask 与位宽的关系
不同数据类型对应元素在 preg 中的位宽不一致,类型转换时 mask_reg 按源操作数筛选有效元素。当源和目的位宽不同时,单条指令的计算量以位宽更大的数据类型为准,layout决定位宽小的元素在寄存器中的排布:
- 16 位宽 → 32 位宽:小位宽元素按
layout展开到 32 位寄存器中的偶数/奇数半区(图 1 展示了 mask_reg 与 layout 同时作用下的转换过程); - 32 位宽 → 16 位宽:每两个有效位宽元素压缩到一个 32 位寄存器(图 2);
- FP4 特例:
DT_FP4E2M1、DT_FP4E1M2与DT_BF16之间的转换,指令按每 2 个元素为一对读写,大转小时 preg 有效位以偶数位为准(图 3、图 4)。
相关的四张转换过程示意图(astype_b16_to_b32_conversion.jpg、astype_b32_to_b16_conversion.jpg、astype_fp4x2_e2m1_to_bf16_conversion.jpg、astype_bf16_to_fp4x2_e2m1_conversion.jpg)存放在 docs/zh/api/pro_api/figures 目录,可对照理解位宽变化时的元素排布。
支持的数据类型转换矩阵
astype支持的数据类型转换组合非常多,以下按转换类别整理(完整表格请参见 astype.md 的「约束说明」):
表1 支持的数据类型转换(节选)
| src | dst |
|---|---|
| DT_INT4 | DT_INT16、DT_FP16、DT_BF16 |
| DT_INT8 | DT_INT16、DT_FP16、DT_INT32 |
| DT_UINT8 | DT_UINT16、DT_FP16、DT_UINT32 |
| DT_FP4E2M1 / DT_FP4E1M2 | DT_BF16 |
| DT_HF8 | DT_FP16、DT_FP32 |
| DT_FP8E8M0 | DT_BF16 |
| DT_FP8E5M2 / DT_FP8E4M3FN | DT_FP32 |
| DT_FP16 | DT_INT4、DT_INT8、DT_UINT8、DT_HF8、DT_INT16、DT_BF16、DT_INT32、DT_FP32 |
| DT_BF16 | DT_FP4E2M1、DT_FP4E1M2、DT_FP8E8M0、DT_FP16、DT_INT32、DT_FP32 |
| DT_FP32 | DT_HF8、DT_FP8E5M2、DT_FP8E4M3FN、DT_INT16、DT_FP16、DT_BF16、DT_INT32、DT_INT64 |
| DT_INT64 | DT_INT32、DT_FP32 |
从源码 docstring(python/pypto_pro/language/_vf_api.py L722-L735)可以确认,FP8/FP4 低精度转换路径与文档表完全一致,二者互为印证。
不同场景下的参数可用性
表2 浮点转整数(节选)
| src | dst | layout | saturate | mode | round_mode |
|---|---|---|---|---|---|
| DT_FP16 | DT_INT4 | ZERO/ONE/TWO/THREE | OFF/ON | ZEROING | CAST_RINT/CAST_ROUND/CAST_FLOOR/CAST_CEIL/CAST_TRUNC |
| DT_FP16 | DT_INT8 | ZERO/ONE | OFF/ON | ZEROING | CAST_RINT/CAST_ROUND/CAST_FLOOR/CAST_CEIL/CAST_TRUNC |
| DT_FP16 | DT_INT16 | UNKNOWN | OFF/ON | ZEROING | CAST_RINT/CAST_ROUND/CAST_FLOOR/CAST_CEIL/CAST_TRUNC |
| DT_FP16 | DT_INT32 | ZERO/ONE | UNKNOWN | ZEROING | CAST_RINT/CAST_ROUND/CAST_FLOOR/CAST_CEIL/CAST_TRUNC |
| DT_BF16 | DT_INT32 | ZERO/ONE | OFF/ON | ZEROING | CAST_RINT/CAST_ROUND/CAST_FLOOR/CAST_CEIL/CAST_TRUNC |
| DT_FP32 | DT_INT16 / DT_INT64 | ZERO/ONE | OFF/ON | ZEROING | CAST_RINT/CAST_ROUND/CAST_FLOOR/CAST_CEIL/CAST_TRUNC |
| DT_FP32 | DT_INT32 | UNKNOWN | OFF/ON | ZEROING | CAST_RINT/CAST_ROUND/CAST_FLOOR/CAST_CEIL/CAST_TRUNC |
表3 浮点转浮点(节选)
| src | dst | layout | saturate | round_mode |
|---|---|---|---|---|
| DT_HF8 | DT_FP16 | ZERO/ONE | UNKNOWN | UNKNOWN |
| DT_HF8 | DT_FP32 | ZERO/ONE/TWO/THREE | UNKNOWN | UNKNOWN |
| DT_FP8E4M3FN / DT_FP8E5M2 | DT_FP32 | ZERO/ONE/TWO/THREE | UNKNOWN | UNKNOWN |
| DT_FP8E8M0 | DT_BF16 | ZERO/ONE | UNKNOWN | UNKNOWN |
| DT_FP4E2M1 / DT_FP4E1M2 | DT_BF16 | ZERO/ONE/TWO/THREE | UNKNOWN | UNKNOWN |
| DT_FP16 | DT_HF8 | ZERO/ONE | OFF/ON | CAST_ROUND/CAST_HYBRID |
| DT_FP16 | DT_BF16 | UNKNOWN | UNKNOWN | CAST_RINT/CAST_ROUND/CAST_FLOOR/CAST_CEIL/CAST_TRUNC |
| DT_BF16 | DT_FP4E2M1 / DT_FP4E1M2 | ZERO/ONE/TWO/THREE | UNKNOWN | CAST_RINT/CAST_ROUND/CAST_FLOOR/CAST_CEIL/CAST_TRUNC |
| DT_FP32 | DT_FP8E4M3FN / DT_FP8E5M2 | ZERO/ONE/TWO/THREE | OFF/ON | CAST_RINT |
| DT_FP32 | DT_FP16 | ZERO/ONE | OFF/ON | CAST_ODD/CAST_RINT/CAST_ROUND/CAST_FLOOR/CAST_CEIL/CAST_TRUNC |
| DT_FP32 | DT_BF16 | ZERO/ONE | OFF/ON | CAST_RINT/CAST_ROUND/CAST_FLOOR/CAST_CEIL/CAST_TRUNC |
表4 整数转浮点(节选)
| src | dst | layout | saturate | round_mode |
|---|---|---|---|---|
| DT_INT4 | DT_FP16 / DT_BF16 | ZERO/ONE/TWO/THREE | UNKNOWN | UNKNOWN |
| DT_INT8 / DT_UINT8 | DT_FP16 | ZERO/ONE | UNKNOWN | UNKNOWN |
| DT_INT16 | DT_FP16 | UNKNOWN | UNKNOWN | CAST_RINT/CAST_ROUND/CAST_FLOOR/CAST_CEIL/CAST_TRUNC |
| DT_INT16 | DT_FP32 | ZERO/ONE | UNKNOWN | UNKNOWN |
| DT_INT32 | DT_FP32 | UNKNOWN | UNKNOWN | CAST_RINT/CAST_ROUND/CAST_FLOOR/CAST_CEIL/CAST_TRUNC |
| DT_INT64 | DT_FP32 | ZERO/ONE | UNKNOWN | CAST_RINT/CAST_ROUND/CAST_FLOOR/CAST_CEIL/CAST_TRUNC |
表5 整数转整数(节选)
| src | dst | layout | saturate | mode | round_mode |
|---|---|---|---|---|---|
| DT_INT4 | DT_INT16 | ZERO/ONE/TWO/THREE | UNKNOWN | ZEROING | UNKNOWN |
| DT_INT8 | DT_INT16 | ZERO/ONE | UNKNOWN | ZEROING | UNKNOWN |
| DT_INT8 | DT_INT32 | ZERO/ONE/TWO/THREE | UNKNOWN | ZEROING | UNKNOWN |
| DT_INT16 | DT_INT4 / DT_UINT8 | ZERO/ONE(/TWO/THREE) | OFF/ON | ZEROING | UNKNOWN |
| DT_INT16 | DT_INT32 / DT_UINT32 | ZERO/ONE | UNKNOWN | ZEROING | UNKNOWN |
| DT_INT32 | DT_UINT8 | ZERO/ONE/TWO/THREE | OFF/ON | ZEROING | UNKNOWN |
| DT_INT32 | DT_INT16 / DT_UINT16 | ZERO/ONE | OFF/ON | ZEROING | UNKNOWN |
| DT_INT32 | DT_INT64 | ZERO/ONE | UNKNOWN | ZEROING | UNKNOWN |
| DT_UINT32 | DT_UINT8 | ZERO/ONE/TWO/THREE | OFF/ON | ZEROING | UNKNOWN |
| DT_INT64 | DT_INT32 | ZERO/ONE | OFF/ON | ZEROING | UNKNOWN |
UNKNOWN 的含义:
layout为 UNKNOWN 表示源与目的位宽相同、无需指定(可省略);saturate为 UNKNOWN 表示该路径不涉及饱和/非饱和选择(可省略);round_mode为 UNKNOWN 表示该路径不涉及精度损失(可省略)。
饱和与不饱和模式的边界行为
饱和/不饱和是astype最容易踩坑的点,文档明确给出了四种场景下的行为差异:
| 场景 | 不饱和模式 | 饱和模式 |
|---|---|---|
| 浮点转整数 | 输入超过输出类型最值时截断为目标格式数据宽度(保留最低有效位),例如 half 输入 257 → uint8_t 输出 1;输入 ±inf 返回输出类型对应最值;输入 nan 返回 0。 | 输入超过最值时返回输出类型最值,例如 half 输入 257 → uint8 输出 255,half 输入 -inf → uint8_t 输出 0;输入 nan 返回 0。 |
| 浮点转浮点 | 输入 nan 输出 nan;输入 ±inf 输出 ±inf。 | 输入 nan 输出 0;输入超过最值时返回输出类型最值。 |
| 整数转浮点 | 不支持不饱和模式 | 输入 nan 输出 0;输入超过最值返回输出类型最值。该场景默认饱和模式,无需配置。 |
| 整数转整数 | 输入截断为目标数据宽度,例如 int32_t 输入 256 → uint8_t 输出 0。 | 输入超出目标数据范围时饱和为目标数据最值。 |
浮点转浮点的特殊约束(务必留意):
- 输出为FP32时只支持不饱和模式;
- 不饱和模式下输出为FP8E4M3FN时(该类型无 inf 表示),溢出输出为 nan;
- 饱和模式下输出为FP8E5M2/FP8E4M3FN时,输入 nan 默认输出 0;若
CTRL[50] = 1'b1则输出 nan; - FP4E2M1/FP4E1M2 没有 inf 和 nan 的定义:BF16→FP4 转换中,输入 inf 或超出 FP4 最值范围时返回对应符号的 FP4 最值,输入 nan 时 FP4 输出 0;
- FP8E8M0:输入 BF16 ±inf 或绝对值超出 FP8E8M0 最大值时,返回最大值
0b11111110;输入 BF16 nan 输出 FP8E8M0 nan =0b11111111。
整数转整数的特殊约束:对于窄类型(如 INT16)转宽类型(如 UINT32),只支持饱和模式,输入负数会被饱和成 0。
vf.astype 实战:五种典型转换的内核写法
文档提供了六个可直接运行的示例(BF16→FP32、FP32→FP16、FP32→INT32、FP32→FP8E4M3FN、BF16→FP4E2M1、FP16→HF8、INT64),这里选取最具代表性的五种讲解模式,其余示例可在 astype.md 中查看完整代码。
模式一:BF16 → FP32(窄转宽,需指定 layout)
import os import pypto_pro.language as pl import torch import torch_npu @pl.vector_function def example_vf(src_tile, dst_tile): preg = vf.create_mask(pattern=pl.MaskPattern.ALL, dtype=pl.DT_FP32) reg_a = vf.load_align(src_tile, 0) reg_bf16 = vf.astype(reg_a, preg, dtype=pl.DT_BF16, layout=pl.CastLayout.ZERO) reg_f32 = vf.astype(reg_bf16, preg, dtype=pl.DT_FP32) vf.store_align(dst_tile, reg_f32, preg) @pl.jit() def example_kernel( a: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_FP32], out: pl.Tensor[[pl.DYNAMIC, pl.DYNAMIC], pl.DT_FP32], ): tf = pl.TileType(shape=[1, 64], dtype=pl.DT_FP32, target_memory=pl.MemorySpace.Vec) in_a_grp = pl.make_tile_group(type=tf, addrs=0x0, mutex_ids=[0]) in_a = in_a_grp.current() t_out_grp = pl.make_tile_group(type=tf, addrs=0x100, mutex_ids=[1]) t_out = t_out_grp.current() with pl.section_vector(): pl.load(in_a, a, [0, 0]) example_vf(in_a, t_out) pl.store(out, t_out, [0, 0]) def test_example(): device_id = int(os.environ.get("TILE_FWK_DEVICE_ID", 0)) device = f"npu:{device_id}" core_nums = 1 torch.npu.set_device(device) a = torch.randn([1, 64], device=device, dtype=torch.float32) out = torch.empty([1, 64], device=device, dtype=torch.float32) example_kernelNone, core_nums torch.npu.synchronize() torch.testing.assert_close(out, a.to(torch.bfloat16).to(torch.float32), rtol=1e-3, atol=1e-3) if __name__ == "__main__": test_example() print("PASSED")要点:FP32 寄存器(64 个元素)转 BF16 后只占偶数位置(CastLayout.ZERO),再转回 FP32 即可与原始数据对比,验证 round-trip 正确性。
模式二:FP32 → FP16(含舍入的浮点转浮点)
结构与模式一完全相同,区别仅在转换参数上:reg_f16 = vf.astype(reg_a, preg, dtype=pl.DT_FP16, layout=pl.CastLayout.ZERO)。验证时使用a.to(torch.float16).to(torch.float32)与输出对比(rtol=1e-3, atol=1e-3),证明转换遵循 FP16 舍入语义。
模式三:FP32 → INT32(指定 round_mode 的浮点转整数)
@pl.vector_function def example_vf_round(src_tile, dst_tile): preg = vf.create_mask(pattern=pl.MaskPattern.ALL, dtype=pl.DT_FP32) reg_a = vf.load_align(src_tile, 0) reg_i = vf.astype(reg_a, preg, dtype=pl.DT_INT32, round_mode=pl.VFRoundMode.CAST_RINT) reg_f = vf.astype(reg_i, preg, dtype=pl.DT_FP32) vf.store_align(dst_tile, reg_f, preg)验证采用torch.testing.assert_close(out, a.to(torch.int32).to(torch.float32), rtol=0, atol=1.0),容忍 ±1 的舍入差异,说明CAST_RINT是四舍五入语义而非截断。
模式四:FP32 → FP8E4M3FN(量化场景:layout + round_mode + saturate 齐用)
@pl.vector_function def example_vf_fp8(src_tile, dst_tile): preg = vf.create_mask(pattern=pl.MaskPattern.ALL, dtype=pl.DT_FP32) reg_a = vf.load_align(src_tile, 0) reg_f8 = vf.astype(reg_a, preg, dtype=pl.DT_FP8E4M3FN, layout=pl.CastLayout.ZERO, round_mode=pl.VFRoundMode.CAST_RINT, saturate=pl.SaturateMode.ON) reg_f32 = vf.astype(reg_f8, preg, dtype=pl.DT_FP32) vf.store_align(dst_tile, reg_f32, preg)验证时只对比偶数索引列out[:, ::2]与expected[:, ::2],这正是 layout 排布(FP8 存于偶数位置)在结果上的体现——窄类型转换后,奇数位置是无效元素。
模式五:BF16 → FP4E2M1(FP4 特例:每 2 个元素一对读写)
_FP4_E2M1_VALUES = np.array([0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0], dtype=np.float32) _FP4_E2M1_EVEN_MASK = np.array([True, False, True, False, True, False, True, False]) @pl.vector_function def example_vf_fp4(src_tile, dst_tile): preg = vf.create_mask(pattern=pl.MaskPattern.ALL, dtype=pl.DT_BF16) reg_a = vf.load_align(src_tile, 0, dtype=pl.DT_BF16) reg_f4 = vf.astype(reg_a, preg, dtype=pl.DT_FP4E2M1, layout=pl.CastLayout.ZERO, round_mode=pl.VFRoundMode.CAST_RINT) reg_bf16 = vf.astype(reg_f4, preg, dtype=pl.DT_BF16) vf.store_align(dst_tile, reg_bf16, preg)验证逻辑揭示了 FP4 转换的硬件特性:测试在 numpy 中枚举 FP4E2M1 的 8 个可表示值(0.0/0.5/1.0/1.5/2.0/3.0/4.0/6.0),寻找最近邻且优先偶数索引候选(_FP4_E2M1_EVEN_MASK),再与out[:, ::2]对比——这与文档中"FP4 大转小时 preg 有效位以偶数位为准"的说明完全对应。
模式六:INT64 数据类型示例
INT64 是 64 位宽类型,转换路径为INT64 → FP32 → INT64(astype.md L500-L545 有完整代码)。内核使用TileType(shape=[1, 32], dtype=pl.DT_INT64, ...),Tile 基地址用 256 字节对齐(addrs=256),最终以rtol=0, atol=0精确断言 round-trip 无损。
vf.bit_cast:按位重解释,零开销的类型标签切换
语义与使用形式
vf.bit_cast对向量寄存器执行按位类型强转,不进行任何数值转换:仅改变类型标签,比特模式不变。与astype的本质区别在于——astype是真正的数值转换(元素值会变),bit_cast只重解释。
bit_cast(src, dtype: DType) -> dst| 参数 | 输入/输出 | 说明 |
|---|---|---|
src | 输入 | 源操作数,reg_tensor。 |
dtype | 输入 | 目标数据类型,将 src 重解释为该类型。 |
支持两种使用方式:
- 赋值形式:
dst = vf.bit_cast(src, dtype=xxx),将 src 按位重解释后赋给新声明的 dst 寄存器; - 嵌套参数形式:
vf.xor(vf.bit_cast(reg_a, dtype=pl.DT_UINT32), vf.bit_cast(reg_b, dtype=pl.DT_UINT32), preg),作为其他 vf.xxx 调用的参数,满足指令对操作数类型的要求。
核心约束:源与目标操作数位宽必须相同(src与dtype的GetBit()返回值一致),否则行为未定义。从源码看,bit_cast生成的是 C++ 侧的RegTensor<T>&引用强转,例如vxor(reg_c, (RegTensor<float>&)reg_a, (RegTensor<float>&)reg_b, preg, MODE_ZEROING)(见 python/pypto_pro/language/_vf_api.py L1482-L1488),因此它本质上是一条"类型注解"而非真实指令,零开销。
实战一:赋值形式 — FP32 寄存器按位异或
@pl.vector_function def example_vf_bit_cast_assign(src_tile_a, src_tile_b, dst_tile): preg_u32 = vf.create_mask(pattern=pl.MaskPattern.ALL, dtype=pl.DT_UINT32) reg_a = vf.load_align(src_tile_a, 0) reg_b = vf.load_align(src_tile_b, 0) reg_a_u32 = vf.bit_cast(reg_a, dtype=pl.DT_UINT32) reg_b_u32 = vf.bit_cast(reg_b, dtype=pl.DT_UINT32) reg_c = vf.xor(reg_a_u32, reg_b_u32, preg_u32) vf.store_align(dst_tile, reg_c, preg_u32)验证用torch.equal(out.view(torch.int32), a.view(torch.int32) ^ b.view(torch.int32))——把输出与输入均view成 int32 后逐位异或对比,精准验证"比特模式不变"的语义。完整的内核与测试代码(含pl.load/pl.store与section_vector上下文)见 bit_cast.md L56-L106。
实战二:嵌套参数形式 — HF8 寄存器按位或
HF8 是 8 位存储类型,加载后 RegTensor 含 256 个元素;将其bit_cast为 DT_UINT8(同为 256 元素)后即可对 UINT8 寄存器执行vf.or_:
@pl.vector_function def example_vf_hf8_to_uint8(src_tile_a, src_tile_b, dst_tile): preg_b8 = vf.create_mask(pattern=pl.MaskPattern.ALL, dtype=pl.DT_UINT8) reg_a = vf.load_align(src_tile_a, 0, dtype=pl.DT_HF8) reg_b = vf.load_align(src_tile_b, 0, dtype=pl.DT_HF8) reg_c = vf.or_(vf.bit_cast(reg_a, dtype=pl.DT_UINT8), vf.bit_cast(reg_b, dtype=pl.DT_UINT8), preg_b8) vf.store_align(dst_tile, reg_c, preg_b8)测试侧通过torch_npu.npu_dtype_cast生成 HF8 张量,并用a.view(torch.uint8) | b.view(torch.uint8)构造期望值(bit_cast.md L146-L163)。这是bit_cast最典型的应用:浮点类型没有按位运算指令,先重解释为等位宽的整数类型再参与位运算。
vf.truncate:向零取整,保留浮点类型
vf.truncate将源操作数中的浮点元素截断为整数值(保留原数据类型)并存入目的操作数:dstReg_i = trunc(srcReg_i)。例如 3.7 → 3.0、-2.3 → -2.0,即丢弃小数部分、向零取整(示意图见 truncate_function.jpg)。
truncate(src, preg, mode: Optional[MergeMode] = None) -> dst| 参数 | 输入/输出 | 说明 |
|---|---|---|
src | 输入 | 源操作数,reg_tensor,与 dst 数据类型保持一致,支持 DT_FP16、DT_BF16、DT_FP32。 |
preg | 输入 | mask_reg,mask 未筛选的元素在 dst 中置零。DT_FP32 只支持不饱和模式。 |
mode | 输入 | 可选,MergeMode。仅支持MergeMode.ZEROING(默认),MERGING当前不支持。 |
从源码看,truncate映射到硬件vtrc指令的ROUND_Z模式(python/pypto_pro/language/_vf_api.py L1867),即向零舍入(round toward zero),与 C 标准库的trunc、PyTorch 的torch.trunc语义一致。约束说明为"无"——接口本身没有额外的数据类型限制,仅需保证 src 为上述三种浮点类型。
调用示例(完整代码见 truncate.md):
@pl.vector_function def example_vf(src_tile, dst_tile): preg = vf.create_mask(pattern=pl.MaskPattern.ALL, dtype=pl.DT_FP32) src_reg = vf.load_align(src_tile, 0) dst_reg = vf.truncate(src_reg, preg) vf.store_align(dst_tile, dst_reg, preg)验证直接使用torch.testing.assert_close(out, torch.trunc(a), rtol=1e-5, atol=1e-5),与 PyTorch 的torch.trunc一一对应。
注意区分truncate与astype的CAST_TRUNC:前者保留浮点数据类型,只做向零取整(3.7 → 3.0,仍是 FP32);后者是类型转换中的舍入模式选项,输出为目标整数/浮点类型。若需要 FP32 → INT32 的截断转换,应使用vf.astype(..., dtype=pl.DT_INT32, round_mode=pl.VFRoundMode.CAST_TRUNC)。
三个接口的选择决策
| 需求 | 选择 | 理由 |
|---|---|---|
| 改变数据类型并保留数值语义(FP16↔FP32、量化、INT↔FP) | vf.astype | 真正的 vcvt 数值转换,支持 layout/round_mode/saturate 完整控制 |
| 对浮点寄存器做按位运算(与/或/异或、位操作) | vf.bit_cast+ 整数位运算指令 | 仅改类型标签,比特不变,零开销满足指令类型要求 |
| 不改变类型、仅丢弃浮点小数部分 | vf.truncate | vtrc ROUND_Z,向零取整,保留原类型 |
三个接口之间还可以组合使用:典型流水线如astype窄化量化 → 位运算重解释 →astype宽化还原。编写时请始终遵循三条原则:
- 先确认设备:三个接口仅 Ascend 950PR/950DT 支持,A2/A3 系列不可用;
- 位宽不同必配 layout:源与目的位宽不同时按文档各表选择
CastLayout.ZERO/ONE(FP4 场景可用 TWO/THREE),窄化转换结果只出现在指定半区,其余位置无效; - 溢出行为提前约定:浮点转整数、浮点转浮点的溢出/nan/inf 行为由
saturate决定,且不同路径的可用性不同,务必对照 astype.md 的约束表选用,避免未定义行为。
延伸阅读
- type_conversion 目录索引:astype / bit_cast / truncate 三篇接口文档入口;
- reg_tensor:三个接口的源/目的操作数类型;
- mask_reg:preg 掩码的创建与语义;
- CastLayout、VFRoundMode、SaturateMode、MergeMode:astype 各参数的枚举定义;
- 接口声明源码:python/pypto_pro/language/_vf_api.py;
- vf API 解析注册:python/pypto_pro/language/parser/_call_parser.py。
【免费下载链接】pyptoPyPTO(发音: pai p-t-o):Parallel Tensor/Tile Operation编程范式。项目地址: https://gitcode.com/cann/pypto
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考