Triton-distributed 语义

本文档描述 Triton-distributed-ascend 的设计理念和语义模型。

设计理念

Triton-distributed 建立在 以块为中心(Tile-Centric) 的设计理念之上(如 MLSys 2025 论文 所述),这意味着:

  1. 块作为工作单元:计算和通信都围绕块(数据块)组织。每个块是一个自包含的单元,可以独立地计算、传输和同步。

  2. 解耦的计算和通信:通信(数据传输)和计算(GEMM 等)被明确分离,可以由不同的 AI 核心或同一内核的不同部分执行。

  3. 细粒度重叠:通过围绕块组织工作,我们可以实现细粒度的重叠,其中对已经可用的块的计算可以在其他块仍在传输时进行。

核心语义概念

生产者-消费者模型

Triton-distributed 使用 生产者-消费者 模型来实现计算与通信的重叠:

  • 生产者:负责数据传输(例如 AllGather、All-to-All)。生产者将数据复制到共享缓冲区,并在每个块准备好时发出信号。

  • 消费者:负责计算(例如 GEMM)。消费者等待块准备好,然后立即对其进行计算。

# 生产者:传输数据并发送信号
@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 的直接加载/存储,无需显式发送/接收

# 直接访问远程 PE 的内存
remote_ptr = tdl.symm_at(local_ptr, peer_pe)
tl.store(remote_ptr + offset, data)  # 写入对等方的内存

注意:在昇腾 NPU 上,对称内存由 ACLSHMEM 提供(参见 SHMEM 主机端 APISHMEM 设备端 API)。

内核设计模式

AllGather + GEMM 重叠

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 重叠

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)

# 交织块分配,使每个 PE 从其本地数据开始
pid_m = (pid_m + rank * tiles_per_pe) % total_tiles

这确保了:

  • 每个 PE 从本地可用的数据开始计算(无需等待)

  • 当 PE 需要远程数据时,很可能已经传输完成

  • 减少同步开销并提高缓存效率

基于令牌的数据依赖

consume_token 原语建立显式的数据依赖关系:

# 没有 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() 分配缓冲区(参见 SHMEM 主机端 API

  • 在内核中使用 remote_ptr() 访问远程内存(参见 SHMEM 设备端 API

  • 使用 triton_dist.language 中的 symm_at() 进行指针转换

内存约束:

  • UB 大小有限(A2 系列为 192 KB,启用双缓冲时进一步减少)

  • 张量尾轴必须对齐:向量操作为 32 字节,立方+向量操作为 512 字节

  • 使用分块以适应 UB 约束内的数据

同步实现

在昇腾 NPU 上,同步原语映射到 ACLSHMEM 操作:

  • wait() → 使用 ACLSHMEM 信号等待机制

  • notify() → 使用 ACLSHMEM putmem_signalsignal_op

  • 屏障 → ACLSHMEM barrier_all()barrier()

有关调试同步问题,请参见 Kernel 调试

最佳实践

  1. 最小化同步粒度:使用每块信号而不是全局屏障

  2. 积极重叠:一旦任何块准备好就开始计算,不要等待所有块

  3. 使用线程块交织:安排工作以便首先处理本地数据

  4. 尽可能批量信号:如果多个块总是一起访问,它们可以共享一个信号

  5. 考虑通信拓扑:在昇腾集群上,利用网络拓扑(NVLink、RoCE)进行高效的数据放置

  6. 性能分析和调优:使用性能分析工具识别重叠效率和同步瓶颈

  7. 高效使用 ACLSHMEM

    • 尽可能优先使用非阻塞操作(*_nbi

    • 明智地使用 quiet()fence() 来强制执行顺序

    • 考虑使用基于团队的操作进行子集通信

  8. 首先使用屏障同步进行测试:调试时,使用 --enable-hivm-inject-barrier-all-sync=true 来排除细粒度同步错误(参见 Kernel 调试

示例:带重叠的环形 AllReduce

以下是一个概念性示例,说明以块为中心的语义如何实现高效的环形 AllReduce:

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() 确保正确的加载顺序

  • 重叠通信(发送下一个块)与计算(归约当前块)

参考文献