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