pypto_pro.language.Scalar#
产品支持情况#
Ascend 950PR/Ascend 950DT:支持
Atlas A3 训练系列产品/Atlas A3 推理系列产品:不支持
Atlas A2 训练系列产品/Atlas A2 推理系列产品:不支持
功能说明#
带数据类型的标量值类型标注,用于kernel标量参数声明或kernel内部局部变量类型标注。
pypto_pro.language.Scalar不能使用普通Python数值直接构造。显式标量常量通过pypto_pro.language.const创建;Scalar(expr)仅用于内部包装IR表达式。
pypto_pro.language.DT_*类型常量主要用于:
kernel函数签名中声明标量参数(如行数、列数、缩放系数等)
kernel内部局部变量的类型标注(如
pypto_pro.language.get_block_idx()返回值)通过
pypto_pro.language.const创建编译期常量标量时指定数据类型
可用类型常量#
类型常量 |
说明 |
|---|---|
|
布尔标量类型 |
|
8位有符号整数标量类型 |
|
16位有符号整数标量类型 |
|
32位有符号整数标量类型 |
|
64位有符号整数标量类型,常用于坐标和偏移计算 |
|
8位无符号整数标量类型 |
|
16位无符号整数标量类型 |
|
32位无符号整数标量类型 |
|
64位无符号整数标量类型 |
|
16位浮点标量类型 |
|
16位Brain浮点标量类型 |
|
32位浮点标量类型 |
约束说明#
以下低精度类型为存储和张量计算专用类型,不支持用于标量表达式:
DT_FP4、DT_FP8E4M3FN、DT_FP8E5M2、DT_FP8E8M0、DT_FP4E2M1、DT_FP4E1M2、DT_INT4、DT_UINT4、DT_HF4、DT_HF8
调用示例#
import pypto_pro.language as pl
# kernel 签名中声明标量参数
@pl.jit()
def kernel(
a: pl.Tensor[[64, 128], pl.DT_FP16],
rows: pl.DT_INT64,
cols: pl.DT_INT64,
scale: pl.DT_FP32,
out: pl.Tensor[[64, 128], pl.DT_FP16],
):
...
# kernel 内部局部变量类型标注
vidx = pl.get_block_idx()
offset: pl.DT_INT64 = vidx * 64
# kernel 内部显式创建标量常量
scale = pl.const(1.0, pl.DT_FP32)
zero_idx = pl.const(0, pl.DT_INT64)
显式创建标量值的完整说明见pypto_pro.language.const。