7.分布式并行通信模拟
用 Python 模拟实现 Ring AllReduce 的完整过程:给定 N 个节点,每个节点持有一个等长向量,要求通过环形传递实现求和并广播到所有节点,写出每一步的数据传输与累加逻辑。¶
1.1 Ring AllReduce 原理¶
Ring AllReduce 是一种高效的分布式规约算法,它将所有节点排列成一个逻辑环。整个操作分为两个阶段:
-
Scatter-Reduce(分散-规约):将每个节点的完整向量切分成 N 个分块(N 为节点数)。执行 N−1 步,每步每个节点向它的下游邻居发送一个分块,并从上游邻居接收一个分块,然后将接收到的分块累加到自己的对应分块上。经过 N−1 步后,每个节点上都有一个分块包含了所有节点的贡献,即该分块的全局和。但不同节点持有不同的分块。
-
AllGather(全收集):再经过 N−1 步,每个节点将已经规约好的分块向下游发送,同时从上游接收新的已规约分块并替换自己的对应分块(不再累加)。最终所有节点都获得所有分块的全局和,即每个节点都拥有完整的规约结果。
1.2 模拟实现¶
假设每个节点是一个 Python 进程(这里用对象模拟),节点之间通过点对点 send/recv 通信。为了简化,我们用消息队列或直接函数调用模拟。
import numpy as np
class RingAllReduceSimulator:
def __init__(self, num_nodes, vector_length):
self.num_nodes = num_nodes
self.vec_len = vector_length
# 每个节点持有自己的本地向量
self.vectors = [np.random.randn(vector_length) for _ in range(num_nodes)]
# 分块大小(假设均匀分割)
assert vector_length % num_nodes == 0
self.chunk_size = vector_length // num_nodes
def run(self):
# 初始状态:每个节点将向量切成 num_nodes 块
# 使用缓冲区存储当前各节点的数据块
local_chunks = []
for r in range(self.num_nodes):
chunks = []
for i in range(self.num_nodes):
start = i * self.chunk_size
end = start + self.chunk_size
chunks.append(self.vectors[r][start:end].copy())
local_chunks.append(chunks) # local_chunks[rank][i] 是节点 rank 的第 i 块
# ========= Scatter-Reduce 阶段 =========
for step in range(self.num_nodes - 1):
# 每个节点计算本轮要发送给下游的分块索引
send_chunk_idx = (rank - step) % self.num_nodes # 常见公式
recv_chunk_idx = (rank - step - 1) % self.num_nodes
# 模拟通信:每个节点将 send_chunk_idx 块发送给下游邻居
# 下游邻居收到后累加到自己的 recv_chunk_idx 块上
# 注意:所有节点同时进行,这里用临时存储来避免先写后读问题
new_chunks = [chunks.copy() for chunks in local_chunks]
for rank in range(self.num_nodes):
downstream = (rank + 1) % self.num_nodes
upstream = (rank - 1) % self.num_nodes
# 发送本节点的 send_chunk_idx 块给下游
data_to_send = local_chunks[rank][send_chunk_idx].copy()
# 下游接收并累加到自己的 recv_chunk_idx
new_chunks[downstream][recv_chunk_idx] += data_to_send
local_chunks = new_chunks
# ========= AllGather 阶段 =========
for step in range(self.num_nodes - 1):
send_chunk_idx = (rank - step) % self.num_nodes
recv_chunk_idx = (rank - step - 1) % self.num_nodes
new_chunks = [chunks.copy() for chunks in local_chunks]
for rank in range(self.num_nodes):
downstream = (rank + 1) % self.num_nodes
# 发送本节点已规约的 send_chunk_idx 块给下游
data_to_send = local_chunks[rank][send_chunk_idx].copy()
# 下游接收到后直接替换对应的块(不累加)
new_chunks[downstream][recv_chunk_idx] = data_to_send
local_chunks = new_chunks
# 最终每个节点将分块拼接成完整向量
final_vectors = []
for rank in range(self.num_nodes):
final_vec = np.concatenate(local_chunks[rank])
final_vectors.append(final_vec)
return final_vectors
1.3 关键细节¶
-
发送/接收索引:在 Scatter-Reduce 阶段,第
step步时,节点rank发送的块索引为(rank - step) % N,接收的块索引为(rank - step - 1) % N。这样确保了环上数据流动的正确性,每个节点逐渐将不同块的贡献分发出去。 -
累加与替换:Scatter-Reduce 阶段是累加,AllGather 阶段是直接替换。
-
通信同步:真实环境中所有节点同时通信,模拟时通过循环和临时拷贝保证逻辑正确。
实现 AllGather 的朴素模拟:每个节点有一个本地张量,要求通过点对点通信(发送/接收)将所有节点的数据收集到每个节点上,不允许使用集体通信原语,给出伪代码或 Python 代码。¶
2.1 原理¶
AllGather 的目标是让每个节点都获得所有节点的数据。最简单的点对点实现方式是:所有节点先将数据发送给一个“根”节点,根节点拼接后再广播给所有节点。但题目要求不能使用集体通信,所以我们采用环形传递:经过 N−1 步,每个节点将接收到的数据向邻居转发,最终每个节点都能收集到所有数据。
2.2 模拟实现(环形 AllGather)¶
def ring_all_gather(rank, num_nodes, local_data):
"""
模拟环形 AllGather。每个节点持有 local_data (numpy array),
返回收集到的所有节点数据的列表。
"""
import numpy as np
# 将本地数据复制到缓冲区,初始缓冲区只包含本地数据
gathered = [None] * num_nodes
gathered[rank] = local_data.copy()
# 向邻居逐步转发
current_data = local_data.copy()
for step in range(1, num_nodes):
downstream = (rank + 1) % num_nodes
upstream = (rank - 1) % num_nodes
# 发送当前持有的数据给下游,同时从上游接收数据
# 这里用 send/recv 模拟:我们通过预先定义的全局数组来模拟
# 简化实现:在一轮内所有节点同时交换
# 由于是模拟,我们采用同步步骤:每个节点先发送,再接收
send_to = downstream
recv_from = upstream
# 用临时变量保存发送数据,接收覆盖 current_data
# 实际伪代码:send(current_data, to=downstream); recv(current_data, from=upstream)
# 这里无法真正通信,使用上一轮的全局状态模拟:
# 在真实环境中,这一步后 current_data 变成上游节点上一步持有的数据。
# 我们简化:手工计算当前节点在这一步结束时应该拥有的数据。
# 实际上环状传递中,经过 step 步后,节点持有的是最初 rank - step 节点的数据。
source_rank = (rank - step) % num_nodes
current_data = local_data_of[source_rank] # 模拟接收
gathered[source_rank] = current_data.copy()
return gathered
更完整的模拟可以通过维护全局列表并模拟每一轮的发送与接收来实现:
def simulate_ring_all_gather(data_list):
"""data_list: 各节点的 numpy 数组列表"""
n = len(data_list)
# 每个节点当前持有的数据(初始为自己)
hold = [d.copy() for d in data_list]
gathered = [ [None]*n for _ in range(n) ]
for r in range(n):
gathered[r][r] = data_list[r].copy()
# 环形传递 n-1 步
for step in range(1, n):
# 每一步,节点将自己的 hold 发送给下游,同时从上游接收
new_hold = [None]*n
for r in range(n):
downstream = (r + 1) % n
new_hold[downstream] = hold[r] # 发送
hold = new_hold
# 接收后记录
for r in range(n):
source_rank = (r - step) % n
gathered[r][source_rank] = hold[r].copy()
return gathered
2.3 通信复杂度¶
朴素环形 AllGather 需要 N−1 步,总通信量与数据量成正比,适用于小规模节点。
模拟 ReduceScatter 操作:N 个节点,每个节点有一个等长向量,要求最终每个节点得到规约后不同分片的结果,基于 Send/Recv 实现。¶
3.1 原理¶
ReduceScatter 的结果是:将长度为 L 的向量分成 N 块,每个节点最终得到第 ii 块的规约和(所有节点对应块之和)。可通过环形规约类似 Ring AllReduce 的 Scatter-Reduce 阶段实现,但是不再进行 AllGather。经过 N−1 步后,每个节点上已累加了所有节点的一个分块,该分块就是全局和。
3.2 实现(环形 ReduceScatter)¶
def ring_reduce_scatter(vectors, num_nodes):
"""
vectors: list of numpy arrays, 每个节点一个
返回: list of numpy arrays, 每个节点得到对应分片的规约结果
"""
chunk_size = len(vectors[0]) // num_nodes
# 每个节点将自己的向量切分成 num_nodes 块
chunks = []
for v in vectors:
parts = [v[i*chunk_size:(i+1)*chunk_size] for i in range(num_nodes)]
chunks.append(parts)
# Scatter-Reduce 阶段(只累加,不 AllGather)
for step in range(num_nodes - 1):
new_chunks = [ [c.copy() for c in parts] for parts in chunks ] # 深拷贝
for rank in range(num_nodes):
send_idx = (rank - step) % num_nodes
recv_idx = (rank - step - 1) % num_nodes
downstream = (rank + 1) % num_nodes
# 发送 send_idx 块给下游,下游接收后累加到自己的 recv_idx 块
data_to_send = chunks[rank][send_idx].copy()
new_chunks[downstream][recv_idx] += data_to_send
chunks = new_chunks
# 每个节点最终保留自己 rank 对应的分片(根据约定,通常是节点 rank 得到第 rank 块)
results = [ chunks[rank][rank].copy() for rank in range(num_nodes) ]
return results
经过 N−1 步后,节点 rank 的第 rank 块已经累加了所有节点的对应块,因此 results[rank] 即为全局规约的第 rank 分片。
编写一个 Broadcast 的树形实现:定义一棵二叉树覆盖所有节点,根节点将数据沿树层层下发,写出递归或迭代的通信逻辑。¶
4.1 二叉树广播原理¶
将节点组织成一棵完全二叉树,根节点持有要广播的数据。根将数据发送给它的两个子节点;每个子节点收到后再发送给自己的子节点,如此递归。这样广播可以在 O(logN) 步内完成,比线性广播高效。
4.2 迭代实现(基于父子关系)¶
假设节点编号为 0 到 N-1,在二叉树中,节点 i 的左子节点为 2i+1,右子节点为 2i+2,父节点为 (i-1)//2。广播由根节点0发起。
def tree_broadcast(data, num_nodes):
"""
data: 要广播的数据(根节点持有)
num_nodes: 节点总数
模拟返回所有节点最终接收到的数据列表
"""
received = [None] * num_nodes
received[0] = data # 根已有数据
# 按照层级顺序,每个节点向子节点发送
for i in range(num_nodes):
left = 2 * i + 1
right = 2 * i + 2
if left < num_nodes:
received[left] = received[i] # 模拟发送
if right < num_nodes:
received[right] = received[i]
return received
若需要更真实的通信模拟,可以使用递归函数,每个节点在收到数据后调用 send 函数:
def broadcast(node, data, children):
node.receive(data)
for child in children:
send(data, child)
broadcast(child, data, children_of_child)
4.3 通信复杂度¶
树形广播的通信步数为 ⌈log2N⌉,数据传输总量为 O(N)。优势在于延迟低,适合大规模集群。
用 Python 多进程或伪代码模拟 AlltoAll 通信:每个节点有 N 个分块,第 i 块发往节点 i,同时接收来自其他节点的对应块。¶
5.1 原理¶
AlltoAll 是一种全交换通信:每个节点向每个其他节点(包括自身)发送一个独立的数据块,并接收来自每个其他节点的数据块。典型实现是:节点 i 向节点 j 发送块 j,从节点 j 接收块 i。
5.2 模拟实现¶
假设每个节点拥有一个长度为 N 的列表,blocks[r][c] 表示节点 r 要发送给节点 c 的数据块。最终每个节点 c 收集到所有节点发给它的块。
def simulate_alltoall(all_blocks):
"""
all_blocks: 二维列表,all_blocks[r][c] 是节点 r 打算发给节点 c 的数据块。
返回: 每个节点收到的块列表 (received[r][c] 来自节点 c 的块)。
"""
N = len(all_blocks)
received = [ [None]*N for _ in range(N) ]
# 所有节点同时进行:对于每个发送者 r 和接收者 c,将 all_blocks[r][c] 发送到 received[c][r]
for r in range(N):
for c in range(N):
# r 发送给 c
received[c][r] = all_blocks[r][c] # 模拟即时完成
return received
更真实的通信模拟可以用步数表示:如果网络不支持全交换,可通过多轮点对点实现,例如使用环形 AlltoAll 算法,此处从简。
5.3 环形 AlltoAll 的简单模拟¶
在环形拓扑中,通过 N-1 步,每个节点每步向邻居发送一个数据块并转发。
def ring_alltoall(blocks):
N = len(blocks)
# 每个节点当前持有的待发送列表,初始就是要发给各个节点的块
sendbuf = [ list(blocks[r]) for r in range(N) ] # sendbuf[r][c]
recvbuf = [ [None]*N for _ in range(N) ]
for step in range(N-1):
for r in range(N):
downstream = (r + 1) % N
# 决定本轮发送哪一块:通常是 (r - step) 索引
send_idx = (r - step) % N
# 发送 sendbuf[r][send_idx] 给下游,下游将其存入 recvbuf[downstream][send_idx]
recvbuf[downstream][send_idx] = sendbuf[r][send_idx]
# 更新 sendbuf 为刚接收到的数据,继续转发
# 细节略
设计一个基于 AllReduce 的同步 SGD 模拟器:每个节点持有模型梯度,调用 AllReduce 求平均梯度并更新参数,用函数封装通信步骤。¶
6.1 同步 SGD 流程¶
在数据并行训练中,每个 worker 在本地计算出一个 mini-batch 的梯度,然后所有 worker 通过 AllReduce 对这些梯度求平均(或求和后除以 world_size),最后每个 worker 用平均梯度更新模型参数。此过程保证所有 worker 上的模型参数始终保持一致。
6.2 模拟器实现¶
import numpy as np
class SyncSGDSimulator:
def __init__(self, num_workers, param_shape):
self.num_workers = num_workers
# 所有worker共享相同的初始参数(这里简化:每个worker持有一份副本)
self.params = [np.random.randn(*param_shape) for _ in range(num_workers)]
# 模拟AllReduce函数
self.allreduce = RingAllReduceSimulator(num_workers, np.prod(param_shape)).run # 略
def compute_gradients(self, worker_id):
# 模拟计算随机梯度
return np.random.randn(*self.params[0].shape)
def step(self):
# 每个worker计算本地梯度
grads = [self.compute_gradients(i) for i in range(self.num_workers)]
# 对梯度做 AllReduce 求平均
# 简单实现:直接求和再除以 N(假设有高效的AllReduce实现)
avg_grad = sum(grads) / self.num_workers
# 每个worker更新参数(这里统一更新,实际每worker各自更新自己的副本)
for i in range(self.num_workers):
self.params[i] -= 0.01 * avg_grad # 学习率0.01
我们也可以将上一题的 RingAllReduce 用于梯度平均:
def allreduce_gradients(grads):
# 假设 grads 是长度相同的向量列表
summed = ring_allreduce(grads) # 返回每个节点得到的总和(完整向量)
avg = [g / len(grads) for g in summed]
return avg
实现数据并行中的梯度同步:假设有 P 个 worker,各自完成一次前向+反向得到梯度,写出如何通过 AllReduce 汇总梯度并更新模型参数的过程。¶
7.1 数据并行梯度同步的详细步骤¶
-
每个 worker 从参数服务器(或本地模型副本)获取最新模型参数。
-
每个 worker 读取不同的 mini-batch 数据,独立执行前向传播和反向传播,计算出本地梯度 gigi。
-
所有 worker 参与 AllReduce 操作,对梯度求和(或平均)。通常采用平均,即 G=1P∑giG=P1∑gi。
-
每个 worker 使用平均梯度更新本地模型参数(如 SGD: w←w−ηGw←w−ηG)。
-
进入下一轮迭代。
伪代码:
for step in range(num_steps):
data = next_batch(worker_id)
grad = compute_gradient(model, data)
avg_grad = allreduce(grad, op=SUM) / world_size
update_model(model, avg_grad)
在同步 SGD 中,AllReduce 是核心通信操作,通常使用 Ring AllReduce 或 Tree AllReduce 来高效实现。
7.2 模拟代码(结合 AllReduce)¶
def data_parallel_step(workers_params, compute_grad_fn, allreduce_fn, lr=0.01):
# 每个worker计算梯度
grads = [compute_grad_fn(w) for w in workers_params]
# AllReduce 求平均
avg_grad = allreduce_fn(grads) # 返回每个节点得到的全局平均梯度列表
# 更新
for i, param in enumerate(workers_params):
param -= lr * avg_grad[i]
实现模型并行中列切分(Column-wise)线性层的前向和反向通信模拟:输入 X 被广播,权重按列切分在各节点,前向计算后输出需 AllGather 拼接,反向计算后梯度需 ReduceScatter。¶
8.1 列切分线性层原理¶

8.2 前向模拟¶
def column_wise_forward(X, W_parts):
"""
X: 输入矩阵 [batch, d_in]
W_parts: 列表,每个元素是某设备上的权重切片 [d_in, d_out_part]
返回: Y 完整输出
"""
N = len(W_parts)
# 1. 广播 X 到所有设备(这里假设所有设备已拥有 X)
Y_parts = []
for i in range(N):
Yi = X @ W_parts[i] # [batch, d_out_part]
Y_parts.append(Yi)
# 2. AllGather 拼接
Y = np.concatenate(Y_parts, axis=1) # 模拟 AllGather 后的结果
return Y
8.3 反向模拟¶
def column_wise_backward(X, W_parts, dY):
"""
dY: 完整输出梯度 [batch, d_out]
返回: dX [batch, d_in], dW_parts 列表
"""
N = len(W_parts)
part_dim = dY.shape[1] // N
dY_parts = [dY[:, i*part_dim:(i+1)*part_dim] for i in range(N)] # 切分
dX_local = []
dW_parts = []
for i in range(N):
dW_i = X.T @ dY_parts[i] # [d_in, d_out_part]
dX_i = dY_parts[i] @ W_parts[i].T # [batch, d_in]
dW_parts.append(dW_i)
dX_local.append(dX_i)
# 对 dX 进行 ReduceScatter(求和):因为每个设备算出的 dX_i 是部分贡献,需要全局求和得到完整 dX
dX = sum(dX_local) # 模拟 AllReduce SUM
# 在真实分布式环境中,dX 的求和可通过 ReduceScatter 实现:每个设备贡献一部分,然后各设备获得不同分片?列切分下 dX 需要完整,所以用 AllReduce。
return dX, dW_parts
注意:标准列切分的反向对 dX 实际上需要 AllReduce(各设备 dX_i 相加),但有些实现可以通过技巧避免。更准确地说,Megatron-LM 中列切分的前向输出是 AllGather,反向输入梯度是 Split(不通信),而对 X 的梯度需要 AllReduce。此处我们演示了通信原语。
实现模型并行中行切分(Row-wise)线性层的通信模拟:输入 X 按列切分分布在节点,权重按行切分,前向输出需 AllReduce,反向时梯度广播或相应通信。¶
9.1 行切分原理¶

9.2 前向模拟¶
def row_wise_forward(X_parts, W_parts):
"""
X_parts: 各设备上的输入切片列表,每个 [batch, d_in_part]
W_parts: 各设备上的权重切片列表,每个 [d_in_part, d_out]
返回: Y 完整输出 [batch, d_out]
"""
N = len(W_parts)
Y_parts = []
for i in range(N):
Yi = X_parts[i] @ W_parts[i] # [batch, d_out]
Y_parts.append(Yi)
# AllReduce SUM
Y = sum(Y_parts) # 模拟 AllReduce
return Y
9.3 反向模拟¶
def row_wise_backward(X_parts, W_parts, dY):
"""
dY: 完整输出梯度 [batch, d_out]
返回: dX_parts (各设备输入梯度), dW_parts
"""
N = len(W_parts)
dX_parts = []
dW_parts = []
for i in range(N):
dW_i = X_parts[i].T @ dY # [d_in_part, d_out]
dX_i = dY @ W_parts[i].T # [batch, d_in_part]
dW_parts.append(dW_i)
dX_parts.append(dX_i)
# 注意:dY 本身就是完整的,每个设备拿到相同的 dY,所以不需要额外通信。
return dX_parts, dW_parts
9.4 通信总结¶
-
前向:在各设备上计算局部输出,然后通过 AllReduce(求和)得到完整结果。
-
反向:上游梯度
dY已被广播到所有设备(因为前向输出是 AllReduce 得到的,反向时dY是完整的,各设备都持有完整dY),因此各设备独立计算dX_i和dW_i,无需通信。但这要求前向时 AllReduce 后的输出在各个设备上都可用(即 AllReduce 的结果在每个设备上相同),这自然满足。
行切分的前向 AllReduce 是必需的,反向则没有通信,这是它与列切分的主要区别。
模拟 Megatron 张量并行中 MLP 层的通信:第一个线性层列切(f),第二个线性层行切(g),写出前向和反向中 AllReduce 和 ReduceScatter 的调用顺序和内容。¶
10.1 Megatron 张量并行 MLP 架构¶

输入 x 对所有设备相同(来自上一个 AllReduce 或 Identity)。中间激活需要经过 GeLU,而 GeLU 是非线性,需要完整数据,所以第一个线性层的输出在进入 GeLU 前是列切分的,恰好可以直接独立应用 GeLU(因为 GeLU 是元素级操作,不改变分片)。之后第二个线性层行切分,局部输出需 AllReduce 求和得到最终输出。
10.2 前向通信¶
# 假设 x 已经在所有 TP 设备上相同(例如,从上一层 AllReduce 得到)
def mlp_forward(x, A_parts, B_parts):
# x: [batch, seq_len, d](假设每个设备都有完整副本)
# A_parts[i], B_parts[i] 是设备 i 的权重分片
# 每个设备独立计算:
# 第一步:x @ A_i -> shape [batch, seq_len, 4d/N]
h_local = x @ A_parts[rank] # 本地矩阵乘,无通信
# 应用 GeLU(元素级)
h_act = gelu(h_local) # 无通信
# 第二步:h_act @ B_i -> shape [batch, seq_len, d]
y_local = h_act @ B_parts[rank] # 本地矩阵乘,无通信
# AllReduce 求和得到完整输出 y
y = allreduce(y_local, op=SUM) # 每个设备都得到相同的 y
return y
通信:仅一次 AllReduce(求和),在每个设备上将 y_local 求和得到 y。注意 y_local 是部分和,因为行切分的输出在特征维上是完整维度,但在设备间是部分和。
10.3 反向通信¶
def mlp_backward(x, A_parts, B_parts, dy):
# dy: 上游梯度,与 y 同形状,所有设备相同(因为前向 AllReduce 后,每个设备都有 y,反向时 dy 自然相同)
# 每个设备独立计算对 B 的梯度和对 h_act 的局部梯度
# h_act 是前向时本地的激活(形状 [batch, seq_len, 4d/N])
# 计算 dB_local = h_act^T @ dy (本地)
dB_i = h_act.T @ dy # [4d/N, d]
# 计算对 h_act 的局部梯度 dh_local = dy @ B_i^T (形状 [batch, seq_len, 4d/N])
dh_local = dy @ B_parts[rank].T # 本地
# dh_local 此时是列切分的(对应完整的 h_act 被列切分),而完整的 dh 需要沿着列拼接
# 但是在 Megatron 中,由于 A 是列切,完整的 dh 应该被看作列切,每个设备持有对应列的部分,这样可以直接用于 A 的梯度计算。
# 实际上不需要通信:dh_local 可以直接用于计算 dA_i。
# 计算 dA_i = x^T @ dh_local (本地)
dA_i = x.T @ dh_local # [d, 4d/N]
# 现在需要计算 dx:完整 dx = sum_i (dh_local_i @ A_i^T) 这里 dh_local_i 是上面的 dh_local,形状 [batch, seq_len, 4d/N],A_i^T 形状 [4d/N, d]
dx_local = dh_local @ A_parts[rank].T # [batch, seq_len, d] 本地部分贡献
# AllReduce 求和得到完整 dx
dx = allreduce(dx_local, op=SUM) # 每个设备得到相同的 dx
return dx, dA_i, dB_i
通信:一次 AllReduce (dx_local 求和)。
总结:MLP 前向一次 AllReduce(y),反向一次 AllReduce(dx)。激活函数 GeLU 因为是独立元素操作,不需要通信。
写出张量并行中自注意力层的通信模拟:QKV 投影列切分,输出投影行切分,描述每个子层的通信操作及数据流动。¶
11.1 自注意力层的张量并行¶

11.2 前向通信模拟¶
def attention_forward(x, QKV_weights_parts, O_weights_parts):
# x 所有设备相同
# 每个设备计算自己的 Q,K,V (列切)
Q_i = x @ Q_weights_parts[rank] # [B, L, (H/N)*d_head]
K_i = x @ K_weights_parts[rank]
V_i = x @ V_weights_parts[rank]
# 重塑为多头并计算注意力(各设备独立,无通信)
attn_out_i = scaled_dot_product_attention(Q_i, K_i, V_i) # [B, L, (H/N)*d_head]
# 输出投影行切分
y_local = attn_out_i @ O_weights_parts[rank] # [B, L, d]
# AllReduce 求和
y = allreduce(y_local, op=SUM) # 每个设备得到相同的完整输出
return y
11.3 反向通信模拟¶
上游梯度 dydy 所有设备相同(因为前向 AllReduce 后每个设备都有 y)。反向时:
-
对 O 权重的梯度:
dW_O_i = attn_out_i^T @ dy,本地计算。 -
对 attn_out_i 的梯度:
d_attn_i = dy @ W_O_i^T,形状[B, L, (H/N)*d_head]。这个梯度是每个设备对自己负责头的注意力输出的梯度,可以继续本地反向传播通过注意力计算得到 dQ_i, dK_i, dV_i(均本地)。 -
接着计算 dx 的局部贡献:
dx_local = dQ_i @ W_Q_i^T + dK_i @ W_K_i^T + dV_i @ W_V_i^T。每个设备的 dx_local 是完整 dx 的部分和,需要 AllReduce 求和 得到完整 dx。
def attention_backward(x, dy, QKV_weights, O_weights, attn_outs):
# dy 所有设备相同
d_attn_i = dy @ O_weights[rank].T
# 反向传播注意力 (本地)
dQ_i, dK_i, dV_i = attention_backward_local(Q_i, K_i, V_i, d_attn_i)
# 本地权重梯度
dW_O_i = attn_outs[rank].T @ dy
dW_Q_i = x.T @ dQ_i
dW_K_i = x.T @ dK_i
dW_V_i = x.T @ dV_i
# dx 局部贡献
dx_local = dQ_i @ Q_weights[rank].T + dK_i @ K_weights[rank].T + dV_i @ V_weights[rank].T
# AllReduce 求和
dx = allreduce(dx_local, op=SUM)
return dx, dW_O_i, dW_Q_i, dW_K_i, dW_V_i
总结:自注意力在张量并行下前向仅需一次 AllReduce(y),反向一次 AllReduce(dx)。这种设计使得注意力层的通信量非常小(每个 AllReduce 的数据量是 batch*seq_len*d),且利用了注意力的头独立特性。
模拟流水线并行中 1F1B 调度的通信模式:定义多个 stage,每个 stage 位于不同节点,用 Send/Recv 传递激活和梯度,实现 micro-batch 的前向和反向交错执行。¶
12.1 1F1B 调度原理¶
流水线并行将模型按层分为多个 stage,每个 stage 放在一个设备上。1F1B(one forward one backward)调度交替执行前向和反向,以减少流水线气泡。具体:首先预热阶段(warm-up),stage 0 接收输入,执行前向,将激活发送给 stage 1,stage 1 执行前向并发送给 stage 2,以此类推。每个 stage 处理完一个 micro-batch 的前向后,如果收到上游的梯度,就开始执行反向,并将梯度发回给前一个 stage。每个 stage 在稳定状态下保持 1 个前向和 1 个反向交替执行。
12.2 模拟设计¶
假设有 S 个 stage,M 个 micro-batch。每个 stage 是一个对象,拥有其模型部分,以及发送/接收队列(用列表模拟)。通信通过 Send/Recv 进行,相邻 stage 之间传递激活(前向)和梯度(反向)。我们模拟一个迭代的完整过程。
用 Python 模拟,定义 Stage 类,以及调度器控制每一步。这里给出伪代码和关键逻辑。
class PipelineStage:
def __init__(self, stage_id, num_stages):
self.id = stage_id
self.next = stage_id + 1 if stage_id < num_stages - 1 else None
self.prev = stage_id - 1 if stage_id > 0 else None
self.activations = {} # microbatch_id -> activation tensor
self.grads = {}
# 模拟的模型层
def forward(self, x, mb_id):
# 计算,返回激活,保存用于反向
act = self.model(x)
self.activations[mb_id] = act
return act
def backward(self, grad, mb_id):
act = self.activations.pop(mb_id)
# 反向计算,返回对输入的梯度
dx = self.model.backward(act, grad)
return dx
调度器:
def simulate_1f1b(num_microbatches, num_stages):
stages = [PipelineStage(i, num_stages) for i in range(num_stages)]
# 每个 stage 维护一个前向任务队列和反向任务队列
# 简化:使用一个全局步循环,每个步模拟一个时钟周期,在各 stage 上执行一个操作
# 这里更直观:我们按照1F1B算法描述伪代码
warmup_steps = num_stages - 1
# 预热阶段:前向传播逐级启动
for mb in range(num_microbatches):
# stage 0 总是前向第一个 micro-batch,但需要等待前一 micro-batch 的完成
# 使用循环逐步推进
pass
更简单明确的模拟:我们使用一个数组记录每个 stage 当前处理的任务,按照1F1B规则推进。由于纯伪代码篇幅较长,这里概括核心通信:
-
前向:stage i 收到 stage i-1 的激活(对于第一个 micro-batch,stage 0 从输入读取),计算后发送给 stage i+1。
-
反向:stage i 收到 stage i+1 的梯度,计算后发送给 stage i-1。
-
每个 stage 在完成一个前向后,如果已有待处理的反向,则执行反向;否则等待新的前向或反向。
1F1B 的通信特征:每个 micro-batch 在每个 stage 边界上产生两次通信(前向激活、反向梯度),总通信量与 micro-batch 数和 stage 数成正比。
编写基于 PyTorch 分布式包(torch.distributed)的简易通信测试:初始化进程组,完成一轮 AllReduce 和 Barrier,输出每 rank 结果。¶
13.1 实现¶
import torch
import torch.distributed as dist
import os
def init_process(rank, world_size, func):
""" 初始化分布式环境并执行 func """
os.environ['MASTER_ADDR'] = '127.0.0.1'
os.environ['MASTER_PORT'] = '29500'
dist.init_process_group('nccl', rank=rank, world_size=world_size)
func(rank, world_size)
dist.destroy_process_group()
def test_allreduce(rank, world_size):
tensor = torch.tensor([rank + 1.0]).cuda(rank)
print(f"Rank {rank} before AllReduce: {tensor}")
dist.all_reduce(tensor, op=dist.ReduceOp.SUM)
print(f"Rank {rank} after AllReduce: {tensor}")
dist.barrier()
if rank == 0:
print("AllReduce test passed.")
if __name__ == "__main__":
world_size = 4
import torch.multiprocessing as mp
mp.spawn(init_process, args=(world_size, test_allreduce), nprocs=world_size)
此代码演示了基本的 AllReduce 和 Barrier。每个 rank 的本地数值被求和,结果在所有 rank 上相同(SUM 的结果)。
用 NCCL 的 Python 绑定(cupy/nccl 或 mock)模拟 AllReduce 的 Ring 和 Tree 算法选择逻辑,并对比它们的延迟和带宽。¶
14.1 NCCL 算法选择¶
NCCL 根据数据大小和节点拓扑自动选择 AllReduce 算法:小数据量用 Tree,大数据量用 Ring。Tree 延迟低(对数步数),但带宽利用率不如 Ring(Ring 的带宽利用率接近 100%)。我们可以模拟两种算法的通信时间。
14.2 模拟代码(不依赖 NCCL)¶
import math
def ring_allreduce_time(data_size_bytes, num_gpus, bandwidth_Bps, latency_sec):
# Ring 需要 2*(N-1) 步,每一步传输 data_size/N 字节(理想情况)
chunk = data_size_bytes / num_gpus
steps = 2 * (num_gpus - 1)
transfer_time = chunk / bandwidth_Bps * steps
total_latency = latency_sec * steps
return transfer_time + total_latency
def tree_allreduce_time(data_size_bytes, num_gpus, bandwidth_Bps, latency_sec):
# 二叉树:每一步传输 data_size 数据,需要 2*log2(N) 步(reduce + broadcast)
steps = 2 * math.ceil(math.log2(num_gpus))
transfer_time = data_size_bytes / bandwidth_Bps * steps
total_latency = latency_sec * steps
return transfer_time + total_latency
# 示例
size = 1024**2 # 1 MB
n = 8
bw = 10e9 # 10 GB/s
lat = 1e-6 # 1 us
print("Ring:", ring_allreduce_time(size, n, bw, lat))
print("Tree:", tree_allreduce_time(size, n, bw, lat))
实际 NCCL 在数据量小时用 Tree 以减少延迟,数据量大时用 Ring 以最大化带宽。模拟可体现交叉点。
设计一个通信模拟器,输入网络拓扑(如交换机连接)和通信原语,估算完成一次 AllReduce 的时间,要求考虑带宽、延迟和并发。¶
15.1 模拟器设计¶
输入:集群节点数,每个节点的 GPU 数,交换机拓扑(如胖树),链路带宽、延迟,通信原语类型(AllReduce、AllGather等),数据大小。
输出:预估完成时间。
核心思想:将通信原语分解为基本 Send/Recv 操作,根据拓扑计算每个数据包的路由路径,考虑带宽共享和网络拥塞。实际实现较复杂,这里给出简化的模型:
-
采用 Ring 算法时,AllReduce 时间 = 2*(N-1) * (latency + chunk_size / min(bandwidth, 链路瓶颈))。
-
考虑节点内 GPU 通信(NVLink)和节点间通信(IB/RoCE)的带宽差异,可分层计算。
伪代码:
class CommSimulator:
def __init__(self, topology, link_bw, link_lat):
self.topo = topology
self.bw = link_bw
self.lat = link_lat
def estimate_allreduce(self, data_size, algorithm='ring'):
if algorithm == 'ring':
steps = 2 * (self.topo.num_gpus - 1)
# 假设所有链路的有效带宽为最慢链路
bw = min(self.bw.values())
chunk = data_size / self.topo.num_gpus
time = steps * (self.lat + chunk / bw)
return time
更精确需模拟消息级传输。
模拟 ZeRO-1 的优化器状态分片后的通信:在数据并行组内对梯度进行 ReduceScatter,每个节点更新自己分片的优化器状态,再通过 AllGather 获取完整参数。¶
16.1 ZeRO-1 原理¶
ZeRO-1 将优化器状态(Adam 的 m, v)切分到各个 DP 节点上。每个节点只存储与自己负责参数分片对应的优化器状态。训练时:
-
每个节点计算本地梯度(全部参数)。
-
对梯度执行 ReduceScatter,使得每个节点得到自己负责那一部分参数的全局平均梯度(规约结果)。
-
每个节点使用该部分梯度更新自己本地的优化器状态和参数(只更新自己负责的部分)。
-
通过 AllGather 收集所有节点更新后的参数部分,还原完整模型参数。
16.2 模拟¶
def zero1_step(params, grads, optimizer_states, rank, world_size):
"""
params: 完整模型参数列表(每个节点开始时都相同)
grads: 完整梯度列表(每个节点计算得到)
optimizer_states: 每个节点只持有自己负责分片的优化器状态,假设均匀切分
"""
# 1. ReduceScatter 梯度
# 每个梯度被均匀切成 world_size 块,各节点得到对应块的规约和
scattered_grads = []
for g in grads:
chunks = list(torch.chunk(g, world_size, dim=0)) # 按某个维切分
# 对每个块做 allreduce 并只保留对应 rank 的块
# 用 AllReduce + 本地取块模拟 ReduceScatter
# 实际 ReduceScatter 是高效的原语,这里模拟:
reduced = allreduce_chunk(chunks, rank) # 返回该节点负责的那一块
scattered_grads.append(reduced)
# 2. 更新本地优化器状态和参数
for i, g_chunk in enumerate(scattered_grads):
# 取出对应参数分片
param_chunk = get_param_chunk(params[i], rank, world_size)
# 更新 optimizer state 和 param_chunk
update_with_adam(param_chunk, g_chunk, optimizer_states[i])
# 写回分片
set_param_chunk(params[i], rank, param_chunk)
# 3. AllGather 参数,恢复完整模型
for i, p in enumerate(params):
gathered = allgather_chunks(p, world_size) # 收集所有分片
params[i].data.copy_(gathered)
通信量:ReduceScatter 和 AllGather 各一次,每次传输的数据量为梯度/参数总大小。与标准数据并行(一次 AllReduce)通信量相当,但优化器状态显存大幅减少。
模拟 ZeRO-2 的梯度分片:反向计算后梯度已分片,通过 ReduceScatter 归约梯度到对应节点,更新参数后 AllGather 参数,写出通信步骤。¶
17.1 ZeRO-2 原理¶
ZeRO-2 在 ZeRO-1 基础上进一步将梯度也分片。每个节点只存储与自己优化器状态分片对应的梯度分片。反向计算时,每个节点计算完整的梯度,但随后立即通过 ReduceScatter 将规约后的梯度分散到各节点,每个节点只保留自己负责那部分梯度。之后的参数更新和 AllGather 与 ZeRO-1 相同。
17.2 通信步骤模拟¶
def zero2_step(params, grads, optimizer_states, rank, world_size):
# grads 完整,但需要被分片规约
# 1. ReduceScatter 梯度:每个节点贡献完整梯度,得到自己负责分片的规约梯度
scattered_grads = []
for g in grads:
scattered = reduce_scatter_chunk(g, world_size, rank)
scattered_grads.append(scattered)
# 2. 用 scattered_grads 更新本地参数分片和优化器状态
for i, g_chunk in enumerate(scattered_grads):
param_chunk = get_chunk(params[i], rank)
update_adam(param_chunk, g_chunk, optimizer_states[i])
# 3. AllGather 参数
for i, p in enumerate(params):
gathered = allgather_chunks(p)
params[i].copy_(gathered)
ZeRO-2 相比 ZeRO-1 进一步减少了显存(梯度也不全存),通信量仍然是 ReduceScatter + AllGather,与 ZeRO-1 相同。反向计算时需要完整参数,但参数已在上一轮 AllGather 后完整。
模拟 ZeRO-3 的前向参数收集:模型参数被分片,前向时遇到某个层,通过 AllGather 收集该层完整参数,用后立刻释放;反向时再次 AllGather 参数计算梯度,然后 ReduceScatter 梯度。¶
18.1 ZeRO-3 原理¶
ZeRO-3 将模型参数也分片,每个节点只持有一部分参数。在计算某一层的前向时,该节点需要该层的完整参数,因此临时通过 AllGather 从所有节点收集该层的完整参数,计算完后立即释放(除非后续还需用)。反向时同样需要 AllGather 参数来计算梯度,得到完整梯度后,每个节点只保留自己负责的分片部分,通过 ReduceScatter 将梯度规约并分散到对应节点上,用于更新本地参数分片。
18.2 通信步骤模拟¶
def zero3_forward(layers, input_data, rank, world_size):
x = input_data
for layer in layers:
# 该层参数分片在各个节点上,需要收集完整参数
full_param = allgather_layer_param(layer.param_id, world_size) # 每节点提供自己分片,组合成完整参数
# 前向计算
x = layer.forward(x, full_param)
# 释放完整参数(删除或标记为可回收)
del full_param
return x
def zero3_backward(layers, grad_output, rank, world_size):
gy = grad_output
for layer in reversed(layers):
# 再次收集完整参数
full_param = allgather_layer_param(layer.param_id, world_size)
# 反向计算得到对参数的梯度(完整)和对输入的梯度
gx, g_param_full = layer.backward(gy, full_param)
del full_param
# 对参数的完整梯度进行 ReduceScatter,每个节点只保留自己分片
local_grad_chunk = reduce_scatter_chunk(g_param_full, world_size, rank)
# 用 local_grad_chunk 更新本地参数分片
update_local_param(layer, local_grad_chunk)
gy = gx
return gy
ZeRO-3 的通信量较大:每层前向一次 AllGather,反向一次 AllGather + 一次 ReduceScatter。但通过细粒度的分片和及时的释放,使得极大的模型能在有限的显存上训练。
编写一个通信开销分析器,给定模型各层大小、并行策略(DP/TP/PP)、硬件带宽和延迟,估算一次训练迭代的通信总时间。¶
19.1 分析器设计¶

19.2 估算方法¶
-
数据并行(AllReduce 梯度):使用 Ring AllReduce 模型,数据量 = 模型参数量 * 数据精度字节。时间 = 2(DP-1)(L + 参数量/(DP * B_{intra}))(假设同节点)。
-
张量并行:前向/反向中的 AllReduce/ReduceScatter。每个 AllReduce 的数据量是激活大小(如 batchseqd)。层内多次通信。
-
流水线并行:点对点 Send/Recv 激活和梯度,数据量 = batchseqd,延迟为 L + 数据量/B_{inter}(假设跨节点)。1F1B 下每个 micro-batch 在 stage 边界传输 2 次。
编写一个类,接收这些参数,计算汇总时间。简化代码:
class CommAnalyzer:
def __init__(self, dp, tp, pp, intra_bw, inter_bw, latency):
self.dp = dp; self.tp = tp; self.pp = pp
self.intra_bw = intra_bw; self.inter_bw = inter_bw
self.lat = latency
def allreduce_time(self, data_bytes, num_gpus, is_inter=False):
bw = self.inter_bw if is_inter else self.intra_bw
steps = 2 * (num_gpus - 1)
return steps * (self.lat + data_bytes / (num_gpus * bw))
def send_recv_time(self, data_bytes, is_inter=True):
bw = self.inter_bw if is_inter else self.intra_bw
return self.lat + data_bytes / bw
def estimate_iteration(self, model_size, batch_seq_dim, num_layers):
total = 0
# Data parallel: AllReduce gradients
grad_size = model_size * 2 # FP16
total += self.allreduce_time(grad_size, self.dp, is_inter=(self.dp>8)) # 假设跨节点
# Tensor parallel: 每层两次 AllReduce (MLP) + 两次 (Attention) 等等,简化为每层 4 次
msg_size = batch_seq_dim * 2 # FP16
for _ in range(num_layers):
# 假定的 AllReduce 次数
total += 4 * self.allreduce_time(msg_size, self.tp, is_inter=False) # TP 通常同节点
# Pipeline parallel: 每层在 stage 边界的通信 (1F1B)
if self.pp > 1:
# 每个 micro-batch 在每个边界传输 2 次
total += (self.pp-1) * num_microbatches * 2 * self.send_recv_time(batch_seq_dim*2)
return total
模拟数据并行中异步更新(ASGD)的通信:每个 worker 计算梯度后异步发送到参数服务器,参数服务器汇总后更新并返回新参数,worker 不等待完成就继续计算。¶
20.1 异步 SGD 原理¶
在异步数据并行中,存在一个或多个参数服务器(PS),保存全局模型参数。每个 worker 独立从 PS 拉取最新参数,计算本地梯度,然后将梯度推送回 PS,不等待其他 worker。PS 收到梯度后,应用更新到全局模型,并可能将新参数发送给请求的 worker。这会导致梯度延迟(staleness),但能提高并行度和吞吐量。
20.2 模拟实现¶
import threading
import time
import numpy as np
class ParameterServer:
def __init__(self, param_shape):
self.params = np.zeros(param_shape)
self.lock = threading.Lock()
def push_gradient(self, grad, lr=0.01):
with self.lock:
self.params -= lr * grad
def get_params(self):
with self.lock:
return self.params.copy()
class Worker(threading.Thread):
def __init__(self, worker_id, ps, data_loader):
super().__init__()
self.worker_id = worker_id
self.ps = ps
self.data_loader = data_loader
def run(self):
for data in self.data_loader:
# 1. 拉取最新参数(可能不是最新的,有延迟)
params = self.ps.get_params()
# 2. 计算梯度(模拟)
grad = compute_gradient(params, data)
# 3. 异步推送梯度,不等待
self.ps.push_gradient(grad)
# 继续下一轮,无需等待其他worker
模拟时,多个 worker 线程并发运行,参数服务器加锁保证一致性。由于异步特性,worker 可能会用旧参数计算梯度,导致更新冲突,但能加速训练。