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