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" —— 跨节点

平台特定要求(昇腾):

  • 信号指针必须是 int32int64uint64 类型

  • 对于 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_atnotify)需要对称内存——在每个 rank 的地址空间中以相同偏移量分配的内存。 确保使用对称分配原语分配内存(参见 SHMEM 主机端 API)。

平台支持

这些 API 专为昇腾 NPU 设计。如各 API 描述中所述,某些操作具有昇腾特定的行为或约束。