Reverse All2All (EP MoE Distributed Kernel)
Reverse All2All kernel 实现 EP(Expert Parallel)MoE 模型中反向方向的 All2All 通信模式。在标准 MoE All2All 中,token 从其归属 rank 被分发到持有其指定 expert 的 rank(forward All2All)。Reverse All2All 是后半段:expert 计算之后,结果必须从 expert 的归属 rank 送回 token 的原始 rank。
本文档包含两个实现版本:
MTE 版本(通用,见「MTE 实现」一节):kernel 在 Vector Core 上采用 Producer-Consumer 流水线模式,偶数 Vector Core 充当 Producer(跨卡远程 store,通过
tl.store+dl.symm_at对称地址映射),奇数 Vector Core 充当 Consumer(本地读取与输出写入)。UDMA 优化版本(Ascend 950,见「UDMA 优化实现」一节):利用 950 上新增的 UDMA 链路提供更高跨卡带宽,把"跨卡数据交换"与"本地数据重排"拆成两个阶段分别处理。
Reverse All2All 算子语义
在 world_size 个 rank 的 EP MoE 模型中,每个 rank 持有一部分 expert。通信模式为:
Forward All2All (dispatch):
Input: (S, H, D) per rank # 按目标 expert rank 分组的 token
Output: (S/world_size, H * world_size, D) per rank # 来自所有 rank 的 token
Expert Computation:
每个 rank 处理分配给其 expert 的 token
Reverse All2All (combine):
Input: (S * world_size, H, D) per rank # 本地 expert 的输出
Output: (S, H * world_size, D) per rank # 返回原始 token rank 的输出
Reverse All2All kernel 专门处理最后一步。输入 tensor shape 为 (S * world_size, H, D)——其中包含需要送回所有其他 rank 的 expert 输出。reverse All2All 之后,每个 rank 收到来自所有 rank 的输出,重组出 shape 为 (S, H * world_size, D) 的完整输出。
与 forward All2All 的关键区别在于维度交换:world_size 因子从 H 维移到 S 维(或反之),实质上是对跨 rank 通信模式做转置。
EP MoE 场景
在典型 EP MoE 层(如 256 expert 的 DeepSeek-V3)中:
Forward All2All (Dispatch):token 被发送到持有其指定 expert 的 rank。每 rank 的输入
(S, H, D)变为(S/world_size, H * world_size, D)。Expert GEMM:每个 rank 用 GroupedGEMM 或 Fused MoE 为本地 expert 计算 GEMM。
Reverse All2All (Combine):expert 输出送回原始 token rank。每 rank 的输入
(S * world_size, H, D)变为(S, H * world_size, D)。
Reverse All2All kernel 是第三步。基于 per-task 信号同步的 Producer-Consumer 流水线实现了跨卡远程 store(Producer)与本地读取/输出写入(Consumer)的重叠,获得比基于 barrier 方式更高的带宽利用率。
MTE 实现
Producer-Consumer流水
Reverse All2All kernel 使用基于 program_id 的 Vector Core 级 Producer-Consumer 流水线,kernel 以 Vector Core 粒度(而非 AICore 粒度)launch:
Even Vector Cores (role = pid % 2 == 0): Producer
- 读取本地输入 tensor 数据
- 通过 dl.symm_at 将数据写入远端对称内存
- 用 dl.wait (waitValue=0, acquire) 在远程写入前等待信号
- 用 dl.consume_token 强制依赖生效
- 所有远程写入完成后调用 libshmem_device.fence()
- 向每个目标 rank 发送 dl.notify 以通知完成
Odd Vector Cores (role = pid % 2 == 1): Consumer
- 用 dl.wait (acquire) 等待到达的数据信号
- 用 dl.consume_token 强制依赖生效
- 从本地对称内存读取数据
- 将 gather 的数据写入输出 tensor
- 所有读取完成后调用 libshmem_device.fence()
- 向自身发送 dl.notify (rank, signal=0) 以重置信号,供下一迭代使用
这与之前的 sub_vec_id() sub-block 模型不同。在 Vector Core PID 模型中:
logical_core_id =
pid // 2—— 每个 Producer-Consumer 对共享一个 logical_core_idn_role_cores =
ncore // 2—— 分配给每个角色的 core 数偶数 pid 负责所有跨卡远程 store 操作(Producer)
奇数 pid 负责所有本地对称内存读取与输出写入操作(Consumer)
kernel 以 Vector Core 总数 而非 AICore 数 launch:
from triton.backends.ascend.driver import NPUUtils
vec_num = NPUUtils().get_aivector_core_num() # Vector Core 数(非 AICore 数)
kernel_hccl_reverse_a2a_pipelined[vec_num, 1, 1](...)
这使每个 Vector Core 可独立地充当 Producer 或 Consumer,最大化通信操作的并行度。
Signal-Based Synchronization: dl.wait / dl.notify / fence(基于信号的同步)
与基于 barrier 的同步(barrier_all() / barrier_all_vec())不同,Reverse All2All kernel 使用 细粒度 per-task 信号同步:
Producer 端:
dl.wait(remote_signal_ptr, 1, "gpu", "acquire", waitValue=0)—— 等待目标 rank 的信号槽被清零(值 == 0),表示目标 rank 的 consumer 已读完上一个 bufferdl.consume_token(remote_ptr, token)—— 强制依赖生效,使远程写入仅在 wait 完成后发生tl.store(remote_ptr + offset, data, mask=mask)—— 将数据写入远端对称内存libshmem_device.fence()—— 确保所有远程写入在发送通知前可见dl.notify(signal_mem_ptr, target_rank)—— 通知目标 rank 数据可用
Consumer 端:
dl.wait(signal_mem_ptr, 1, "gpu", "acquire")—— 等待来自源 rank 的信号(数据可用)dl.consume_token(local_peer_ptr, token)—— 强制依赖生效,使本地读取仅在 wait 完成后发生tl.load(local_peer_ptr + offset, mask=mask)—— 从本地对称内存读取数据tl.store(output_ptr + offset, data, mask=mask)—— 写入输出 tensorlibshmem_device.fence()—— 确保所有读取在重置信号前完成dl.notify(signal_mem_ptr, rank, 0)—— 重置信号(值=0)以放行下一迭代的 producer
这种 per-task 信号机制比粗粒度 barrier 方式实现更细粒度的 Producer 与 Consumer 重叠。每个 Producer task 一旦自己专属的信号槽被清零即可推进,无需等待所有 task 完成全局 barrier。
Signal Memory Layout(信号内存布局)
信号内存组织为 per-task 的槽位数组:
signal_mem shape: [H * world_size * buffer_num * 8] (int64, 每槽 8 字节)
Layout: [buffer_id][num_blocks_d][H * rank_size][rank_factor]
- buffer_id: 标识双缓冲槽
- num_blocks_d * H * rank_size: 每个 buffer 迭代的总 task 数
- rank_factor: 标识 producer/consumer 对(producer 用 rank,consumer 用 r)
Producer 使用: remote_signal_ptr + buffer_id * num_blocks_d * H * rank_size * 8
+ num_blocks_d * H * rank * 8 + task_idx * 8
Consumer 使用: signal_mem_ptr + buffer_id * num_blocks_d * H * rank_size * 8
+ num_blocks_d * H * r * 8 + task_idx * 8
Triton-Distributed Reverse All2All 算子
import triton
import triton.language as tl
import triton_dist.language as dl
from triton_dist.language.extra import libshmem_device
from triton.backends.ascend.driver import NPUUtils
@triton.jit
def kernel_hccl_reverse_a2a_pipelined(
a_ptr, # 输入 tensor: shape (S_total, H, D)
c_ptr, # 输出 tensor: shape (S, H * world_size, D)
peer_mem_ptr, # 对称内存指针
signal_mem_ptr, # 信号内存指针 (int64)
rank, # 当前 rank ID(host 侧)
rank_size, # rank 总数(host 侧)
buffer_num, # 双缓冲数量
S, H, D, # tensor 维度 (S = S_total / rank_size)
stride_as, stride_ah, stride_ad, # 输入 strides
stride_cs, stride_ch, stride_cd, # 输出 strides
COMM_BLOCK_S: tl.constexpr, # S 维 tile 大小
COMM_BLOCK_D: tl.constexpr, # D 维 tile 大小
):
ncore = tl.num_programs(axis=0)
pid = tl.program_id(axis=0)
# Producer(偶数 PID)/ Consumer(奇数 PID)角色分配
role = pid % 2
logical_core_id = pid // 2
n_role_cores = ncore // 2
num_blocks_s = tl.cdiv(S, COMM_BLOCK_S)
num_blocks_d = tl.cdiv(D, COMM_BLOCK_D)
# 对称内存布局: [buffer_num, S_block, H * rank_size, D]
buffer_chunk_size = (H * rank_size) * D
stride_ps = H * rank_size * D
stride_ph = D
stride_pd = 1
# 外层循环: 遍历 S 维的 block
for global_id_s in range(0, num_blocks_s):
buffer_id = global_id_s % buffer_num
# ---- Producer Stage(偶数 Vector Core): 跨卡远程 store ----
if role == 0:
total_prod_tasks = num_blocks_d * H * rank_size
for task_idx in range(logical_core_id, total_prod_tasks, n_role_cores):
tmp = task_idx
rank_loop_id = tmp % rank_size
tmp //= rank_size
h_id = tmp % H
block_id_d = tmp // H
target_rank = (rank + rank_loop_id) % rank_size
# 非对齐 S 与 D 维的边界 mask
offs_s = global_id_s * COMM_BLOCK_S + tl.arange(0, COMM_BLOCK_S)
offs_d = block_id_d * COMM_BLOCK_D + tl.arange(0, COMM_BLOCK_D)
mask = (offs_s < S)[:, None] & (offs_d < D)[None, :]
# 从本地 A tensor 读取(行 offset = target_rank * S)
a_s = target_rank * S + offs_s
a_offs = a_s[:, None] * stride_as + h_id * stride_ah + offs_d[None, :] * stride_ad
a_data = tl.load(a_ptr + a_offs, mask=mask, other=0.0)
# 等待信号: 目标 rank 的 consumer 已完成上一次读取
remote_ptr = dl.symm_at(peer_mem_ptr, target_rank)
remote_signal_ptr = dl.symm_at(signal_mem_ptr, target_rank)
token = dl.wait(
remote_signal_ptr
+ buffer_id * num_blocks_d * H * rank_size * 8
+ num_blocks_d * H * rank * 8
+ tmp * 8,
1, "gpu", "acquire", waitValue=0,
)
remote_ptr_dummy = dl.consume_token(remote_ptr, token)
# 写入远端对称内存
peer_h_write = rank * H + h_id
peer_offs_write = (
buffer_id * (COMM_BLOCK_S * buffer_chunk_size)
+ tl.arange(0, COMM_BLOCK_S)[:, None] * stride_ps
+ peer_h_write * stride_ph
+ offs_d[None, :] * stride_pd
)
tl.store(remote_ptr_dummy + peer_offs_write, a_data, mask=mask)
# 在 notify 前确保所有远程写入可见
libshmem_device.fence()
for task_idx in range(logical_core_id, total_prod_tasks, n_role_cores):
tmp = task_idx
rank_loop_id = tmp % rank_size
tmp //= rank_size
target_rank = (rank + rank_loop_id) % rank_size
dl.notify(
signal_mem_ptr
+ buffer_id * num_blocks_d * H * rank_size * 8
+ num_blocks_d * H * rank * 8
+ tmp * 8,
target_rank,
)
# ---- Consumer Stage(奇数 Vector Core): 从本地 Shmem 读取并 store 到 C ----
if role == 1:
local_peer_ptr = dl.symm_at(peer_mem_ptr, rank)
total_cons_tasks = num_blocks_d * H * rank_size
for task_idx in range(logical_core_id, total_cons_tasks, n_role_cores):
tmp = task_idx
r = tmp % rank_size
tmp //= rank_size
h_id = tmp % H
block_id_d = tmp // H
# 边界 mask
offs_s = global_id_s * COMM_BLOCK_S + tl.arange(0, COMM_BLOCK_S)
offs_d = block_id_d * COMM_BLOCK_D + tl.arange(0, COMM_BLOCK_D)
mask = (offs_s < S)[:, None] & (offs_d < D)[None, :]
# 等待信号: 源 rank 的 producer 已完成写入
peer_h_read = r * H + h_id
peer_offs_read = (
buffer_id * (COMM_BLOCK_S * buffer_chunk_size)
+ tl.arange(0, COMM_BLOCK_S)[:, None] * stride_ps
+ peer_h_read * stride_ph
+ offs_d[None, :] * stride_pd
)
token = dl.wait(
signal_mem_ptr
+ buffer_id * num_blocks_d * H * rank_size * 8
+ num_blocks_d * H * r * 8
+ tmp * 8,
1, "gpu", "acquire",
)
local_peer_ptr = dl.consume_token(local_peer_ptr, token)
peer_data = tl.load(local_peer_ptr + peer_offs_read, mask=mask, other=0.0)
# 写入本地 C tensor(转置: r * H + h_id 变为列索引)
c_h_idx = r * H + h_id
c_offs = offs_s[:, None] * stride_cs + c_h_idx * stride_ch + offs_d[None, :] * stride_cd
tl.store(c_ptr + c_offs, peer_data, mask=mask)
# 在重置信号前确保所有读取完成
libshmem_device.fence()
for task_idx in range(logical_core_id, total_cons_tasks, n_role_cores):
tmp = task_idx
r = tmp % rank_size
tmp //= rank_size
h_id = tmp % H
block_id_d = tmp // H
dl.notify(
signal_mem_ptr
+ buffer_id * num_blocks_d * H * rank_size * 8
+ num_blocks_d * H * r * 8
+ tmp * 8,
rank, 0,
)
# 以 Vector Core 数 launch
def hccl_reverse_a2a_kernel_launcher(A, C, peer_mem, signal_mem, rank, rank_size, buffer_num, COMM_BLOCK_S, COMM_BLOCK_D):
S_total, H, D = A.shape
S = S_total // rank_size
vec_num = NPUUtils().get_aivector_core_num()
kernel_hccl_reverse_a2a_pipelined[vec_num, 1, 1](
A, C, peer_mem, signal_mem, rank, rank_size, buffer_num, S, H, D,
A.stride(0), A.stride(1), A.stride(2),
C.stride(0), C.stride(1), C.stride(2),
COMM_BLOCK_S=COMM_BLOCK_S, COMM_BLOCK_D=COMM_BLOCK_D,
)
Symmetric Memory Layout(对称内存布局)
对 Reverse All2All,对称内存布局在 block 级(而非完整 tensor 级)组织:
peer_mem layout per buffer slot: [COMM_BLOCK_S, H * rank_size, D]
- 每个 S 维 block 有 COMM_BLOCK_S 行
- H * rank_size 列: 每个 rank 的 H 个 expert 各占一个列 slot
- D 维深度
完整 peer_mem 大小: COMM_BLOCK_S * (H * world_size) * D * buffer_num
(作为该总大小的扁平 1D tensor 分配)
stride_ps = H * rank_size * D (buffer slot 内沿 S-block 行的 stride)
stride_ph = D (沿 H 列的 stride)
stride_pd = 1 (沿 D 深度的 stride)
Producer 写入远端 rank 对称内存时,H 维 offset 为 rank * H + h_id(写入为本 rank 数据预留的 slot)。Consumer 从本地对称内存读取时,offset 为 r * H + h_id(读取 rank r 写入数据的 slot)。
这种 block 级布局支持 流水线双缓冲:每个 buffer slot 一次只持 COMM_BLOCK_S 行,而非整个 S * world_size 行,显著降低对称内存占用。
性能考量
Per-task 信号 vs 全局 barrier:
dl.wait/dl.notify/fence机制比barrier_all_vec()实现更细粒度同步。每个 task 一旦自身信号满足即可推进,无需等待所有 task 完成,减少空闲时间、提升吞吐。动态 mask:kernel 用边界 mask
(offs_s < S)[:, None] & (offs_d < D)[None, :]正确处理非对齐 S 与 D 维,避免越界读/写。双缓冲:
buffer_num=2的 Producer-Consumer 模式允许 Producer 为迭代i+1写数据的同时 Consumer 读迭代i的数据,重叠两阶段。
UDMA 优化实现(Ascend 950)
Ascend 950 引入了 UDMA 链路,相比原有 MTE 链路(tl.load/tl.store + dl.symm_at 对称地址映射)能提供更高的跨卡带宽,但 UDMA 只能通过 libshmem_device.putmem/getmem 这类 RMA 原语访问。本节是 Reverse All2All(EP MoE combine 阶段)在 950 上的 UDMA 优化版本。
核心思路:把"跨卡数据交换"和"本地数据重排"拆成两个性质不同的阶段,分别用最适合的链路处理——跨卡整块传输交给 UDMA(一次性、大块、一核一 pe),本地数据搬运(包括无法用 UDMA 表达的"发给自己"这一特例)留给 MTE 的细粒度双缓冲流水线。
两阶段数据拷贝
Phase 1 — UDMA 跨卡整块交换(一次性,核↔目标 rank 一一映射):
每个 rank 把「本地 A 中属于目标 rank r 的整段 (S, H, D) 数据」一次性 putmem 给 r。
一个核只负责一个目标 rank,且一次调用搬运该 rank 的全部数据(不按 s/h/d 拆 tile)。
Phase 2 — 本地 MTE 自环搬运(细粒度、双缓冲,仅处理"发给自己"这一特例):
UDMA putmem 无法把数据"发给自己"(dst 必须是远端 shmem 地址),
所以 r == rank(自己) 这部分数据改用本地 MTE tiled 拷贝写入对称内存的本地 slot,
按 S-block 双缓冲,与后续 Consumer 读取阶段重叠。
单次 barrier_all() 同步:
等待 Phase 1 的跨卡 UDMA 写入与 Phase 2 的本地 MTE 写入全部完成。
Consumer 阶段(细粒度,MTE 读 + 输出写):
对 r != rank:直接从 Phase 1 已整块到达的对称内存读取(无需再等双缓冲,因为整块数据已一次到位)。
对 r == rank:从 Phase 2 写入的双缓冲 slot 读取(与 Phase 2 的双缓冲节奏对应)。
Phase 1 用 UDMA 换取更高的跨卡带宽;Phase 2 和 Consumer 阶段仍用 MTE,因为它们需要细粒度 tile 级并行与双缓冲流水线重叠,这正是 UDMA 的限制所不允许、 MTE 支持(卡内数据拷贝)的场景。
vector core角色划分
Kernel 沿用 Vector Core 级 sub-block 编程模型(sub_vec_id()/sub_vec_num()),并进一步把偶数子块(global_vec_id % 2 == 0)按物理核 id 分成两组角色:
role A(偶数子块,由 range(pid, rank_size, ncore) 条带化实现):UDMA 发送者
当 ncore >= rank_size(常见情况,如 950 上 ncore 远大于 rank_size)时,
该循环等价于 pid < rank_size 的核各分到唯一一个 target_rank == pid;
ncore < rank_size 时同一核可能处理多个 target_rank,但每个 target_rank
仍只被唯一一个核处理(Constraint 2 的一核一 pe 语义由步长 ncore 保证)。
每个 target_rank 一次 putmem 发完整段数据。
role B(偶数子块,pid >= rank_size 时 UDMA 循环为空,转而参与):本地 MTE 自环 Producer
仅处理 r == rank 的本地数据,按 (s_block, h, d) 细粒度 tile 双缓冲写入对称内存本地 slot。
role C(奇数子块):Consumer
按 (s_block, h, d, r) 细粒度 tile 读取对称内存(r==rank 走双缓冲 slot,r!=rank 走 Phase 1 整块到达的 slot),
写入输出 tensor。
这种角色划分让 UDMA 发送(role A)与本地 MTE 填充(role B)在同一个偶数子块分组内并行执行、互不冲突(各自处理的 (rank, s_block) 组合不重叠),随后统一用一次 barrier_all() 与 Consumer(role C)同步。
Triton-Distributed 算子实现
import triton
import triton.language as tl
import triton_dist.language as dl
from triton_dist.language.extra import libshmem_device
from triton.language.extra.cann.extension import sub_vec_id, sub_vec_num
from triton.backends.ascend.driver import NPUUtils
@triton.jit
def kernel_reverse_a2a_udma(
a_ptr, c_ptr, peer_mem_ptr,
rank, rank_size, buffer_num,
S, H, D,
stride_as, stride_ah, stride_ad,
stride_cs, stride_ch, stride_cd,
COMM_BLOCK_S: tl.constexpr,
COMM_BLOCK_D: tl.constexpr,
elem_bytes: tl.constexpr,
):
vec_num_per_aicore = 2
ncore = tl.num_programs(axis=0) * sub_vec_num() // vec_num_per_aicore
pid = tl.program_id(axis=0) * sub_vec_num() // vec_num_per_aicore
global_vec_id = pid * sub_vec_num() + sub_vec_id()
num_blocks_s = tl.cdiv(S, COMM_BLOCK_S)
num_blocks_d = tl.cdiv(D, COMM_BLOCK_D)
stride_ps = H * D # per-rank slot layout: [S, H, D]
# ---- Phase 1 (role A): one core per target rank, one-shot UDMA putmem ----
# Constraint 2: 步长必须为 ncore,确保同一 target_rank 只被唯一一个 pid 处理
# (无步长的 range(pid, rank_size) 会让多个 pid 重复处理同一 target_rank,见 udma-programming.md Constraint 2)
if global_vec_id % 2 == 0:
for target_rank in range(pid, rank_size, ncore):
if target_rank != rank:
# Constraint 3: 一次调用搬运该 target 的整段 (S, H, D),不按 tile 拆分
libshmem_device.putmem(
peer_mem_ptr + rank * S * H * D, # dst: target_rank 的对称内存, "来自 rank" 的 slot
a_ptr + target_rank * S * H * D, # src: 本地 A 中属于 target_rank 的整段数据
S * H * D * elem_bytes,
target_rank, # pe: putmem 的目标 rank
)
# ---- Phase 2 (role B): local self-copy via tiled MTE double buffer ----
for global_id_s in range(0, num_blocks_s):
buffer_id = global_id_s % buffer_num
if global_vec_id % 2 == 0 and pid >= rank_size:
total_tasks = num_blocks_d * H
logical_pid = pid - rank_size
logical_ncores = ncore - rank_size
for task_idx in range(logical_pid, total_tasks, logical_ncores):
h_id = task_idx % H
block_id_d = task_idx // H
offs_s = global_id_s * COMM_BLOCK_S + tl.arange(0, COMM_BLOCK_S)
offs_d = block_id_d * COMM_BLOCK_D + tl.arange(0, COMM_BLOCK_D)
mask = (offs_s < S)[:, None] & (offs_d < D)[None, :]
a_s = rank * S + offs_s # 本地数据中"发给自己"的那一段
a_offs = a_s[:, None] * stride_as + h_id * stride_ah + offs_d[None, :] * stride_ad
a_data = tl.load(a_ptr + a_offs, mask=mask, other=0.0)
peer_offs = (
rank * S * H * D
+ buffer_id * (COMM_BLOCK_S * stride_ps)
+ tl.arange(0, COMM_BLOCK_S)[:, None] * stride_ps
+ h_id * D
+ offs_d[None, :]
)
tl.store(peer_mem_ptr + peer_offs, a_data, mask=mask)
# ---- Sync: 等待本轮 S-block 的 UDMA 整块写入 + 本地 MTE 写入均完成 ----
libshmem_device.barrier_all()
# ---- Consumer (role C): 读取对称内存并写输出 ----
if global_vec_id % 2 == 1:
total_tasks = num_blocks_d * H * rank_size
for task_idx in range(pid, total_tasks, ncore):
tmp = task_idx
r = tmp % rank_size
tmp //= rank_size
h_id = tmp % H
block_id_d = tmp // H
offs_s = global_id_s * COMM_BLOCK_S + tl.arange(0, COMM_BLOCK_S)
offs_d = block_id_d * COMM_BLOCK_D + tl.arange(0, COMM_BLOCK_D)
mask = (offs_s < S)[:, None] & (offs_d < D)[None, :]
if r == rank:
# 本地自环数据:读 Phase 2 写入的双缓冲 slot
peer_offs = (
rank * S * H * D
+ buffer_id * (COMM_BLOCK_S * stride_ps)
+ tl.arange(0, COMM_BLOCK_S)[:, None] * stride_ps
+ h_id * D
+ offs_d[None, :]
)
else:
# 远端数据:Phase 1 已一次性整块到达,直接按 s_block 索引读取
peer_offs = (
r * S * H * D
+ global_id_s * COMM_BLOCK_S * stride_ps
+ tl.arange(0, COMM_BLOCK_S)[:, None] * stride_ps
+ h_id * D
+ offs_d[None, :]
)
peer_data = tl.load(peer_mem_ptr + peer_offs, mask=mask, other=0.0)
c_h_idx = r * H + h_id
c_offs = offs_s[:, None] * stride_cs + c_h_idx * stride_ch + offs_d[None, :] * stride_cd
tl.store(c_ptr + c_offs, peer_data, mask=mask)
def launch_reverse_a2a_udma(A, C, peer_mem, rank, rank_size, buffer_num, COMM_BLOCK_S, COMM_BLOCK_D):
S_total, H, D = A.shape
S = S_total // rank_size
aicore_num = NPUUtils().get_aicore_num()
kernel_reverse_a2a_udma[aicore_num, 1, 1](
A, C, peer_mem, rank, rank_size, buffer_num, S, H, D,
A.stride(0), A.stride(1), A.stride(2),
C.stride(0), C.stride(1), C.stride(2),
COMM_BLOCK_S=COMM_BLOCK_S, COMM_BLOCK_D=COMM_BLOCK_D,
elem_bytes=A.element_size(),
)
Symmetric Memory Layout(对称内存布局)
与纯 MTE 版本(对称内存按 [buffer_num, COMM_BLOCK_S, H*rank_size, D] 组织,逐 tile 双缓冲)不同,UDMA 版本的对称内存按 rank slot 组织,只在"本地自环"这一个 slot 内部做双缓冲:
peer_mem 总大小: rank_size * S * H * D
per-rank slot: [S, H, D],slot 起始偏移 = r * S * H * D
- r != rank 的 slot: 由 Phase 1 UDMA putmem 一次性整块写入,Consumer 直接按 s_block 索引读取,无需缓冲
- r == rank 的 slot: 由 Phase 2 本地 MTE 写入,仅在 slot 内偏移
[0, buffer_num * COMM_BLOCK_S * H * D) 的前缀区间做双缓冲循环复用
这种布局把"整块一次到达、无需缓冲"(远端 UDMA 数据)与"逐 block 流水、需要双缓冲"(本地自环数据)区分开,避免对不需要缓冲的远端数据也分配 buffer_num 倍空间,比纯 MTE 版本更省对称内存。
性能考量
UDMA 带宽优势集中在 Phase 1:跨卡数据量为
(rank_size - 1) × S × H × D × elem_bytes,用rank_size - 1次大块putmem完成,充分利用 UDMA 相对 MTE 的带宽优势;tile 越大、调用次数越少,越接近 UDMA 峰值带宽。本地自环仍需 MTE 细粒度双缓冲:
r == rank这部分数据量为S × H × D,占总数据量的1/rank_size,用 MTE tiled 双缓冲处理,与 Consumer 读取重叠,其相对开销随rank_size增大而减小。核数分配的权衡:
pid < rank_size的核专职 UDMA 发送,pid >= rank_size的核专职本地 MTE 填充——当rank_size较大时,参与本地 MTE 填充的核数变少(ncore - rank_size),可能造成本地自环阶段成为瓶颈。单次 barrier_all() 覆盖两条链路:Phase 1(UDMA)与 Phase 2(MTE)写入共用同一个
barrier_all(),意味着 Consumer 必须等两条链路中较慢的一条完成。若 UDMA 与 MTE 延迟差异明显,为两条链路设置独立信号,让 Consumer 按数据来源(r == rankvsr != rank)分别等待对应链路,而不必被最慢的一条拖累。内存对齐:与纯 MTE 版本相同,
(H, D)布局需满足 32B 对齐;UDMA 的整块传输对总字节数(S*H*D*elem_bytes)没有额外对齐要求,但 Phase 2/Consumer 中的 tile 级tl.load/tl.store仍受 MTE 对齐约束。
常见问题
1. 把 MTE 版本的 tile 循环直接套用到 putmem 上
如果照搬纯 MTE 版本按 (s_block, h, d) 拆 tile 后逐 tile 调用 putmem,会导致精度问题,由于目前UDMA仅支持单QP通信,多核并发读写同一rank时,单QP的信号量槽位不足以保证多核通信全部完成。对策:Phase 1 必须重新设计为"一核一 pe、一次调用整段数据",不能是 MTE tile 循环的逐行替换。
2. 试图用 putmem 处理"发给自己"的数据
putmem 的 dst 必须是远端 shmem 地址,无法表达"发给自己"。对策:这部分数据永远走本地 MTE(tl.load/tl.store),如 Phase 2 所示。
3. 对已整块到达的远端数据仍套用双缓冲索引
远端数据经 UDMA 一次性到达后已是完整的 [S, H, D],Consumer 读取时应直接用 global_id_s * COMM_BLOCK_S 索引到该整块内部,而不是像本地自环数据那样用 buffer_id 做双缓冲偏移——两者的对称内存寻址方式不同,混用会读到错误偏移。对策:Consumer 按 r == rank 与 r != rank 分别选择偏移公式(如代码示例所示)。
Wait/Notify 优化
上述 UDMA + MTE 双链路实现使用 单次 barrier_all() 同步 Phase 1(UDMA 跨卡写入)、Phase 2(本地 MTE 自环写入)与 Consumer 读取,这意味着 Consumer 必须等待所有 Producer task(包括所有 rank 的 UDMA 发送 + 本地 MTE 自环)全部完成才能开始读取,存在全局同步瓶颈。
优化目标
将 barrier_all() 替换为 细粒度 wait/notify 信号量同步,让每个 Consumer task 只需等待自己依赖的 Producer 信号即可开始读取,而不会被最慢的 Producer 阻塞。
核心改动
1. 双信号内存区域设计
signal_mem: flat int64 array, 分为两个独立区域
Region A — 跨卡 UDMA 信号 [num_blocks_s, rank_size]:
槽语义:(global_id_s, sender_rank) 表示"sender_rank 已完成 S-block global_id_s 的 UDMA 发送"
槽总数:num_blocks_s * rank_size
槽地址:signal_mem_ptr + (global_id_s * rank_size + sender_rank) * SIG_STRIDE
Notifier:UDMA Sender(偶数核,logical_core_id < rank_size)
Waiter:Consumer(奇数核),等待 r != rank 的远端数据
Buffer 复用:无(每个 (block, rank) 对使用唯一地址,不需要复位)
Credit 回传:不需要
Region B — 本地 MTE 自环双缓冲信号 [buffer_num, tasks_per_block]:
槽语义:(buffer_id, task_idx) 表示"本地 MTE 已完成 buffer_id 的 task_idx 写入"
槽总数:buffer_num * num_blocks_d * H
槽地址:signal_mem_ptr + region_b_base + (buffer_id * tasks_per_block + task_idx) * SIG_STRIDE
Notifier:MTE Producer(偶数核,logical_core_id >= rank_size)
Waiter:Consumer(奇数核),等待 r == rank 的本地自环数据
Buffer 复用:是(num_blocks_s / buffer_num 轮)
Credit 回传:必须(Consumer 读完后 notify(..., 0) 重置信号)
其中 SIG_STRIDE = 8(一个 cache line),避免 false sharing。
region_b_base = num_blocks_s * rank_size * SIG_STRIDE。
关键设计原则:
两条链路独立信号:UDMA 与 MTE 各有独立的信号区域,互不干扰,Consumer 按数据来源(
r == rankvsr != rank)分别等待对应信号UDMA 信号无需复位:每个
(block, sender)对使用唯一地址,写一次用一次,waiter 数量无限制MTE 信号需双向握手:Producer 等待信号清零(
wait(waitValue=0))→ 写入 → notify(1);Consumer 等待信号置位(wait())→ 读取 → notify(0) 回传 credit
2. UDMA Sender Stage(无需等待信号,发送后立即 notify)
# Phase 1 — UDMA 跨卡整块发送(偶数核,logical_core_id < rank_size)
if global_vec_id % 2 == 0 and logical_core_id < rank_size:
for udma_target_rank in range(logical_core_id, rank_size, n_role_cores):
if udma_target_rank != rank:
# Step 1: UDMA putmem(同步语义,返回即数据已达远端)
libshmem_device.putmem(
peer_mem_ptr + rank * S * H * D + global_id_s * COMM_BLOCK_S * stride_ps,
a_ptr + udma_target_rank * S * H * D + global_id_s * COMM_BLOCK_S * stride_as,
cur_block_s * H * D * elem_bytes,
udma_target_rank,
)
# Step 2: 向每个目标 rank 发送完成信号(UDMA 同步,无需 fence)
for udma_target_rank in range(logical_core_id, rank_size, n_role_cores):
if udma_target_rank != rank:
dl.notify(
signal_mem_ptr + (global_id_s * rank_size + rank) * SIG_STRIDE,
udma_target_rank.to(tl.int32),
)
关键点:
UDMA
putmem是同步的(调用返回即数据已达远端),notify 前不需要fence()每个 sender_rank 向目标 rank 的信号槽
(global_id_s, rank)发送信号,表示"我这个 rank 的这个 S-block 已发完"
3. MTE Producer Stage(双向握手:等待清零 → 写入 → 发送信号)
# Phase 2 — 本地 MTE 自环写入(偶数核,logical_core_id >= rank_size)
if global_vec_id % 2 == 0 and logical_core_id >= rank_size:
logical_mte_pid = logical_core_id - rank_size
logical_mte_ncores = n_role_cores - rank_size
for task_idx in range(logical_mte_pid, tasks_per_block, logical_mte_ncores):
h_id = task_idx % H
block_id_d = task_idx // H
# ... 计算 offsets 和 mask ...
# Step 1: 等待 Consumer 已读完上一轮(信号值为 0)
self_sig_ptr = (
signal_mem_ptr + region_b_base
+ (buffer_id * tasks_per_block + task_idx) * SIG_STRIDE
)
token = dl.wait(self_sig_ptr, 1, "gpu", "acquire", waitValue=0)
store_ptr = dl.consume_token(peer_mem_ptr, token)
# Step 2: 加载本地数据并写入对称内存
a_data = tl.load(a_ptr + a_offs, mask=mask, other=0.0)
peer_offs_write = ( ... )
tl.store(store_ptr + peer_offs_write, a_data, mask=mask)
# Step 3: fence 确保所有 MTE 写入可见
libshmem_device.fence()
# Step 4: 向每个 task 发送完成信号
for task_idx in range(logical_mte_pid, tasks_per_block, logical_mte_ncores):
self_sig_ptr = (
signal_mem_ptr + region_b_base
+ (buffer_id * tasks_per_block + task_idx) * SIG_STRIDE
)
dl.notify(self_sig_ptr, rank)
关键点:
MTE 需要
fence()确保写入可见,且fence()必须在notify()之前双向握手:Producer 先
wait(waitValue=0)确保 Consumer 已读完上一轮,再写入,再notify(1)
4. Consumer Stage(按数据来源分别等待信号 + Credit 回传)
# Consumer(奇数核)
if global_vec_id % 2 == 1:
total_cons_tasks = tasks_per_block * rank_size
for task_idx in range(logical_core_id, total_cons_tasks, n_role_cores):
tmp = task_idx
r = tmp % rank_size
tmp //= rank_size
h_id = tmp % H
block_id_d = tmp // H
# ... 计算 offsets 和 mask ...
if r == rank:
# 本地自环数据:等待 MTE Producer 信号
self_sig_ptr = (
signal_mem_ptr + region_b_base
+ (buffer_id * tasks_per_block + tmp) * SIG_STRIDE
)
token = dl.wait(self_sig_ptr, 1, "gpu", "acquire")
load_ptr = dl.consume_token(peer_mem_ptr, token)
peer_offs_read = ( ... ) # 双缓冲地址
peer_data = tl.load(load_ptr + peer_offs_read, mask=mask, other=0.0)
else:
# 远端数据:等待 UDMA Sender 信号
remote_sig_ptr = (
signal_mem_ptr + (global_id_s * rank_size + r) * SIG_STRIDE
)
token = dl.wait(remote_sig_ptr, 1, "gpu", "acquire")
load_ptr = dl.consume_token(peer_mem_ptr, token)
peer_offs_read = ( ... ) # 整块地址
peer_data = tl.load(load_ptr + peer_offs_read, mask=mask, other=0.0)
# 写入输出 tensor
c_offs = ( ... )
tl.store(c_ptr + c_offs, peer_data, mask=mask)
# Step: fence 确保所有读取完成
libshmem_device.fence()
# Step: Credit 回传(仅对本地自环数据)
for task_idx in range(logical_core_id, total_cons_tasks, n_role_cores):
tmp = task_idx
r = tmp % rank_size
tmp //= rank_size
if r == rank:
self_sig_ptr = (
signal_mem_ptr + region_b_base
+ (buffer_id * tasks_per_block + tmp) * SIG_STRIDE
)
dl.notify(self_sig_ptr, rank, 0) # 重置为 0,释放下一轮 Producer
性能收益
同步模式 |
Consumer 等待时间 |
全局同步点数 |
|---|---|---|
|
|
每个 S-block 一次 |
|
|
无全局同步 |
Host侧修改
1. Signal Memory Allocation(信号量内存分配)
num_blocks_s = -(-S // COMM_BLOCK_S)
num_blocks_d = -(-D // COMM_BLOCK_D)
signal_mem_size = (
num_blocks_s * rank_size + buffer_num * num_blocks_d * H
) * 8 # int64, 8 bytes per slot
signal_mem = ash.aclshmem_create_tensor(
[signal_mem_size], dtype=torch.int64, device_id=rank
)
2. kernel launch前重置信号量
for _ in range(iters):
signal_mem.fill_(0) # 重置所有信号为 0
dist.barrier() # 确保所有 rank 都已清零
launcher(...)
torch.npu.synchronize()
dist.barrier() # 确保所有 rank 都已完成
关键点:
每次 launch 前必须
fill_(0),否则残留信号会让下一次wait直接通过(race condition)清零前后各一个
dist.barrier():前者防止本 rank 清零擦掉对端信号,后者防止下一轮清零擦掉本轮尾部信号