triton_dist.language API
本模块为运行在昇腾 NPU 上的 Triton kernel 提供分布式计算原语,支持分布式系统中多个处理单元间的通信与同步。
所有 API 可通过 triton_dist.language 或简写 tdl 访问:
import triton_dist.language as tdl
关于原始 Triton-distributed 项目的更多背景信息,请参考 triton-distributed 文档。
分布式操作
wait
- wait(barrierPtrs, numBarriers, scope: str, semantic: str, waitValue: int = 1) tensor
等待屏障指针直到满足指定条件。
- 参数:
barrierPtrs -- 指向内存中屏障位置的指针或指针张量
numBarriers -- 需要等待的屏障数量
scope -- 内存作用域(注意:triton-distributed-ascend 目前不支持作用域变体)
semantic -- 内存语义/顺序(例如
"acquire"、"relaxed")waitValue -- 在屏障位置等待的值(默认:1)
- 返回:
表示完成状态的 int32 张量
示例:
# 使用 acquire 语义等待屏障 status = tdl.wait(barrier_ptr, num_barriers=1, scope="", semantic="acquire", waitValue=1)
源文件:
python/triton_dist/language/distributed_ops.py:58-85
consume_token
- consume_token(value, token) tensor or tensor_descriptor
消费一个 token 并将其与值关联,建立顺序依赖关系。
- 参数:
value -- 要与 token 关联的张量或张量描述符
token -- 整数 token 值(必须是整数类型)
- 返回:
附加了 token 依赖的输入值(类型与输入相同)
此操作用于在分布式流水线中建立顺序约束。返回值保持与输入相同的形状和类型。
源文件:
python/triton_dist/language/distributed_ops.py:89-97
rank
- rank(axis: int = -1) tensor
获取当前处理单元的 rank(ID)。
- 参数:
axis -- 查询 rank 的轴(默认:-1 表示全局 rank)
- 返回:
包含当前 rank 的 int32 张量
示例:
my_rank = tdl.rank() # 获取全局 rank rank_on_axis0 = tdl.rank(axis=0) # 获取特定轴上的 rank
源文件:
python/triton_dist/language/distributed_ops.py:101-103
num_ranks
- num_ranks(axis: int = -1) tensor
获取系统中的总 rank 数(处理单元数)。
- 参数:
axis -- 查询 rank 数量的轴(默认:-1 表示全局)
- 返回:
包含总 rank 数的 int32 张量
示例:
total_ranks = tdl.num_ranks() # 获取总 rank 数 ranks_on_axis0 = tdl.num_ranks(axis=0) # 获取特定轴上的 rank 数
源文件:
python/triton_dist/language/distributed_ops.py:107-109
symm_at
- symm_at(ptr, rank) tensor
获取指定远程 rank 上指针的对称地址。
- 参数:
ptr -- 指向对称内存的标量指针(必须是标量指针,而非块)
rank -- 要获取其对称地址的目标 rank
- 返回:
包含指定 rank 上对称指针的张量(类型与输入指针相同)
仅支持标量指针。要求内存在所有 rank 上被分配为对称内存。
示例:
# 获取 rank 1 上的对称地址 remote_ptr = tdl.symm_at(local_ptr, rank=1)
源文件:
python/triton_dist/language/distributed_ops.py:113-116
notify
- notify(ptr, rank, signal: int = 1, sig_op: str = 'set', comm_scope: str = 'inter_node') tensor
通过更新远程内存位置向远程 rank 发送通知信号。
- 参数:
ptr -- 指向信号位置的标量指针(必须是标量指针)
rank -- 要通知的目标 rank
signal -- 要发送的信号值(默认:1)
sig_op -- 信号操作类型(默认:
"set")comm_scope -- 通信作用域(默认:
"inter_node")
- 返回:
空张量(操作有副作用但无返回值)
信号操作:
"set"—— 设置目标位置的值"add"—— 将值加到目标位置
通信作用域:
"intra_node"—— 单节点内"inter_node"—— 跨节点
平台特定要求(昇腾):
信号指针必须是
int32、int64或uint64类型对于
int64/uint64类型,仅支持sig_op="set"对于
int32类型,支持"set"和"add"两种操作
示例:
# 通过设置信号为 1 通知 rank 1 tdl.notify(signal_ptr, rank=1, signal=1, sig_op="set", comm_scope="inter_node") # 通过递增信号通知 rank 2 tdl.notify(signal_ptr, rank=2, signal=1, sig_op="add", comm_scope="intra_node")
源文件:
python/triton_dist/language/distributed_ops.py:120-146
外部库调用
extern_call
- extern_call(lib_name: str, lib_path: str, args: list, arg_type_symbol_dict: dict, is_pure: bool) tensor
在 Triton kernel 内部调用外部库函数。
- 参数:
lib_name -- 外部库名称
lib_path -- 库的文件系统路径
args -- 传递给函数的参数列表
arg_type_symbol_dict -- 将参数类型元组映射到 (函数名, 返回类型) 对的字典
is_pure -- 函数是否为纯函数(无副作用)
- 返回:
包含外部函数返回值的张量
所有参数必须是标量类型(不支持块/张量参数)。函数根据参数类型分发到相应的外部符号。 此功能在设备端 SHMEM 操作中内部使用。
示例:
result = tdl.extern_call( "libshmem_device", "", [ptr, value, pe], { (tl.pointer_type(tl.int32), tl.int32, tl.int32): ("aclshmem_int32_p", ()), }, is_pure=False )
源文件:
python/triton_dist/language/core.py:100-146
注意事项
内存顺序与同步
本模块中的分布式操作与硬件内存层次结构和互连交互。使用屏障操作、等待原语和通知时:
根据内存一致性需求选择适当的
semantic参数需要观察其他线程的写操作时使用
acquire语义向其他线程发布数据时使用
release语义
对称内存
若干操作(symm_at、notify)需要对称内存——在每个 rank 的地址空间中以相同偏移量分配的内存。
确保使用对称分配原语分配内存(参见 SHMEM 主机端 API)。
平台支持
这些 API 专为昇腾 NPU 设计。如各 API 描述中所述,某些操作具有昇腾特定的行为或约束。