模型训练代码从单卡迁移到多卡需要改的地方
模型训练代码从单卡迁移到多卡需要改的地方
flyfish
需要改以下几个部分,单卡情况代码怎么写,多卡情况代码怎么写,有个对比
导入了 DDP、dist、DistributedSampler
初始化了DDP环境,拿到local_rank/rank/world_size
DEVICE绑定了local_rank
DataLoader换成了DistributedSampler,返回了train_sampler
每个epoch调用了 train_sampler.set_epoch(epoch)
模型用DDP包装,开了find_unused_parameters=True
所有 model.features / model.classifier 都加了 .module.
验证函数做了全局指标聚合
早停信号广播到了所有卡
模型保存用了 model.module...,去掉了前缀
所有打印、保存文件、可视化都只在rank0执行
最后调用了 dist.destroy_process_group()
用torchrun启动脚本
导入
位置:所有 import 的地方
新增内容:
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import DistributedSampler # 原来只有DataLoader,新增采样器
全局配置前:初始化DDP环境
位置:torch.backends.cudnn 配置之前
新增代码:
# 初始化分布式环境
def setup_ddp():
local_rank = int(os.environ["LOCAL_RANK"])
rank = int(os.environ["RANK"])
world_size = int(os.environ["WORLD_SIZE"])
torch.cuda.set_device(local_rank)
dist.init_process_group(
backend="nccl",
device_id=torch.device(f"cuda:{local_rank}")
)
return local_rank, rank, world_size
# 封装只在主进程打印的工具
def print_log(*args, **kwargs):
if rank == 0:
print(*args, **kwargs)
# 执行初始化
local_rank, rank, world_size = setup_ddp()
作用:建立多卡通信、给每个进程分配对应显卡、统一日志输出。
全局设备变量
位置:原来的 DEVICE = ...
原写法:
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
改后写法:
DEVICE = torch.device(f"cuda:{local_rank}")
作用:每个进程独占一张显卡,避免多进程抢0号卡。
数据加载函数 build_dataloader
4.1 替换采样器
位置:DataLoader 创建处
原写法:
train_loader = DataLoader(
train_dataset, BATCH_SIZE, shuffle=True,
...
)
val_loader = DataLoader(
val_dataset, BATCH_SIZE, shuffle=False,
...
)
改后写法:
# 新增:分布式采样器
train_sampler = DistributedSampler(train_dataset, shuffle=True)
val_sampler = DistributedSampler(val_dataset, shuffle=False)
train_loader = DataLoader(
train_dataset, BATCH_SIZE,
sampler=train_sampler, # 替换shuffle=True
num_workers=NUM_WORKERS, pin_memory=True,
prefetch_factor=4, persistent_workers=True, drop_last=True
)
val_loader = DataLoader(
val_dataset, BATCH_SIZE,
sampler=val_sampler, # 替换shuffle=False
num_workers=NUM_WORKERS, pin_memory=True,
prefetch_factor=4, persistent_workers=True
)
4.2 函数返回值增加sampler
原写法:
return train_loader, val_loader
改后写法:
return train_loader, val_loader, train_sampler
作用:训练时需要用sampler设置epoch打乱数据。
训练阶段函数:冻结/解冻层
位置:所有 model.features/model.classifier` 的地方
5.1 冻结解冻代码
原写法:
for param in model.features.parameters():
param.requires_grad = False
for param in model.classifier.parameters():
param.requires_grad = True
改后写法:
for param in model.module.features.parameters():
param.requires_grad = False
for param in model.module.classifier.parameters():
param.requires_grad = True
5.2 优化器传参
原写法:
optimizer = optim.AdamW(model.classifier.parameters(), ...)
改后写法:
optimizer = optim.AdamW(model.module.classifier.parameters(), ...)
5.3 每个epoch开头加打乱设置
原写法:每个epoch循环里第一行没有sampler操作
改后写法:每个epoch循环第一行加:
train_sampler.set_epoch(epoch)
作用:保证每个epoch数据打乱方式不同,避免每轮数据顺序完全一样。
验证函数 validate
单卡指标 变 全局聚合指标
- 把统计量从普通数值改成GPU张量,方便跨卡聚合
- 用
all_reduce汇总loss、正确数、总数 - 用
all_gather收集所有预测结果,主进程算混淆矩阵 - 只有主进程打印结果
改后逻辑:
@torch.no_grad()
def validate(model, val_loader, criterion):
model.eval()
# 改成张量,方便跨卡聚合
total_loss = torch.tensor(0.0, device=DEVICE)
correct = torch.tensor(0, device=DEVICE)
total = torch.tensor(0, device=DEVICE)
all_preds = []
all_labels = []
for images, labels, _ in tqdm(val_loader, desc="正在验证", disable=(rank!=0)):
images = images.to(DEVICE, non_blocking=True)
labels = labels.to(DEVICE, non_blocking=True)
outputs = model(images)
loss = criterion(outputs, labels)
total_loss += loss
_, preds = torch.max(outputs, 1)
correct += (preds == labels).sum()
total += labels.size(0)
all_preds.append(preds)
all_labels.append(labels)
all_preds = torch.cat(all_preds)
all_labels = torch.cat(all_labels)
# 全局聚合
dist.all_reduce(total_loss, op=dist.ReduceOp.SUM)
dist.all_reduce(correct, op=dist.ReduceOp.SUM)
dist.all_reduce(total, op=dist.ReduceOp.SUM)
avg_loss = total_loss.item() / world_size / len(val_loader)
acc = correct.item() / total.item()
# 收集所有卡的预测到主进程
gathered_preds = [torch.zeros_like(all_preds) for _ in range(world_size)]
gathered_labels = [torch.zeros_like(all_labels) for _ in range(world_size)]
dist.all_gather(gathered_preds, all_preds)
dist.all_gather(gathered_labels, all_labels)
if rank == 0:
final_preds = torch.cat(gathered_preds).cpu().numpy()
final_labels = torch.cat(gathered_labels).cpu().numpy()
print(f"\n验证集:Loss={avg_loss:.4f}, 准确率={acc:.4f}")
print("混淆矩阵:\n", confusion_matrix(final_labels, final_preds))
print(classification_report(final_labels, final_preds, target_names=CLASS_NAMES, digits=4))
dist.barrier()
return avg_loss, acc
模型保存与加载
7.1 保存权重
位置:训练阶段里保存最优模型的地方
原写法:
torch.save(model.state_dict(), ...)
改后写法(保存原生无前缀权重,兼容单卡部署):
# 如果开了torch.compile,用这个
torch.save(model.module._orig_mod.state_dict(), os.path.join(SAVE_DIR, "stage1_best.pth"))
# 如果没开compile,用这个
# torch.save(model.module.state_dict(), os.path.join(SAVE_DIR, "stage1_best.pth"))
7.2 加载最优权重
原写法:
model.load_state_dict(torch.load(...))
改后写法:
if rank == 0:
state_dict = torch.load(os.path.join(SAVE_DIR, "stage1_best.pth"), map_location=DEVICE)
model.module._orig_mod.load_state_dict(state_dict)
dist.barrier()
# 广播参数到所有卡,保证模型一致
for param in model.module.parameters():
dist.broadcast(param.data, src=0)
早停逻辑同步
位置:早停判断
原写法:只有rank0判断早停、break
问题:只有主卡退出,其他卡还在跑,程序会死锁
改后写法:
stop_flag = torch.tensor(0, device=DEVICE)
if rank == 0:
if val_acc > best_acc:
best_acc = val_acc
early_stop_counter = 0
torch.save(model.module._orig_mod.state_dict(), os.path.join(SAVE_DIR, "stage2_best.pth"))
else:
early_stop_counter += 1
if early_stop_counter >= EARLY_STOP_PATIENCE:
stop_flag = torch.tensor(1, device=DEVICE)
# 广播早停信号,所有卡一起判断
dist.broadcast(stop_flag, src=0)
dist.barrier()
if stop_flag.item() == 1:
break
所有IO/可视化/推理:只让主进程执行
以下操作全部加 if rank == 0 判断,避免多卡重复执行、文件冲突:
- 创建保存目录
os.makedirs(SAVE_DIR) - 打印数据集统计、训练日志
- 生成Grad-CAM热力图
grad_cam_check - TTA推理
tta_infer - 保存CSV、写文件、计算耗时
- 高损失样本筛查
check_high_loss_samples
可以直接把所有 print 替换成前面封装的 print_log。
仅主进程打印日志
def print_log(*args, **kwargs):
if rank == 0:
print(*args, **kwargs)
主函数:模型DDP包装
位置:model = build_model() 之后
原写法:
model = build_model()
model = torch.compile(model, mode="default")
改后写法:
model = build_model()
# 可选:强制参数内存连续,消除梯度布局警告
for param in model.parameters():
param.data = param.data.contiguous()
model = torch.compile(model, mode="default")
# DDP包装
model = DDP(
model,
device_ids=[local_rank],
output_device=local_rank,
find_unused_parameters=True # 有冻结层必须开,否则报错
)
程序末尾:清理分布式环境
位置:主函数最后一行
新增代码:
dist.destroy_process_group()
作用:正常释放通信资源,避免残留进程和警告。
启动命令
原启动方式:
python train.py
多卡启动方式:
# 2张卡
torchrun --nproc_per_node=2 train.py
# 4张卡
torchrun --nproc_per_node=4 train.py
更多推荐

所有评论(0)