# 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_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_id - **n_role_cores** = `ncore // 2` —— 分配给每个角色的 core 数 - 偶数 pid 负责所有跨卡远程 store 操作(Producer) - 奇数 pid 负责所有本地对称内存读取与输出写入操作(Consumer) kernel 以 **Vector Core 总数** 而非 AICore 数 launch: ```python 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 算子 ```python 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 算子实现 ```python 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 == 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 == 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) ```python # 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(双向握手:等待清零 → 写入 → 发送信号) ```python # 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 回传) ```python # 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(信号量内存分配) ```python 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前重置信号量 ```python 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 清零擦掉对端信号,后者防止下一轮清零擦掉本轮尾部信号