13.多 GPU / 分布式训练
第 13 章 — 多 GPU / 分布式训练
第 12 章用
set_seed保证了"同配置 = 同结果",但所有实验仍然只在一张卡上跑:模型一大就显存不够,数据一多就训练太慢,GPU 利用率上不去。本章解决"怎么让多张 GPU 一起干活":先给一个几行就能上手的DataParallel(单进程多卡,教学够用),再给 PyTorch 官方推荐的标准方案DistributedDataParallel(多进程、每进程一张卡,生产可用),并用config里的开关一键切换三种模式。
一、本章要解决的问题
- 之前:代码只能在一张卡上训练,
DataParallel/DDP这些词只听过没见过。 - 现在:模型太大装不进单卡显存、训练太慢等不起——我们需要让多张 GPU 协作,同时不破坏第 12 章的可复现性、不写出重复的 checkpoint/日志。
- 判断标准:改一行
config,就能在"单卡 → DataParallel → DDP"之间切换;跑 DDP 时终端只打印一份日志、只生成一份 best.pth,且训练速度随卡数提升。
二、核心概念速览
下面 10 个概念是本章代码的全部"生词"。先花 10 分钟读完,再看代码会轻松很多。
1. 单卡瓶颈:一个厨师炒菜 vs 多个厨师分工
一张 GPU 就是一位厨师:洗菜(数据预处理)、切菜(前向传播)、炒菜(反向传播)、装盘(更新参数)全由他一人完成。再快的厨师,锅(显存)只有一口,菜(数据)再多也只能一锅一锅炒——这就是"单卡瓶颈":算力、显存、数据搬运速度三者相互制约,总吞吐量上不去。多卡训练就是多请几位厨师,每人一口锅,各炒一份菜(各自处理数据分片),炒完把"调味"(梯度)汇总一下,保证每口锅的味道(模型参数)保持一致。
2. DataParallel 与 DDP 的本质区别
DataParallel(DP)像"一位大厨带几个帮厨":只有一个主进程、主卡当总指挥,每批数据被主卡拆开后分发出去,算完再汇总回主卡——主卡又搬数据又汇总,成了瓶颈,而且靠线程同步,速度有限。DistributedDataParallel(DDP)像"几位平级的大厨":启动 N 个进程、每个进程独占一张卡,各自维护一份完整的模型副本,独立处理自己的数据分片,只在反向传播后做一次梯度同步。DDP 没有主卡瓶颈、扩展性好,是官方推荐做法;DP 简单但仅作教学了解。
3. 进程与 rank:谁是几号员工
进程就是"一个正在运行的程序实例"。DDP 会同时启动 N 个进程,大家各干各的但必须协作,于是给每个进程发一个工号 rank(0、1、2…),rank 0 通常当"班长"(主进程),负责日志、存档这类只能做一次的事;world_size 是员工总数(进程数)。有了工号,程序才能判断"这件事该不该我来做",也才能给每个进程分配互不重叠的数据分片。
4. init_process_group:先开个碰头会
DDP 开工前,所有进程必须互相认识,torch.distributed.init_process_group 就是这场"碰头会":它告诉每个进程"你是谁(rank)、一共几人(world_size)、用什么方式通信(backend)、去哪找组织者(init_method 里的 TCP 地址)"。会开完,大家才知道和谁同步、等谁。训练结束再调 destroy_process_group 散会,释放通信资源。
5. DistributedSampler:数据分片 + set_epoch
DDP 下如果还用普通 DataLoader,每个进程都会读到同一批数据,等于"几位厨师炒同一锅菜"。DistributedSampler 先把训练集按 world_size 切成互不重叠的 N 份,每个进程只看到自己的那 1/N。但分片内的洗牌是"固定伪随机"的,如果每个 epoch 不重新洗,永远看到相同顺序——所以每个 epoch 开始前必须调 train_sampler.set_epoch(epoch),让各 epoch 的分片内顺序不同,数据增强才不会失效。
6. 梯度同步(allreduce):对账
每个进程独立前向反向后,它算出的梯度只反映自己那 1/N 的数据。为了让所有进程的模型保持一致,需要"对账":allreduce 把所有进程里同一参数的梯度求和再取平均,然后广播给每个人。这样每个进程拿到相同的平均梯度、更新出相同的参数,模型才始终"同味"。这个操作在 DDP 中由框架在 loss.backward() 时自动触发,不需要你手写。
7. 为什么只有 rank 0 存档/打日志/验证
如果每个进程都写 checkpoint,会产出 N 份重复文件,还可能同时抢写同一个文件互相踩踏;如果大家都打日志,日志会乱成一锅粥。更重要的是死锁风险:分布式是"所有进程一起等",如果每个进程都去写文件而文件系统卡住,整个训练就全部挂起。所以约定:存档、打日志、跑测试这类"只能做一次"的职责全部交给 rank 0,其他进程静默干活。
8. broadcast:把验证分数广播给所有人
验证集只有 rank 0 真正跑一遍,算出 val_loss 和 val_acc。但其他进程也要用这个分数决定"要不要更新 best、要不要早停"——数据不能只留在 rank 0 手里。于是 rank 0 把两个指标打包进一个张量,用 torch.distributed.broadcast(tensor, src=0) 广播给所有进程,确保每个人都看到同一个分数,才能做出完全一致的决策(都存 best、都停)。
9. Windows 上 gloo backend 与 spawn 的注意点
分布式通信有两种常用 backend:nccl(NVIDIA 官方,最快,但只支持 Linux + GPU)和 gloo(跨平台通用,CPU/GPU 都能跑)。Windows 上没有 nccl,所以本章统一用 gloo 保证"任何机器都能跑通"。另外,Windows 上不能用 fork 式多进程,代码里用 torch.multiprocessing.spawn 启动多个 worker;spawn 要求启动代码必须放在 if __name__ == "__main__": 保护之下,否则子进程会再次执行整份文件,造成无限递归启动。
10. DDP 下 batch_size 是每卡的值
配置里的 batch_size: 64 在 DDP 模式下指的是"每张卡每步喂 64 张",不是全局 64。全局 batch = batch_size × world_size:2 卡就是 128,4 卡就是 256。这意味着单卡调好的 batch_size 在开 DDP 后每卡保持原值,全局吞吐变大,通常学习率也要跟着放大——本章为了聚焦分布式本身不调 lr,但记住这个换算关系,调参时才不会懵。
三、解决思路
- 方案 A(教学)
DataParallel:nn.DataParallel(model)一行包装,PyTorch 自动把每个 batch 拆到多卡并行。改动最小,但主卡要做数据分发和结果汇总,多卡时负载不均、扩展性差,了解即可,生产不推荐。 - 方案 B(标准)
DistributedDataParallel(DDP):用spawn启动world_size个进程,每个进程:- 先
init_process_group建立进程组(Windows/CPU 用 gloo); - 用
DistributedSampler按 rank 拿自己的数据分片,每 epoch 调set_epoch重洗; - 用
DistributedDataParallel包装模型(反向传播时自动 allreduce 同步梯度); - 只在 rank 0 存档、写日志、跑验证,验证结果
broadcast给所有人统一决策。
- 先
- 一个开关控制三种模式:
config里distributed.enabled管 DDP,distributed.data_parallel管 DP,两个都关就是单卡——同一套工程三种跑法。 - 不做什么:不引入
torchrun/ddp命令行启动器(那是生产环境的标准姿势),本章聚焦"理解原理 + 能用python train.py直接跑通"。
trade-off:DDP 的进程模型明显更复杂(spawn 启动、每进程独立配置、同步语义),换来接近线性的加速和稳定的多卡表现;gloo 在 CPU/Windows 上能跑通教学示例,但真实大模型训练必须上 Linux + nccl。验证只在 rank 0 跑会稍微拖慢那个进程,对 CIFAR-10 来说毫无感觉。
四、代码变更
相对第 12 章的改动(第 12 章本身是第 11 章 + utils/seed.py 的工程):
+ utils/dist.py # 新增:进程组初始化/清理、主进程判断
+
# config/config.yaml
seed: 42
device: auto
+ distributed: # 多 GPU / 分布式训练开关
+ enabled: false # true = DDP 多进程;false = 不走 DDP
+ data_parallel: false # enabled=false 且 true 时,用简单版 DataParallel
+ world_size: 2 # DDP 进程数(= 卡数)
+ backend: gloo # Windows/CPU 用 gloo;Linux 多卡可换 nccl
data:
root: ./data
batch_size: 64 # DDP 下这是每卡 batch,全局 = 64 × world_size
...
# dataset/datasets.py
- from torch.utils.data import DataLoader, random_split
+ from torch.utils.data import DataLoader, DistributedSampler, random_split
- def build_dataloaders(cfg, worker_init_fn=None):
+ def build_dataloaders(cfg, worker_init_fn=None, rank=0, world_size=1):
...
- return (make_loader(train_dataset, shuffle=True),
- make_loader(val_dataset, shuffle=False),
- make_loader(test_dataset, shuffle=False))
+ train_sampler = None
+ if world_size > 1:
+ train_sampler = DistributedSampler(
+ train_dataset, num_replicas=world_size, rank=rank, shuffle=True)
+ train_loader = DataLoader(train_dataset, batch_size=cfg.data.batch_size,
+ shuffle=(train_sampler is None),
+ sampler=train_sampler, ...)
+ ...
+ return train_loader, val_loader, test_loader, train_sampler
# engine/trainer.py
- def __init__(self, cfg, model, optimizer, scheduler, scaler,
- train_loader, val_loader, test_loader, logger, writer):
+ def __init__(self, cfg, model, optimizer, scheduler, scaler,
+ train_loader, val_loader, test_loader, train_sampler,
+ logger, writer, rank=0, world_size=1):
...
+ self.is_main = (rank == 0)
+ if self.train_sampler is not None:
+ self.train_sampler.set_epoch(epoch) # 每 epoch 重洗数据分片
+ ...
+ def _dist_validate(self): # 验证只在 rank 0 算,再 broadcast
+ ...
+ val_loss, val_acc = self._dist_validate()
+ if not self.is_main: # 存档/日志/测试只让 rank 0 干
+ return
# train.py
+ from utils.dist import dist_setup, dist_cleanup, is_main_process
+ def main_worker(rank, world_size, args, use_ddp): # 每个进程跑一次
+ ...
+ set_seed(cfg.seed + rank) # 各进程随机序列互相独立
+ if use_ddp:
+ dist_setup(rank, world_size, backend=cfg.distributed.backend)
+ ...
+ model = nn.parallel.DistributedDataParallel(model, device_ids=[rank])
+ ...
+ dist_cleanup()
+ def main():
+ if use_ddp:
+ torch.multiprocessing.spawn(main_worker, ...) # spawn 启动 N 个进程
+ else:
+ main_worker(0, 1, args, False)五、完整代码
本章完整代码 = 第 12 章工程 + 以下新增/修改文件(未列出的文件,如 models/classifier.py、utils/checkpoint.py、utils/metrics.py、utils/early_stopping.py、utils/logger.py、utils/seed.py,与第 12 章完全相同,requirements.txt 也不需要改动)。
utils/dist.py(新增):
"""分布式工具:进程组初始化/清理,主进程判断。"""
import torch
def dist_setup(rank, world_size, backend="gloo"):
"""初始化分布式进程组。gloo 通用(Windows/CPU 均可);Linux 多卡可换 nccl。"""
torch.distributed.init_process_group(
backend=backend, # 通信后端:gloo / nccl
init_method="tcp://127.0.0.1:23456", # 通过本机 TCP 端口"碰头",大家连过来互相认识
rank=rank, # 我的工号(第几个进程)
world_size=world_size, # 一共有几个进程
)
def dist_cleanup():
"""训练结束释放进程组资源。"""
torch.distributed.destroy_process_group()
def is_main_process(rank):
"""只有 rank 0 执行存档/日志/验证等"唯一职责"。"""
return rank == 0config/config.yaml(修改,新增 distributed 段):
# CIFAR-10 图像分类实验配置
seed: 42 # 全局随机种子(DDP 下每个进程自动加 rank 偏移)
device: auto # auto / cuda / cpu
distributed: # —— 本章新增:多 GPU / 分布式训练开关 ——
enabled: false # true = DDP 多进程训练;false = 不走 DDP
data_parallel: false # enabled=false 且 true 时,用简单版 DataParallel
world_size: 2 # DDP 进程数(= 卡数),spawn 启动几个 worker
backend: gloo # gloo 跨平台通用(Windows/CPU 均可);Linux 多卡可换 nccl
data:
root: ./data
batch_size: 64 # DDP 下这是"每卡"的 batch,全局 = batch_size × world_size
num_workers: 2
val_ratio: 0.1
augmentation:
random_crop_padding: 4
horizontal_flip: true
color_jitter: [0.2, 0.2, 0.2, 0.1]
model:
num_classes: 10
train:
epochs: 30
lr: 0.001
lr_scheduler: step
lr_step_size: 15
lr_gamma: 0.1
patience: 7
use_amp: true # AMP 混合精度(仅 GPU 生效)
checkpoint:
dir: ./checkpoints
log:
dir: ./runs/cifar10dataset/datasets.py(修改,支持按 rank 切分数据):
"""数据集与数据加载:build_transforms 构建变换,build_dataloaders 组装三个 loader。"""
from torch.utils.data import DataLoader, DistributedSampler, random_split
from torchvision import datasets, transforms
def build_transforms(cfg, train=False):
"""按配置构建训练/验证两套变换。训练集做增强,评估集只归一化。"""
aug = cfg.data.augmentation
transforms_list = []
if train:
if aug.random_crop_padding > 0:
transforms_list.append(
transforms.RandomCrop(32, padding=aug.random_crop_padding))
if aug.horizontal_flip:
transforms_list.append(transforms.RandomHorizontalFlip())
if aug.color_jitter is not None:
transforms_list.append(transforms.ColorJitter(*aug.color_jitter))
transforms_list += [
transforms.ToTensor(),
transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2470, 0.2435, 0.2616)),
]
return transforms.Compose(transforms_list)
def build_dataloaders(cfg, worker_init_fn=None, rank=0, world_size=1):
"""返回 (train_loader, val_loader, test_loader, train_sampler)。
DDP 时 train_loader 使用 DistributedSampler 按 rank 切分数据。"""
train_dataset = datasets.CIFAR10(root=cfg.data.root, train=True, download=True,
transform=build_transforms(cfg, train=True))
test_dataset = datasets.CIFAR10(root=cfg.data.root, train=False, download=True,
transform=build_transforms(cfg, train=False))
val_size = int(len(train_dataset) * cfg.data.val_ratio)
train_dataset, val_dataset = random_split( # 训练集里再切出验证集
train_dataset, [len(train_dataset) - val_size, val_size]
)
train_sampler = None
if world_size > 1: # 只有 DDP 才需要分布式采样器
train_sampler = DistributedSampler( # 数据集切成 world_size 份,本进程只看第 rank 份
train_dataset, num_replicas=world_size, rank=rank, shuffle=True)
train_loader = DataLoader(train_dataset, batch_size=cfg.data.batch_size,
shuffle=(train_sampler is None), # 有 sampler 时必须关 shuffle,否则冲突
sampler=train_sampler,
num_workers=cfg.data.num_workers,
worker_init_fn=worker_init_fn)
val_loader = DataLoader(val_dataset, batch_size=cfg.data.batch_size,
shuffle=False, num_workers=cfg.data.num_workers,
worker_init_fn=worker_init_fn)
test_loader = DataLoader(test_dataset, batch_size=cfg.data.batch_size,
shuffle=False, num_workers=cfg.data.num_workers,
worker_init_fn=worker_init_fn)
return train_loader, val_loader, test_loader, train_samplerengine/trainer.py(修改,支持 DDP 的验证/存档去重):
"""Trainer:训练/验证/存档/早停/恢复/指标记录的完整封装,支持 DDP。"""
import os
import torch
import torch.nn as nn
from tqdm import tqdm
from utils.checkpoint import save_checkpoint, load_checkpoint
from utils.metrics import MetricsTracker
from utils.early_stopping import EarlyStopping
class Trainer:
def __init__(self, cfg, model, optimizer, scheduler, scaler,
train_loader, val_loader, test_loader, train_sampler,
logger, writer, rank=0, world_size=1):
self.cfg = cfg
self.model = model
self.optimizer = optimizer
self.scheduler = scheduler
self.scaler = scaler
self.train_loader = train_loader
self.val_loader = val_loader
self.test_loader = test_loader
self.train_sampler = train_sampler # 只有 DDP 时才非 None
self.logger = logger
self.writer = writer
self.rank = rank # 我的工号
self.world_size = world_size # 总进程数
self.is_main = (rank == 0) # 存档/日志/验证这些"唯一职责"只让 rank 0 干
self.device = next(model.parameters()).device
self.criterion = nn.CrossEntropyLoss()
self.metrics = MetricsTracker()
self.early_stopping = EarlyStopping(patience=cfg.train.patience)
self.start_epoch = 1
self.best_acc = 0.0
def _unwrap(self):
"""取回可存档的原始模型(DDP 包装后多一层 .module)。"""
return self.model.module if hasattr(self.model, "module") else self.model
def _train_one_epoch(self, epoch):
if self.train_sampler is not None:
self.train_sampler.set_epoch(epoch) # 每个 epoch 重打乱本进程的数据分片
self.model.train()
total_loss, correct, total = 0.0, 0, 0
for images, labels in tqdm(self.train_loader, desc="Train", leave=False):
images, labels = images.to(self.device), labels.to(self.device)
self.optimizer.zero_grad()
with torch.amp.autocast(device_type=self.device.type,
enabled=self.scaler.is_enabled()):
outputs = self.model(images)
loss = self.criterion(outputs, labels)
self.scaler.scale(loss).backward() # DDP 会在这里自动触发梯度 allreduce 同步
self.scaler.step(self.optimizer)
self.scaler.update()
total_loss += loss.item() * images.size(0)
correct += (outputs.argmax(dim=1) == labels).sum().item()
total += images.size(0)
return total_loss / total, correct / total
def _validate(self, loader):
self.model.eval()
total_loss, correct, total = 0.0, 0, 0
with torch.no_grad():
for images, labels in tqdm(loader, desc="Val", leave=False):
images, labels = images.to(self.device), labels.to(self.device)
outputs = self.model(images)
loss = self.criterion(outputs, labels)
total_loss += loss.item() * images.size(0)
correct += (outputs.argmax(dim=1) == labels).sum().item()
total += images.size(0)
return total_loss / total, correct / total
def _dist_validate(self):
"""验证只在 rank 0 计算,然后广播给所有进程。"""
val_loss, val_acc = 0.0, 0.0
if self.is_main:
val_loss, val_acc = self._validate(self.val_loader) # 只有 rank 0 真正跑验证集
if self.world_size > 1:
tensor = torch.tensor([val_loss, val_acc], device=self.device)
torch.distributed.broadcast(tensor, src=0) # 把验证指标广播给所有进程
val_loss, val_acc = tensor.tolist()
return val_loss, val_acc
def _build_state(self, epoch):
return {
"model_state": self._unwrap().state_dict(), # 用原始模型存档,去掉 DDP 的 .module 前缀
"optimizer_state": self.optimizer.state_dict(),
"scheduler_state": self.scheduler.state_dict() if self.scheduler else None,
"scaler_state": self.scaler.state_dict(),
"epoch": epoch,
"best_acc": self.best_acc,
"metrics": self.metrics.to_dict(),
}
def _save(self, epoch, is_best):
if not self.is_main: # 非 rank 0 直接跳过,避免写重复 checkpoint
return
ckpt_dir = self.cfg.checkpoint.dir
state = self._build_state(epoch)
save_checkpoint(state, os.path.join(ckpt_dir, "last.pth"))
if is_best:
save_checkpoint(state, os.path.join(ckpt_dir, "best.pth"))
self.logger.info(f"[Checkpoint] best.pth 已更新(epoch {epoch})")
def _resume(self, resume_path):
ckpt = load_checkpoint(resume_path, self.device)
self._unwrap().load_state_dict(ckpt["model_state"])
self.optimizer.load_state_dict(ckpt["optimizer_state"])
if self.scheduler is not None and ckpt.get("scheduler_state") is not None:
self.scheduler.load_state_dict(ckpt["scheduler_state"])
self.scaler.load_state_dict(ckpt["scaler_state"])
self.start_epoch = ckpt["epoch"] + 1
self.best_acc = ckpt.get("best_acc", 0.0)
self.metrics.load(ckpt.get("metrics", {}))
if self.is_main:
self.logger.info(
f"[Resume] 从 epoch {ckpt['epoch']} 恢复,历史最优 val_acc={self.best_acc:.4f}")
def train(self, resume=None):
if resume:
self._resume(resume)
if self.is_main:
self.logger.info("=" * 60)
self.logger.info(f"最终生效配置(rank {self.rank},world_size {self.world_size}):")
for k, v in self.cfg.items():
self.logger.info(f" {k}: {v}")
self.logger.info("=" * 60)
for epoch in range(self.start_epoch, self.cfg.train.epochs + 1):
train_loss, train_acc = self._train_one_epoch(epoch)
val_loss, val_acc = self._dist_validate() # DDP 下所有进程拿到同一个验证分数
if self.scheduler is not None:
self.scheduler.step()
current_lr = self.optimizer.param_groups[0]["lr"]
if self.is_main: # 指标记录/打日志只让 rank 0 做
self.metrics.update(train_loss=train_loss, train_acc=train_acc,
val_loss=val_loss, val_acc=val_acc)
self.logger.info(
f"Epoch {epoch:02d}/{self.cfg.train.epochs} | "
f"Train loss {train_loss:.4f} acc {train_acc:.4f} | "
f"Val loss {val_loss:.4f} acc {val_acc:.4f} | "
f"lr {current_lr:.2e}")
self.writer.add_scalar("train/loss", train_loss, epoch)
self.writer.add_scalar("train/acc", train_acc, epoch)
self.writer.add_scalar("val/loss", val_loss, epoch)
self.writer.add_scalar("val/acc", val_acc, epoch)
self.writer.add_scalar("lr", current_lr, epoch)
self._save(epoch, is_best=(val_acc > self.best_acc))
self.best_acc = max(self.best_acc, val_acc)
self.early_stopping(val_acc) # 所有进程拿到同一个 val_acc,早停判断必然一致
if self.early_stopping.early_stop:
if self.is_main:
self.logger.info(
f"[EarlyStopping] 连续 {self.cfg.train.patience} 轮无提升,"
f"在 epoch {epoch} 停止")
break
if self.is_main:
test_loss, test_acc = self._validate(self.test_loader) # 测试集也只跑一遍
self.logger.info(f"[Test] loss {test_loss:.4f} acc {test_acc:.4f}")train.py(修改,支持三种模式切换):
"""训练入口:支持单卡、DataParallel、DistributedDataParallel 三种模式。"""
import argparse
import logging
import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.tensorboard import SummaryWriter
import yaml
from dataset.datasets import build_dataloaders
from models.classifier import build_model
from engine.trainer import Trainer
from utils.logger import setup_logger
from utils.seed import set_seed, worker_init_fn
from utils.dist import dist_setup, dist_cleanup, is_main_process # 本章新增的分布式工具
def parse_args():
parser = argparse.ArgumentParser(description="CIFAR-10 图像分类训练")
parser.add_argument("--config", type=str, default="config/config.yaml",
help="YAML 配置文件路径")
parser.add_argument("--lr", type=float, default=None, help="覆盖 train.lr")
parser.add_argument("--epochs", type=int, default=None, help="覆盖 train.epochs")
parser.add_argument("--batch-size", type=int, default=None, help="覆盖 data.batch_size")
parser.add_argument("--resume", type=str, default=None, help="从指定 checkpoint 恢复训练")
parser.add_argument("--seed", type=int, default=None, help="覆盖 seed")
return parser.parse_args()
class Config(dict):
def __getattr__(self, key):
try:
return self[key]
except KeyError as e:
raise AttributeError(key) from e
def _to_config(obj):
if isinstance(obj, dict):
return Config({k: _to_config(v) for k, v in obj.items()})
if isinstance(obj, list):
return [_to_config(v) for v in obj]
return obj
def load_config(path):
with open(path, "r", encoding="utf-8") as f:
return _to_config(yaml.safe_load(f))
def build_scheduler(optimizer, cfg):
if cfg.train.lr_scheduler == "step":
return optim.lr_scheduler.StepLR(
optimizer, step_size=cfg.train.lr_step_size, gamma=cfg.train.lr_gamma)
if cfg.train.lr_scheduler == "cosine":
return optim.lr_scheduler.CosineAnnealingLR(
optimizer, T_max=cfg.train.epochs)
return None
def main_worker(rank, world_size, args, use_ddp):
"""每个进程都执行一次本函数(单卡时只跑一次,rank=0)。"""
cfg = load_config(args.config)
if args.lr is not None:
cfg.train.lr = args.lr
if args.epochs is not None:
cfg.train.epochs = args.epochs
if args.batch_size is not None:
cfg.data.batch_size = args.batch_size
if args.seed is not None:
cfg.seed = args.seed
set_seed(cfg.seed + rank) # 各进程随机序列互相独立,避免所有卡看到相同顺序
if use_ddp:
dist_setup(rank, world_size, backend=cfg.distributed.backend) # 开会:建立进程组
if is_main_process(rank):
logger = setup_logger(cfg.log.dir) # 只有 rank 0 建日志和 TensorBoard
writer = SummaryWriter(log_dir=cfg.log.dir)
logger.info(f"Using device: {cfg.device},world_size={world_size}")
else:
logger = logging.getLogger("train") # 其他进程只留一个静默 logger,不刷屏
logger.setLevel(logging.ERROR)
writer = None
train_loader, val_loader, test_loader, train_sampler = build_dataloaders(
cfg, worker_init_fn=worker_init_fn, rank=rank, world_size=world_size)
model = build_model(cfg)
if use_ddp:
if torch.cuda.is_available():
model = model.to(rank) # 每个进程把模型放到自己的卡上(按 rank 编卡号)
model = nn.parallel.DistributedDataParallel(model, device_ids=[rank])
else:
model = nn.parallel.DistributedDataParallel(model) # CPU 也能跑 DDP(gloo)
else:
if cfg.distributed.data_parallel and torch.cuda.device_count() > 1:
model = nn.DataParallel(model) # 单进程包一层,自动把 batch 拆到多卡
model = model.to(rank) if torch.cuda.is_available() else model
optimizer = optim.Adam(model.parameters(), lr=cfg.train.lr)
scheduler = build_scheduler(optimizer, cfg)
use_amp = cfg.train.use_amp and torch.cuda.is_available()
scaler = torch.amp.GradScaler("cuda", enabled=use_amp)
if use_amp and is_main_process(rank):
logger.info("AMP 已启用(FP16 混合精度)")
trainer = Trainer(cfg, model, optimizer, scheduler, scaler,
train_loader, val_loader, test_loader, train_sampler,
logger, writer, rank=rank, world_size=world_size)
trainer.train(resume=args.resume)
if writer is not None:
writer.close()
if use_ddp:
dist_cleanup() # 散会:释放进程组资源
def main():
args = parse_args()
cfg = load_config(args.config)
use_ddp = cfg.distributed.enabled
world_size = cfg.distributed.world_size if use_ddp else 1
if use_ddp:
torch.multiprocessing.spawn( # spawn 启动 world_size 个进程(Windows 上也安全)
main_worker, args=(world_size, args, True), nprocs=world_size)
else:
main_worker(0, 1, args, False) # 单卡 / DataParallel 都只跑一个进程
if __name__ == "__main__":
main()运行方式:单卡 / DataParallel / DDP
三种模式只改 config/config.yaml 里的两个开关,命令行都是 python train.py:
① 单卡(默认):distributed.enabled: false 且 distributed.data_parallel: false。行为与第 12 章完全一致,一个进程、一张卡。
python train.py --epochs 5② DataParallel(教学用,需 ≥2 张 GPU):把 data_parallel: true(enabled 保持 false):
distributed:
enabled: false
data_parallel: truepython train.py --epochs 5仍只启动一个进程,模型被 nn.DataParallel 包一层,batch 自动拆到多卡。可以看到日志打印顺序与单卡相同,但每步吞吐变大。
③ DDP(标准做法):把 enabled: true,并把 world_size 设成可用的 GPU 数:
distributed:
enabled: true
data_parallel: false
world_size: 2
backend: gloopython train.py --epochs 5spawn 会启动 2 个进程(rank 0 和 rank 1)。注意此时 batch_size: 64 是每卡的 batch,全局 batch = 128;终端里只会看到 rank 0 打印的一份日志,checkpoints/ 里也只会有一份 best.pth。
注意事项:
- 只有 1 张 GPU(或只有 CPU)时跑
world_size: 2的 DDP 也能"跑通"——gloo 在 CPU 上同样支持多进程,只是多个进程共享同一块 GPU 显存,速度不升反降。想快速验证 DDP 逻辑、又不想买多卡,可以故意用 CPU 跑 2 进程当"分布式冒烟测试"。 - 启动 DDP 后每个进程是独立进程,没有 GPU 的机器上 CPU 训练也能演示全部同步逻辑;生产环境请用 Linux + nccl。
spawn启动的代码必须受if __name__ == "__main__":保护,否则子进程会再次执行main()造成无限递归。
六、本章小结
- 学到了什么
- 多卡加速有两条路:
DataParallel(单进程多卡,简单但有主卡瓶颈)和DistributedDataParallel(多进程每卡一份模型副本,标准做法)。 - DDP 的四个核心动作:
init_process_group建组、DistributedSampler切数据、DDP包装模型(自动 allreduce 梯度)、rank 0 承担全部"唯一职责"(存档/日志/验证)。 - 分布式下"同步决策"靠
broadcast:验证只在 rank 0 算,分数广播给所有人,早停/存档判断才一致。 - DDP 下
batch_size是每卡的值,全局 batch = batch_size × world_size;每 epoch 记得set_epoch让分片内顺序不同。 - Windows/CPU 上必须用
gloobackend +torch.multiprocessing.spawn;Linux 多卡生产环境换nccl。
- 多卡加速有两条路:
- 常见坑
- 忘记
set_epoch(epoch):每个 epoch 数据顺序完全相同,数据增强基本失效,模型学不到新分布。 shuffle=True和sampler同时存在:PyTorch 直接报错,DDP 下必须shuffle=(train_sampler is None)。- 所有进程都写 checkpoint/日志:文件互相踩踏甚至死锁,必须用
is_main守卫。 - 存档用了
self.model.state_dict():DDP 包装后带module.前缀,必须_unwrap()取原始模型。 - 在 Windows 上把
backend写成nccl:直接报错,只有 gloo 可用。 spawn代码没放在if __name__ == "__main__"里:无限递归启动子进程,进程数爆炸。
- 忘记
- 下一章预告:多卡训完的模型还"活在训练脚本里",无法对任意图片做单图预测,也没法系统评估各类别表现——第 14 章将补齐工程的"输出侧":
predict.py单图推理与analyze.py混淆矩阵/每类指标分析。
七、动手练习
- 三种模式对比:在一台多卡机器上,分别以单卡、
data_parallel: true、enabled: true(world_size=2)跑--epochs 5,对比单 epoch 耗时与最终 val_acc;再用nvidia-smi观察训练中每张卡的显存占用和利用率,体会 DDP 的负载均衡优势。 - 观察数据分片:临时在
main_worker里打印len(train_loader.dataset)与第一个 batch 的标签序列,分别用单卡和 DDP(world_size=2)跑,确认"每个进程只看 1/N 的数据、且分片不重叠"。 - 注释掉 set_epoch:把
_train_one_epoch里的set_epoch那两行注释掉,用 DDP 跑 3 个 epoch,打印每个 epoch 第一个 batch 的标签序列——会发现 3 个 epoch 的顺序完全相同,说明数据增强白做了。 - 撤销 is_main 守卫:在
_save里临时去掉if not self.is_main: return,DDP 跑 1 个 epoch,观察checkpoints/下出现什么、日志是否有冲突,体会"唯一职责"的意义。 - CPU 冒烟测试:没有 GPU 的机器上把
enabled: true, world_size: 2跑通(gloo 支持 CPU 多进程),确认"能跑"不等于"更快"——用time对比单卡 CPU 与双进程 CPU 的耗时,量化通信开销。
