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 后缀

fp16

half

fp32

float

bf16

bfloat16

int8

int8

int16

int16

int32

int32

int64

int64

uint8

uint8

uint16

uint16

uint32

uint32

uint64

uint64

常量

ACLSHMEMSignalOp

信号操作的信号操作类型:

class ACLSHMEMSignalOp
SET = 0

将信号设置为指定值

ADD = 1

将值加到信号上

ACLSHMEMCmpOp

signal_wait_until 的比较操作:

class ACLSHMEMCmpOp
EQ = 0

等于

NE = 1

不等于

GT = 2

大于

GE = 3

大于或等于

LT = 4

小于

LE = 5

小于或等于

ACLSHMEMTeam

团队标识符:

class ACLSHMEMTeam
INVALID = -1

无效团队

WORLD = 0

全局团队(所有 PE)

使用示例

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_opcmp_ 的无效枚举值将引发带有详细消息的 ValueError,列出有效选项。

性能考虑

  • 非阻塞操作(*_nbi)可通过重叠通信与计算来提高性能

  • 使用 barrier_vecbarrier_all_vec 可能比标准屏障获得更好的性能

  • 在可能的情况下将多个小传输批处理为较大传输,以分摊通信开销

  • 使用 extern_call 基础设施从 Triton kernel 调用这些操作