SHMEM 设备端 API
本模块提供可在运行于昇腾 NPU 上的 Triton kernel 内调用的设备端 ACLSHMEM 操作。这些 API 支持对远程内存访问、同步以及 PE 间通信的细粒度控制。
所有设备端 API 可通过以下方式访问:
from triton_dist.language.extra.ascend import libaclshmem_device as shmem
源文件: python/triton_dist/language/extra/ascend/libaclshmem_device.py
关于底层 ACLSHMEM 库的更多信息,请参考 ACLSHMEM 文档。
拓扑与团队信息
my_pe
- my_pe() int32
获取当前 PE 的处理单元(PE)ID。
- 返回:
int32 —— 当前 PE 编号(0 到 n_pes - 1)
示例:
pe_id = shmem.my_pe()
n_pes
- n_pes() int32
获取系统中 PE 的总数。
- 返回:
int32 —— 处理单元总数
示例:
total_pes = shmem.n_pes()
team_my_pe
- team_my_pe(team: int32) int32
获取特定团队内的 PE 编号。
- 参数:
team -- 团队句柄(int32)
- 返回:
int32 —— 指定团队内的 PE 编号
team_n_pes
- team_n_pes(team: int32) int32
获取特定团队中的 PE 数量。
- 参数:
team -- 团队句柄(int32)
- 返回:
int32 —— 团队中的 PE 数量
team_translate_pe
- team_translate_pe(src_team: int32, pe_in_src_team: int32, dest_team: int32) int32
将 PE 编号从一个团队上下文转换到另一个团队上下文。
- 参数:
src_team -- 源团队句柄
pe_in_src_team -- 源团队中的 PE 编号
dest_team -- 目标团队句柄
- 返回:
int32 —— 目标团队中对应的 PE 编号
远程内存访问(RMA)
remote_ptr
- remote_ptr(local_ptr, pe) pointer
获取远程 PE 上对称内存对象的指针。
- 参数:
local_ptr -- 指向本地对称内存的指针
pe -- 目标 PE 编号(int32 或 uint32)
- 返回:
与 local_ptr 类型相同的指针,指向远程 PE 上的对称对象
支持所有数据类型。本地指针必须引用对称内存。
int_p
- int_p(dest, value: int32, pe: int32) void
向远程 PE 写入单个 int32 值(阻塞操作)。
- 参数:
dest -- 远程 PE 上的目标指针(必须是 int32*)
value -- 要写入的 int32 值
pe -- 目标 PE 编号
示例:
shmem.int_p(remote_flag, tl.int32(1), target_pe)
getmem
- getmem(dest, source, bytes: uint32, pe: int32) void
从远程内存阻塞获取(读取)操作。
- 参数:
dest -- 本地目标缓冲区指针
source -- 远程 PE 上的源指针
bytes -- 要传输的字节数(转换为 uint32)
pe -- 源 PE 编号
支持的类型: fp16、fp32、int8、int16、int32、int64、uint8、uint16、uint32、uint64、bf16
putmem
- putmem(dest, source, bytes: uint32, pe: int32) void
向远程内存阻塞写入(put)操作。
- 参数:
dest -- 远程 PE 上的目标指针
source -- 本地源缓冲区指针
bytes -- 要传输的字节数(转换为 uint32)
pe -- 目标 PE 编号
支持的类型: fp16、fp32、int8、int16、int32、int64、uint8、uint16、uint32、uint64、bf16
getmem_nbi
- getmem_nbi(dest, source, bytes: uint32, pe: int32, qp_id: int = 0) void
从远程内存非阻塞获取(读取)操作。
- 参数:
dest -- 本地目标缓冲区指针
source -- 远程 PE 上的源指针
bytes -- 要传输的字节数
pe -- 源 PE 编号
qp_id -- 队列对 ID(默认:0,当前未使用)
操作可能异步完成。使用
quiet()或fence()确保在访问结果前操作完成。
putmem_nbi
- putmem_nbi(dest, source, bytes: uint32, pe: int32, qp_id: int = 0) void
向远程内存非阻塞写入(put)操作。
- 参数:
dest -- 远程 PE 上的目标指针
source -- 本地源缓冲区指针
bytes -- 要传输的字节数
pe -- 目标 PE 编号
qp_id -- 队列对 ID(默认:0,当前未使用)
操作可能异步完成。使用
quiet()或fence()确保操作完成。
信号操作
putmem_signal
- putmem_signal(dest, source, nbytes: uint32, sig_addr, signal: int32, sig_op: int32, pe: int32) void
阻塞的 put 操作,完成时原子性更新信号。
- 参数:
dest -- 远程 PE 上的目标指针
source -- 本地源缓冲区指针
nbytes -- 要传输的字节数
sig_addr -- 远程 PE 上信号位置的指针(必须是 int32*)
signal -- 要应用的信号值
sig_op -- 信号操作(参见 ACLSHMEMSignalOp 枚举)
pe -- 目标 PE 编号
信号操作:
使用
triton_dist.language.extra.ascend.aclshmem_constants.ACLSHMEMSignalOp中的常量:SET—— 原子性将信号位置设置为信号值ADD—— 原子性将信号值加到现有值上
示例:
from triton_dist.language.extra.ascend.aclshmem_constants import ACLSHMEMSignalOp shmem.putmem_signal(remote_buf, local_buf, nbytes, signal_addr, tl.int32(1), ACLSHMEMSignalOp.SET, target_pe)
putmem_signal_nbi
- putmem_signal_nbi(dest, source, nbytes: uint32, sig_addr, signal: int32, sig_op: int32, pe: int32) void
非阻塞的 put 操作,完成时原子性更新信号。
- 参数:
dest -- 远程 PE 上的目标指针
source -- 本地源缓冲区指针
nbytes -- 要传输的字节数
sig_addr -- 远程 PE 上信号位置的指针(必须是 int32*)
signal -- 要应用的信号值
sig_op -- 信号操作(SET 或 ADD)
pe -- 目标 PE 编号
信号操作在数据传输完成时发生。使用
quiet()确保所有非阻塞操作已完成。
signal_op
- signal_op(sig_addr, signal: int32, sig_op: int32, pe: int32) void
在远程 PE 上执行原子信号操作,不传输数据。
- 参数:
sig_addr -- 远程 PE 上信号位置的指针(必须是 int32*)
signal -- 要应用的信号值
sig_op -- 信号操作(SET 或 ADD)
pe -- 目标 PE 编号
signal_wait_until
- signal_wait_until(sig_addr, cmp_: int32, cmp_val: int32) int32
等待直到信号位置满足比较条件。
- 参数:
sig_addr -- 本地信号位置的指针(必须是 int32*)
cmp -- 比较操作(参见 ACLSHMEMCmpOp 枚举)
cmp_val -- 要比较的值
- 返回:
int32 —— 满足条件的观察值
比较操作:
使用
triton_dist.language.extra.ascend.aclshmem_constants.ACLSHMEMCmpOp中的常量:EQ—— 等于NE—— 不等于GT—— 大于GE—— 大于或等于LT—— 小于LE—— 小于或等于
示例:
from triton_dist.language.extra.ascend.aclshmem_constants import ACLSHMEMCmpOp # 等待直到信号 >= 1 value = shmem.signal_wait_until(signal_ptr, ACLSHMEMCmpOp.GE, tl.int32(1))
同步
barrier_all
- barrier_all() void
系统中所有 PE 的全局屏障。
所有 PE 必须调用此操作。执行阻塞直到所有 PE 到达屏障。
barrier_all_vec
- barrier_all_vec() void
所有 PE 的向量化全局屏障(优化实现)。
昇腾特定的优化屏障实现。语义上等价于
barrier_all,但可能具有更好的性能。
barrier
- barrier(team: int32) void
团队特定的屏障同步。
- 参数:
team -- 团队句柄
团队中的所有 PE 必须调用此操作。执行阻塞直到所有团队成员到达屏障。
barrier_vec
- barrier_vec(team: int32) void
向量化团队屏障(优化实现)。
- 参数:
team -- 团队句柄
quiet
- quiet() void
等待本地 PE 发起的所有未完成 RMA 操作完成。
确保所有先前的非阻塞 put 和 get 操作已完成。不保证远程可见性(使用
fence()进行顺序控制)。
fence
- fence() void
内存屏障,确保 RMA 操作的顺序。
确保所有先前的 RMA 操作在后续 RMA 操作之前排序。提供内存顺序保证而不等待完成。
支持的数据类型
RMA 操作支持以下数据类型:
Triton 类型 |
C/Kernel 后缀 |
|---|---|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
常量
ACLSHMEMSignalOp
信号操作的信号操作类型:
ACLSHMEMCmpOp
signal_wait_until 的比较操作:
ACLSHMEMTeam
团队标识符:
使用示例
import triton
import triton.language as tl
from triton_dist.language.extra.ascend import libaclshmem_device as shmem
from triton_dist.language.extra.ascend.aclshmem_constants import ACLSHMEMSignalOp, ACLSHMEMCmpOp
@triton.jit
def distributed_kernel(data_ptr, signal_ptr, result_ptr, BLOCK_SIZE: tl.constexpr):
pe = shmem.my_pe()
n_pe = shmem.n_pes()
# 在开始时同步所有 PE
shmem.barrier_all()
# 每个 PE 处理本地数据
local_data = tl.load(data_ptr + pe * BLOCK_SIZE)
processed = local_data * 2
# 将结果发送到 PE 0
if pe > 0:
remote_ptr = shmem.remote_ptr(result_ptr, tl.int32(0))
shmem.putmem_signal(
remote_ptr + pe * BLOCK_SIZE,
tl.pointer(processed),
BLOCK_SIZE * 4, # 每个 float32 4 字节
signal_ptr,
tl.int32(1),
ACLSHMEMSignalOp.ADD,
tl.int32(0)
)
else:
# PE 0 等待所有其他 PE
shmem.signal_wait_until(signal_ptr, ACLSHMEMCmpOp.EQ, tl.int32(n_pe - 1))
# 最终屏障
shmem.barrier_all()
注意事项
对称内存要求
大多数 RMA 操作需要对称内存——在所有 PE 上以相同虚拟地址偏移量分配的内存。使用主机端的
aclshmem_malloc() 分配对称内存(参见 SHMEM 主机端 API)。
错误处理
sig_op 或 cmp_ 的无效枚举值将引发带有详细消息的 ValueError,列出有效选项。
性能考虑
非阻塞操作(
*_nbi)可通过重叠通信与计算来提高性能使用
barrier_vec和barrier_all_vec可能比标准屏障获得更好的性能在可能的情况下将多个小传输批处理为较大传输,以分摊通信开销
使用
extern_call基础设施从 Triton kernel 调用这些操作