模型训练代码从单卡迁移到多卡需要改的地方

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

单卡指标 变 全局聚合指标

  1. 把统计量从普通数值改成GPU张量,方便跨卡聚合
  2. all_reduce 汇总loss、正确数、总数
  3. all_gather 收集所有预测结果,主进程算混淆矩阵
  4. 只有主进程打印结果

改后逻辑

@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 判断,避免多卡重复执行、文件冲突:

  1. 创建保存目录 os.makedirs(SAVE_DIR)
  2. 打印数据集统计、训练日志
  3. 生成Grad-CAM热力图 grad_cam_check
  4. TTA推理 tta_infer
  5. 保存CSV、写文件、计算耗时
  6. 高损失样本筛查 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
Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐