pypto_pro.language.Ptr#
产品支持情况#
Ascend 950PR/Ascend 950DT:支持
Atlas A3 训练系列产品/Atlas A3 推理系列产品:不支持
Atlas A2 训练系列产品/Atlas A2 推理系列产品:不支持
功能说明#
指向某种元素类型的全局内存裸指针类型标注,对应PTO MLIR的!pto.ptr<dtype>。
pypto_pro.language.Ptr主要用于:
kernel函数签名中声明GM裸指针参数
配合
pypto_pro.language.make_ptr创建不同元素类型的指针视图配合
pypto_pro.language.addptr做指针偏移配合
pypto_pro.language.make_tensor从裸指针构造tensor view
函数原型#
pypto_pro.language.Ptr[dtype]
参数类型#
参数 |
输入/输出 |
说明 |
|---|---|---|
|
输入 |
指针指向的元素数据类型 |
参数范围#
参数 |
输入/输出 |
说明 |
|---|---|---|
|
输入 |
|
调用示例#
下面是一个完整kernel:在kernel签名中用pypto_pro.language.Ptr声明workspace裸指针参数,配合addptr偏移到后半段、make_tensor包装成tensor view,完成a*2写回out。vector kernel开auto_mutex,同步由make_tile_group自动管理。
import pypto_pro.language as pl
@pl.jit(auto_mutex=True)
def workspace_kernel(
a: pl.Tensor[[64, 128], pl.DT_FP16],
workspace: pl.Ptr[pl.DT_FP16],
out: pl.Tensor[[64, 128], pl.DT_FP16],
):
ws_buf_ptr = pl.addptr(workspace, 64 * 128)
ws_buf = pl.make_tensor(ws_buf_ptr, [64, 128], [128, 1])
tt = pl.TileType(shape=[64, 128], dtype=pl.DT_FP16, target_memory=pl.MemorySpace.Vec)
tile = pl.make_tile_group(type=tt, addrs=0x0000, mutex_ids=[0])
with pl.section_vector():
t = tile.current()
pl.load(t, a, [0, 0])
pl.add(t, t, t)
pl.store(ws_buf, t, [0, 0])
pl.load(t, ws_buf, [0, 0])
pl.store(out, t, [0, 0])