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)中:

  1. Forward All2All (Dispatch):token 被发送到持有其指定 expert 的 rank。每 rank 的输入 (S, H, D) 变为 (S/world_size, H * world_size, D)

  2. Expert GEMM:每个 rank 用 GroupedGEMM 或 Fused MoE 为本地 expert 计算 GEMM。

  3. 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_idVector 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_id

  • n_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 端:

  1. dl.wait(remote_signal_ptr, 1, "gpu", "acquire", waitValue=0) —— 等待目标 rank 的信号槽被清零(值 == 0),表示目标 rank 的 consumer 已读完上一个 buffer

  2. dl.consume_token(remote_ptr, token) —— 强制依赖生效,使远程写入仅在 wait 完成后发生

  3. tl.store(remote_ptr + offset, data, mask=mask) —— 将数据写入远端对称内存

  4. libshmem_device.fence() —— 确保所有远程写入在发送通知前可见

  5. dl.notify(signal_mem_ptr, target_rank) —— 通知目标 rank 数据可用

Consumer 端:

  1. dl.wait(signal_mem_ptr, 1, "gpu", "acquire") —— 等待来自源 rank 的信号(数据可用)

  2. dl.consume_token(local_peer_ptr, token) —— 强制依赖生效,使本地读取仅在 wait 完成后发生

  3. tl.load(local_peer_ptr + offset, mask=mask) —— 从本地对称内存读取数据

  4. tl.store(output_ptr + offset, data, mask=mask) —— 写入输出 tensor

  5. libshmem_device.fence() —— 确保所有读取在重置信号前完成

  6. 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 全局 barrierdl.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 == rank vs r != 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 == rankr != 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 == rank vs r != 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 等待时间

全局同步点数

barrier_all() (baseline)

max(所有 UDMA 发送, 所有 MTE 自环)

每个 S-block 一次

wait/notify (optimized)

t_UDMA[r](远端数据)或 t_MTE[task](本地数据)

无全局同步

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 清零擦掉对端信号,后者防止下一轮清零擦掉本轮尾部信号