系列第三篇,深入讲解 All-Reduce、Ring-AllReduce、ZeRO 分片策略,以及 NCCL/RCCL 实现梯度同步的底层原理和代码实现。
CUDA 使用 NCCL(NVIDIA Collective Communications Library),而 ROCm 使用 RCCL(ROCm Collective Communications Library)。幸运的是,它们的 API 签名完全一致,只需改前缀(nccl 换 rccl)。我们可以用预处理器宏将它们统一起来。
// Unified header selection
#ifdef __HIP_PLATFORM_AMD__
#include <rccl/rccl.h>
#define COMM_ID rcclUniqueId
#define COMM_INIT rcclCommInitRank
#define COMM_ALL_REDUCE rcclAllReduce
#define COMM_GET_ERROR rcclGetErrorString
#else
#include <nccl.h>
#define COMM_ID ncclUniqueId
#define COMM_INIT ncclCommInitRank
#define COMM_ALL_REDUCE ncclAllReduce
#define COMM_GET_ERROR ncclGetErrorString
#endif
初始化通信器需要一个唯一 ID,然后广播给所有 rank(通常通过 MPI,但单节点可以用简单的环境变量):
COMM_ID id;
if (rank == 0) { // rank 0 generates the ID
COMM_GET_UNIQUE_ID(&id);
}
// In a real cluster, you broadcast this via MPI_Bcast.
// For single-node, we just pass it directly.
COMM_COMM_T comm;
COMM_INIT(&comm, world_size, id, rank);
在数据并行(Data Parallelism, DP)中,每个 GPU 都持有模型的完整副本。我们向每个 GPU 喂入不同的 micro-batch,计算本地梯度(dW_local),然后在所有 GPU 上对它们求平均。
数学操作是:dW_global = (1 / world_size) * Σ dW_local_i
这正是一个带 SUM 算子的 All-Reduce 操作(我们稍后除以 world_size,或者支持 AVG 的话直接用)。
一种朴素的做法是使用一个中心"服务器" GPU 接收所有梯度、求平均后再发回去。这会造成瓶颈。我们改用 Ring-AllReduce 算法,它是带宽最优的。
Ring 算法在 N 个以逻辑环连接的 GPU 上分两个阶段工作:
Scatter-Reduce(N-1 步):每个 GPU 将自己梯度的一块发送给下一个 GPU,同时接收另一块,累加求和。
All-Gather(N-1 步):GPU 将已经聚合好的块沿环转发,直到每个 GPU 都拥有完整的和。
不用从零写,我们使用高度优化的 COMM_ALL_REDUCE:
void sync_gradients(float* d_gradient, int num_elements,
ncclComm_t comm, hipStream_t stream) {
// In-place All-Reduce: sums gradients across all GPUs
// Result is stored in d_gradient on every rank.
COMM_ALL_REDUCE((const void*)d_gradient, // sendbuff
(void*)d_gradient, // recvbuff (in-place)
num_elements,
COMM_FLOAT, // data type
COMM_SUM, // operation
comm, stream);
// Average the sum to get the mean gradient
int world_size;
COMM_COMM_COUNT(comm, &world_size);
float inv_world = 1.0f / world_size;
scale_kernel<<<(num_elements+255)/256, 256, 0, stream>>>(d_gradient, inv_world, num_elements);
}
关于 AMD 的说明:RCCL 使用完全相同的函数签名。只需在链接标志中将 nccl 替换为 rccl(-lrccl vs -lnccl)。
数据并行很棒,但每个 GPU 仍然存储整个模型、优化器状态和梯度。对于 175B 参数,仅权重就需要约 3.2TB(FP32)。微软提出的 ZeRO 将这些组件分片到各个 GPU 上,以减少内存占用。
在第 2 部分中,我们为 AdamW 存储了 m(动量)和 v(方差)。它们和模型权重一样大。在阶段 1 中,我们对这些优化器状态进行分区:GPU 0 持有参数 0 到 N/2 的优化器状态,GPU 1 持有剩下的部分。
对我们的 AdamW kernel 的实现修改:不再在整个参数上启动 kernel,而是只在本地区分上启动。
__global__ void adamw_update_sharded_kernel(float* W, float* dW, float* m, float* v,
int total_params, int rank, int world_size,
...) {
int idx = blockIdx.x * blockDim.x + threadIdx.x + rank * (total_params / world_size);
if (idx >= total_params) return;
// ... (same update logic as Part 2)
}
在反向传播期间,我们必须 All-Gather 更新后的权重,以便每个 GPU 在下一个前向传播之前都有完整的更新后模型。这在优化器之后增加了一个通信步骤。
阶段 2 更进一步:它还分片梯度(dW)。在 All-Reduce 之前,我们只 reduce 属于本 GPU 分区的梯度。这将通信量减少多达一半。
带 ZeRO-2 的训练步骤伪代码:
// Forward: All GPUs have full weights (via All-Gather after previous step)
forward_pass();
// Backward: Compute local gradients (full size)
backward_pass();
// Reduce-Scatter: Each GPU sums only its assigned partition of the gradients
reduce_scatter_gradients(dW, local_partition, ...);
// Update: Update only local weights and local optimizer states
adamw_update_local(W_local, dW_local, m_local, v_local, ...);
// All-Gather: Sync the updated weights so everyone has the full model
all_gather_weights(W_full, W_local, ...);
这是 DeepSpeed 和 FairScale 等现代框架的核心。
对于真正巨大的模型(例如 > 1 万亿参数),即使 ZeRO-3 也不够。我们需要模型并行(Model Parallelism, MP),即线性层本身被拆分到各个 GPU 上。
张量并行(Tensor Parallelism, TP):按行或按列拆分权重矩阵。例如 Y = X * W,其中 W 按列拆分。这需要在每个线性层之后进行一次 All-Reduce。
流水线并行(Pipeline Parallelism, PP):将层拆分到各个 GPU 上(例如 GPU 0 持有层 1-10,GPU 1 持有层 11-20)。这需要通过 P2P(点对点)通信在 GPU 之间发送激活值(前向)和梯度(反向)。
以下是在 NCCL/RCCL 中启动 P2P 发送/接收的方式:
// GPU 0 sends activations to GPU 1
if (rank == 0) {
COMM_SEND(d_activation, size, comm, 1, stream);
} else if (rank == 1) {
COMM_RECV(d_activation, size, comm, 0, stream);
}
3D 并行范式(DP + TP + PP)是训练 GPT-4 和 Llama 3 的秘诀。你在节点内应用 TP(使用 NVLink/Infinity Fabric),跨节点应用 PP,跨节点组应用 DP。
通信是昂贵的。隐藏它的最好方式是重叠。当 GPU 正在计算第 10 层的反向传播时,我们可以在后台开始第 1 层梯度的 All-Reduce。
我们使用 Stream 和基于 Event 的同步来实现这一点:
hipStream_t compute_stream, comm_stream;
hipStreamCreate(&compute_stream);
hipStreamCreate(&comm_stream);
// During backward pass:
for (int layer = 0; layer < num_layers; ++layer) {
// Compute gradient for this layer on compute_stream
backward_layer<<<..., compute_stream>>>(...);
// Record an event when gradient is ready
hipEventRecord(grad_ready_event, compute_stream);
// Make comm_stream wait for the event
hipStreamWaitEvent(comm_stream, grad_ready_event, 0);
// Launch All-Reduce on comm_stream (async)
COMM_ALL_REDUCE(..., comm_stream);
}
这样,GPU 几乎不会空闲等待网络数据包;计算填满了那些空隙。
将我们学到的所有东西——ZeRO-2 分片、重叠通信和混合精度——结合起来,以下是一个分布式环境下单训练步骤的框架:
void distributed_train_step(float* d_model_weights, ... , int rank, int world_size,
ncclComm_t comm, hipStream_t comp_stream, hipStream_t comm_stream) {
// 1. Forward Pass (Full model on each GPU)
forward_pass(d_model_weights, d_activations, ..., comp_stream);
// 2. Backward Pass (Compute full gradients)
backward_pass(d_model_weights, d_gradients, ..., comp_stream);
// 3. Reduce-Scatter Gradients (ZeRO-2)
hipEventRecord(grad_ready, comp_stream);
hipStreamWaitEvent(comm_stream, grad_ready, 0);
int local_size = total_params / world_size;
COMM_REDUCE_SCATTER(d_gradients, d_gradients_local, ... , comm_stream);
hipEventRecord(comm_done, comm_stream);
hipStreamWaitEvent(comp_stream, comm_done, 0);
// 4. Update local weights & optimizer states (AdamW)
adamw_update_local<<<..., comp_stream>>>(d_local_weights, d_gradients_local,
d_momentum_local, d_variance_local, ...);
// 5. All-Gather to sync the full model
hipEventRecord(update_done, comp_stream);
hipStreamWaitEvent(comm_stream, update_done, 0);
COMM_ALL_GATHER(d_local_weights, d_full_weights, ... , comm_stream);
// 6. Synchronize the main stream at the very end
hipStreamSynchronize(comp_stream);
}
我们刚刚冲过了 GPU 优化马拉松的终点线。从第 1 部分的简单向量加法开始,我们构建了单 GPU transformer,在第 2 部分添加了带混合精度的完整训练栈,最后通过实现分布式数据并行、ZeRO 分片和使用 NCCL/RCCL 重叠通信,击碎了单 GPU 的天花板。
你现在拥有了驱动各大 AI 公司训练集群的架构蓝图。无论你是在家里的 4-GPU 工作站还是在 4000-GPU 超级计算机上运行,原则完全相同:减少通信,最大化计算,分片一切可以分片的东西。
接下来是什么?如果对第 4 部分有需求,我们可以探索 FP8 量化以实现更快的训练、融合多头注意力(Flash Attention)kernel 以进一步优化内存带宽,或者分布式 Checkpointing 以保存分片模型而不使文件系统崩溃。
请在下方评论告诉我你最感兴趣的是什么!
感谢你陪我完成这次深度探索。如果你在实现这些 kernel 时遇到任何错误,或者有特定场景(例如在 AMD MI300X 上使用 RCCL vs 在 NVIDIA H100 上使用 NCCL)想要我详细说明,请留言——我亲自回复所有评论。
下次见,保持你的 warp 收敛、带宽饱和!
下一场冒险再见!🚀