Triton-distributed 语义
本文档描述 Triton-distributed-ascend 的设计理念和语义模型。
设计理念
Triton-distributed 建立在 以块为中心(Tile-Centric) 的设计理念之上(如 MLSys 2025 论文 所述),这意味着:
块作为工作单元:计算和通信都围绕块(数据块)组织。每个块是一个自包含的单元,可以独立地计算、传输和同步。
解耦的计算和通信:通信(数据传输)和计算(GEMM 等)被明确分离,可以由不同的 AI 核心或同一内核的不同部分执行。
细粒度重叠:通过围绕块组织工作,我们可以实现细粒度的重叠,其中对已经可用的块的计算可以在其他块仍在传输时进行。
核心语义概念
生产者-消费者模型
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 主机端 API 和 SHMEM 设备端 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] → ... │
└─────────────────────────────────────────────────────────────┘
生产者通过 AllGather 从其他 PE 传输块
每个块传输完成后,生产者向消费者发送信号
消费者等待每个块,然后立即开始 GEMM 计算
计算和通信重叠,隐藏通信延迟
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] → ... │
└─────────────────────────────────────────────────────────────┘
生产者(GEMM)计算输出块
每个块计算完成后,生产者向消费者发送信号
消费者(ReduceScatter)等待块并执行归约 + 分散
当计算产生结果时进行通信,最大化重叠
线程块交织
为了最大化重叠并最小化同步,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 之后执行
这对正确性至关重要,因为:
编译器会积极地重新排序指令以提高性能
wait 和 load 之间没有语法依赖关系
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()→ 使用 ACLSHMEMputmem_signal或signal_op屏障 → ACLSHMEM
barrier_all()或barrier()
有关调试同步问题,请参见 Kernel 调试。
最佳实践
最小化同步粒度:使用每块信号而不是全局屏障
积极重叠:一旦任何块准备好就开始计算,不要等待所有块
使用线程块交织:安排工作以便首先处理本地数据
尽可能批量信号:如果多个块总是一起访问,它们可以共享一个信号
考虑通信拓扑:在昇腾集群上,利用网络拓扑(NVLink、RoCE)进行高效的数据放置
性能分析和调优:使用性能分析工具识别重叠效率和同步瓶颈
高效使用 ACLSHMEM:
尽可能优先使用非阻塞操作(
*_nbi)明智地使用
quiet()或fence()来强制执行顺序考虑使用基于团队的操作进行子集通信
首先使用屏障同步进行测试:调试时,使用
--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()确保正确的加载顺序重叠通信(发送下一个块)与计算(归约当前块)
参考文献
triton_dist.language API — 分布式原语的 API 参考
SHMEM 设备端 API — ACLSHMEM 设备端操作
SHMEM 主机端 API — ACLSHMEM 主机端初始化和内存管理