Triton-distributed 语义 ============================ 本文档描述 Triton-distributed-ascend 的设计理念和语义模型。 设计理念 ----------------- Triton-distributed 建立在 **以块为中心(Tile-Centric)** 的设计理念之上(如 `MLSys 2025 论文 `_ 所述),这意味着: 1. **块作为工作单元**:计算和通信都围绕块(数据块)组织。每个块是一个自包含的单元,可以独立地计算、传输和同步。 2. **解耦的计算和通信**:通信(数据传输)和计算(GEMM 等)被明确分离,可以由不同的 AI 核心或同一内核的不同部分执行。 3. **细粒度重叠**:通过围绕块组织工作,我们可以实现细粒度的重叠,其中对已经可用的块的计算可以在其他块仍在传输时进行。 核心语义概念 ---------------------- 生产者-消费者模型 ~~~~~~~~~~~~~~~~~~~~~~~ Triton-distributed 使用 **生产者-消费者** 模型来实现计算与通信的重叠: - **生产者**:负责数据传输(例如 AllGather、All-to-All)。生产者将数据复制到共享缓冲区,并在每个块准备好时发出信号。 - **消费者**:负责计算(例如 GEMM)。消费者等待块准备好,然后立即对其进行计算。 .. code-block:: python # 生产者:传输数据并发送信号 @triton.jit def producer_kernel(...): # 将块传输到共享缓冲区 tl.store(remote_buffer_ptr + tile_offset, data) # 发出块已准备好的信号 tdl.notify(signal_ptr, peer_rank, signal=tile_id, sig_op="set") # 消费者:等待数据并计算 @triton.jit def consumer_kernel(...): # 等待块准备好 token = tdl.wait(signal_ptr + tile_id, 1) # 消费令牌以建立数据依赖 data_ptr = tdl.consume_token(data_ptr, token) # 现在可以安全地对块进行计算 result = compute(tl.load(data_ptr)) 基于信号的同步 ~~~~~~~~~~~~~~~~~~~~~~~~~~~~ 同步模型基于 **信号** 而不是屏障: - ``wait(ptr, n)``:等待直到 ptr 处的 n 个信号达到预期值 - ``notify(ptr, rank, signal, sig_op)``:向特定秩发送信号 - ``consume_token(value, token)``:在等待和内存访问之间建立数据依赖 **关键见解**:与同步所有秩的全局屏障不同,基于信号的同步允许细粒度的块级协调。这使得: - 不同的块可以独立进行 - 通信和计算之间的最大重叠 - 避免"落后者效应",即慢速秩阻塞所有人 对称内存模型 ~~~~~~~~~~~~~~~~~~~~~~ Triton-distributed-ascend 使用 **对称内存**(ACLSHMEM)进行跨秩通信: - 所有处理单元(Processing Elements,PE)在相同的虚拟地址分配内存 - ``symm_at(ptr, rank)``:将本地指针映射到另一个 PE 上的相应地址 - 支持跨 PE 的直接加载/存储,无需显式发送/接收 .. code-block:: python # 直接访问远程 PE 的内存 remote_ptr = tdl.symm_at(local_ptr, peer_pe) tl.store(remote_ptr + offset, data) # 写入对等方的内存 注意:在昇腾 NPU 上,对称内存由 ACLSHMEM 提供(参见 :doc:`shmem_host` 和 :doc:`shmem_device`)。 内核设计模式 ---------------------- AllGather + GEMM 重叠 ~~~~~~~~~~~~~~~~~~~~~~~~~~~~ .. code-block:: text Time → ┌─────────────────────────────────────────────────────────────┐ │ Producer (AllGather) │ │ [Tile 0] → [Tile 1] → [Tile 2] → [Tile 3] → ... │ │ ↓ ↓ ↓ ↓ │ │ signal signal signal signal │ │ ↓ ↓ ↓ ↓ │ │ Consumer (GEMM) │ │ [Tile 0] → [Tile 1] → [Tile 2] → [Tile 3] → ... │ └─────────────────────────────────────────────────────────────┘ 1. 生产者通过 AllGather 从其他 PE 传输块 2. 每个块传输完成后,生产者向消费者发送信号 3. 消费者等待每个块,然后立即开始 GEMM 计算 4. 计算和通信重叠,隐藏通信延迟 GEMM + ReduceScatter 重叠 ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ .. code-block:: text Time → ┌─────────────────────────────────────────────────────────────┐ │ Producer (GEMM) │ │ [Tile 0] → [Tile 1] → [Tile 2] → [Tile 3] → ... │ │ ↓ ↓ ↓ ↓ │ │ signal signal signal signal │ │ ↓ ↓ ↓ ↓ │ │ Consumer (ReduceScatter) │ │ [Tile 0] → [Tile 1] → [Tile 2] → [Tile 3] → ... │ └─────────────────────────────────────────────────────────────┘ 1. 生产者(GEMM)计算输出块 2. 每个块计算完成后,生产者向消费者发送信号 3. 消费者(ReduceScatter)等待块并执行归约 + 分散 4. 当计算产生结果时进行通信,最大化重叠 线程块交织 ~~~~~~~~~~~~~~~~~~~~~ 为了最大化重叠并最小化同步,Triton-distributed 使用 **线程块交织(threadblock swizzling)**: .. code-block:: python # 交织块分配,使每个 PE 从其本地数据开始 pid_m = (pid_m + rank * tiles_per_pe) % total_tiles 这确保了: - 每个 PE 从本地可用的数据开始计算(无需等待) - 当 PE 需要远程数据时,很可能已经传输完成 - 减少同步开销并提高缓存效率 基于令牌的数据依赖 ~~~~~~~~~~~~~~~~~~~~~~~~~~~ ``consume_token`` 原语建立显式的数据依赖关系: .. code-block:: python # 没有 consume_token:编译器可能在 wait 之前重新排序 load token = tdl.wait(signal_ptr, 1) data = tl.load(data_ptr) # 错误:可能在 wait 之前执行! # 使用 consume_token:显式依赖 token = tdl.wait(signal_ptr, 1) data_ptr = tdl.consume_token(data_ptr, token) # 建立依赖关系 data = tl.load(data_ptr) # 保证在 wait 之后执行 这对正确性至关重要,因为: 1. 编译器会积极地重新排序指令以提高性能 2. wait 和 load 之间没有语法依赖关系 3. ``consume_token`` 创建显式的数据依赖关系以防止重新排序 昇腾特定注意事项 ------------------------------- 内存层次结构 ~~~~~~~~~~~~~~~~ 昇腾 NPU 具有针对 AI 工作负载优化的独特内存层次结构: - **全局内存(GM/HBM)**:主设备内存,具有高带宽 - **UB(统一缓冲区)**:向量核心(AIV)使用的本地暂存/寄存器类内存,用于逐元素、归约和向量操作。大小:Atlas 800T/I A2 系列为 192 KB - **L0 缓冲区**:专用于立方核心(AIC)进行矩阵乘法的专用超快速输入和累加缓冲区: - **L0A**:矩阵乘法的输入缓冲区 A - **L0B**:矩阵乘法的输入缓冲区 B - **L0C**:矩阵乘法结果的累加缓冲区 - **对称内存**:通过 ACLSHMEM 从全局内存分配,可跨 PE 访问以进行分布式通信 **架构概述:** 每个 AI 核心由一个立方核心(AIC)和两个向量核心(AIV)组成: - **立方核心(AIC)**:使用 L0A、L0B、L0C 缓冲区执行 ``tl.dot`` 操作 - **向量核心(AIV)**:使用 UB 执行逐元素操作、归约和聚集/分散 当使用对称内存进行分布式操作时: - 在主机上使用 ``aclshmem_malloc()`` 分配缓冲区(参见 :doc:`shmem_host`) - 在内核中使用 ``remote_ptr()`` 访问远程内存(参见 :doc:`shmem_device`) - 使用 ``triton_dist.language`` 中的 ``symm_at()`` 进行指针转换 **内存约束:** - UB 大小有限(A2 系列为 192 KB,启用双缓冲时进一步减少) - 张量尾轴必须对齐:向量操作为 32 字节,立方+向量操作为 512 字节 - 使用分块以适应 UB 约束内的数据 同步实现 ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ 在昇腾 NPU 上,同步原语映射到 ACLSHMEM 操作: - ``wait()`` → 使用 ACLSHMEM 信号等待机制 - ``notify()`` → 使用 ACLSHMEM ``putmem_signal`` 或 ``signal_op`` - 屏障 → ACLSHMEM ``barrier_all()`` 或 ``barrier()`` 有关调试同步问题,请参见 :doc:`../developer-guide/kernel_debugging`。 最佳实践 -------------- 1. **最小化同步粒度**:使用每块信号而不是全局屏障 2. **积极重叠**:一旦任何块准备好就开始计算,不要等待所有块 3. **使用线程块交织**:安排工作以便首先处理本地数据 4. **尽可能批量信号**:如果多个块总是一起访问,它们可以共享一个信号 5. **考虑通信拓扑**:在昇腾集群上,利用网络拓扑(NVLink、RoCE)进行高效的数据放置 6. **性能分析和调优**:使用性能分析工具识别重叠效率和同步瓶颈 7. **高效使用 ACLSHMEM**: - 尽可能优先使用非阻塞操作(``*_nbi``) - 明智地使用 ``quiet()`` 或 ``fence()`` 来强制执行顺序 - 考虑使用基于团队的操作进行子集通信 8. **首先使用屏障同步进行测试**:调试时,使用 ``--enable-hivm-inject-barrier-all-sync=true`` 来排除细粒度同步错误(参见 :doc:`../developer-guide/kernel_debugging`) 示例:带重叠的环形 AllReduce ------------------------------------- 以下是一个概念性示例,说明以块为中心的语义如何实现高效的环形 AllReduce: .. code-block:: python import triton import triton.language as tl import triton_dist.language as tdl @triton.jit def ring_allreduce_kernel( local_data_ptr, remote_data_ptr, signal_ptr, rank: tl.constexpr, world_size: tl.constexpr, num_tiles: tl.constexpr, ): pid = tl.program_id(0) # 阶段 1:Reduce-Scatter(发送到下一个,从上一个接收) for step in range(world_size - 1): send_rank = (rank + 1) % world_size recv_rank = (rank - 1 + world_size) % world_size # 确定在此步骤中发送/接收哪个块 tile_id = (rank - step + world_size) % world_size if pid == tile_id: # 将块发送到下一个秩 remote_ptr = tdl.symm_at(remote_data_ptr, send_rank) data = tl.load(local_data_ptr + tile_id * TILE_SIZE) tl.store(remote_ptr + tile_id * TILE_SIZE, data) tdl.notify(signal_ptr, send_rank, signal=step, sig_op="set") # 等待来自上一个秩的传入块 token = tdl.wait(signal_ptr + step, 1) recv_ptr = tdl.consume_token(remote_data_ptr, token) incoming_data = tl.load(recv_ptr + tile_id * TILE_SIZE) # 与本地数据归约 local_data = tl.load(local_data_ptr + tile_id * TILE_SIZE) reduced = local_data + incoming_data tl.store(local_data_ptr + tile_id * TILE_SIZE, reduced) # 阶段 2:AllGather(传播归约后的块) for step in range(world_size - 1): send_rank = (rank + 1) % world_size recv_rank = (rank - 1 + world_size) % world_size # 确定在此步骤中发送哪个块 # 每个秩发送它刚刚完成归约的块 send_tile_id = (rank - step + world_size) % world_size # 确定在此步骤中接收哪个块 recv_tile_id = (rank - step - 1 + world_size) % world_size if pid == send_tile_id: # 将归约后的块发送到下一个秩 remote_ptr = tdl.symm_at(remote_data_ptr, send_rank) data = tl.load(local_data_ptr + send_tile_id * TILE_SIZE) tl.store(remote_ptr + send_tile_id * TILE_SIZE, data) tdl.notify(signal_ptr, send_rank, signal=step + world_size - 1, sig_op="set") # 等待来自上一个秩的传入块 token = tdl.wait(signal_ptr + step + world_size - 1, 1) recv_ptr = tdl.consume_token(remote_data_ptr, token) incoming_data = tl.load(recv_ptr + recv_tile_id * TILE_SIZE) # 存储接收到的归约块 tl.store(local_data_ptr + recv_tile_id * TILE_SIZE, incoming_data) **关键点**: - 每个块使用 ``wait()``/``notify()`` 独立同步 - 块一旦到达就可以进行归约(无需屏障) - 使用 ``consume_token()`` 确保正确的加载顺序 - 重叠通信(发送下一个块)与计算(归约当前块) 参考文献 ---------- - `TileLink: Generating Efficient Compute-Communication Overlapping Kernels using Tile-Centric Primitives (MLSys 2025) `_ - `Triton-distributed: Programming Overlapping Kernels on Distributed AI Systems with the Triton Compiler `_ - :doc:`triton_dist_language` — 分布式原语的 API 参考 - :doc:`shmem_device` — ACLSHMEM 设备端操作 - :doc:`shmem_host` — ACLSHMEM 主机端初始化和内存管理