如何设置分布式训练批次

GPU
小华
2026-08-19

设置“分布式训练批次(Distributed Batch Size)”是深度学习训练中的关键步骤,它直接影响模型收敛速度和显存占用。核心原则是:总批次大小(Global Batch Size) = 单卡批次大小(Per-GPU Batch Size) × 显卡数量(World Size),并且通常需要配合学习率缩放(Linear Scaling Rule)

以下是针对不同框架的详细设置指南和核心概念解析。

一、 核心概念

  1. Global Batch Size (总批次大小): 模型在一个训练迭代(step)中看到的总样本数。
  2. Local Batch Size (单卡批次大小): 每张显卡一次处理的数据量。
  3. 梯度累积 (Gradient Accumulation): 如果单卡显存不足以支撑理想的 Local Batch Size,可以通过多次前向传播累积梯度,再进行一次反向传播。此时:Global Batch Size = Local Batch Size × World Size × Accumulation Steps

二、 PyTorch 设置方法

PyTorch 推荐使用 torchrun(原 torch.distributed.launch)或 accelerate 库。

1. 原生 PyTorch (DDP)

步骤 A: 初始化进程组

import torch
import torch.distributed as dist
import os
# 通常在启动脚本中设置
dist.init_process_group(backend='nccl') # GPU用nccl, CPU用gloo
local_rank = int(os.environ["LOCAL_RANK"])
torch.cuda.set_device(local_rank)

步骤 B: 包装模型

model = YourModel().to(local_rank)
model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[local_rank])

步骤 C: 使用 DistributedSampler (关键)
这是设置批次的核心。Sampler 确保每张卡拿到数据不重复,且覆盖整个数据集。

from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
dataset = YourDataset()
sampler = DistributedSampler(dataset)
# 注意:这里的 batch_size 是单卡的 Batch Size
dataloader = DataLoader(dataset, batch_size=32, sampler=sampler, shuffle=False)
# 注意:使用 DistributedSampler 时,Dataloader 的 shuffle 必须设为 False

步骤 D: 训练循环中的调整

for epoch in range(epochs):
sampler.set_epoch(epoch) # 确保每个epoch数据打乱方式不同
for batch in dataloader:
inputs, labels = batch
inputs = inputs.to(local_rank)
labels = labels.to(local_rank)
outputs = model(inputs)
loss = criterion(outputs, labels)
# 如果使用了 DDP,loss 已经是各卡平均后的结果(取决于 reduction='mean')
loss.backward()
optimizer.step()
optimizer.zero_grad()

启动命令:

torchrun --nproc_per_node=4 train.py
# --nproc_per_node=4 表示使用4张卡

2. 使用 Hugging Face Accelerate (最简单)

Accelerate 自动处理了 DDP 和 Sampler。

from accelerate import Accelerator
accelerator = Accelerator()
# 初始化
model, optimizer, dataloader = accelerator.prepare(model, optimizer, dataloader)
# 训练循环不需要改太多,直接跑
for batch in dataloader:
inputs, labels = batch
outputs = model(inputs)
loss = criterion(outputs, labels)
accelerator.backward(loss) # 使用 accelerator 的 backward
optimizer.step()
optimizer.zero_grad()

启动:

accelerate config # 第一次配置
accelerate launch train.py

三、 TensorFlow / Keras 设置方法

TensorFlow 主要通过 tf.distribute.Strategy 实现。

1. 使用 MirroredStrategy (单机多卡)

import tensorflow as tf
# 设置策略
strategy = tf.distribute.MirroredStrategy()
# 在策略作用域内定义模型和优化器
with strategy.scope():
model = create_model()
optimizer = tf.keras.optimizers.Adam()
# 准备数据
dataset = your_dataset_fn()
# 关键:通过 strategy.experimental_distribute_dataset 包装
dist_dataset = strategy.experimental_distribute_dataset(dataset)
# 定义训练步
@tf.function
def train_step(inputs):
images, labels = inputs
with tf.GradientTape() as tape:
predictions = model(images, training=True)
loss = compute_loss(labels, predictions)
gradients = tape.gradient(loss, model.trainable_variables)
optimizer.apply_gradients(zip(gradients, model.trainable_variables))
# 训练循环
for epoch in range(epochs):
for batch in dist_dataset:
strategy.run(train_step, args=(batch,))

注意: 在 Keras 中,如果你直接 model.fit(),TensorFlow 会自动处理批次分配。你只需要设置 batch_size单卡的批次大小即可。

四、 学习率调整策略 (非常重要)

当你增加 GPU 数量(即增加 Global Batch Size)时,必须调整学习率,否则模型可能不收敛或震荡。
线性缩放规则 (Linear Scaling Rule):
如果 Baseline 配置是:1 张卡,Batch Size = 32,学习率 = 0.001。
现在你有 8 张卡,Global Batch Size = 32 * 8 = 256。
那么新的学习率应为:$0.001 \times 8 = 0.008$。
代码示例 (PyTorch):

base_lr = 0.001
world_size = torch.distributed.get_world_size()
optimizer = torch.optim.Adam(model.parameters(), lr=base_lr * world_size)

Warmup:

当 Global Batch Size 非常大(例如 > 1024)时,直接线性缩放可能也不够稳定,通常需要配合 Warmup(学习率预热)策略,前几百步从 0 慢慢升到目标学习率。

五、 常见问题与检查清单

  1. 数据重复/丢失: 如果不使用 DistributedSampler (PyTorch),每张卡会加载相同的数据,导致等效 Batch Size 没变,且数据浪费。
  2. Batch Normalization:
  • 在 DDP 中,默认的 BatchNorm 是单卡独立计算的。
  • 如果单卡 Batch Size 很小(如 < 4),BN 层会失效。此时应使用 SyncBatchNorm
  • PyTorch: model = torch.nn.SyncBatchNorm.convert_sync_batchnorm(model)
  1. 梯度累积与 DDP:
  • 如果你显存小,想模拟大 Batch。
  • 逻辑:每 4 步更新一次权重。
  • 注意: 在 DDP 中,梯度累积时,loss.backward() 会自动累积,不需要额外处理,只需控制 optimizer.step() 的频率。
  1. 验证集: 验证时通常不需要 DistributedSampler,或者使用 DistributedSampler 但设置 shuffle=False,且不需要同步梯度。

总结公式

场景设置方式
单卡 Batch Size设为 B
DDP 数据加载Sampler 自动切分数据
实际总 Batch Size$B \times N_{gpus}$
学习率设置$LR_{base} \times N_{gpus}$
BN 层处理使用 SyncBatchNorm (如果单卡 B 很小)
亿速云提供售前/售后服务

售前业务咨询

售后技术保障

400-100-2938

7*24小时售后电话

官方微信小程序