You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch中如何用多进程训练任意尺寸图像的微调模型?

解决方案

针对你遇到的任意尺寸图像模型微调时的性能问题,以下是几个可行的解决思路:

1. 基于PyTorch多进程的手动并行处理

multiprocessing.pool无法直接兼容autograd,但可以用PyTorch原生的torch.multiprocessing,通过共享内存让子进程复用模型参数,手动收集输出和梯度:

import torch
import torch.multiprocessing as mp

def process_batch(rank, model, imgs_chunk, output_queue, target_chunk):
    # 子进程复用共享的模型参数,计算当前分片的预测与损失
    preds = [model(im) for im in imgs_chunk]
    chunk_loss = loss_fn(torch.cat(preds), target_chunk)
    chunk_loss.backward()  # 子进程中计算梯度,会同步到共享的模型参数
    output_queue.put(preds)

def parallel_process(model, imgs, targets, num_workers=4):
    model.share_memory()  # 将模型参数放入共享内存,供子进程访问
    # 拆分数据为多个分片
    img_chunks = [imgs[i::num_workers] for i in range(num_workers)]
    target_chunks = [targets[i::num_workers] for i in range(num_workers)]
    
    output_queue = mp.Queue()
    processes = []
    for rank in range(num_workers):
        p = mp.Process(
            target=process_batch,
            args=(rank, model, img_chunks[rank], output_queue, target_chunks[rank])
        )
        p.start()
        processes.append(p)
    
    # 收集所有预测结果
    all_preds = []
    for _ in range(num_workers):
        all_preds.extend(output_queue.get())
    # 等待所有子进程结束
    for p in processes:
        p.join()
    
    return torch.cat(all_preds)

# 使用示例
optimizer.zero_grad()
y_pred = parallel_process(model, imgs, y)
# 主进程执行优化器更新
optimizer.step()

2. 动态分组批量处理

自定义DataLoader的collate_fn,将同尺寸的图像归为一组批量输入,不同尺寸的单独处理,尽可能利用GPU的批量计算优势:

from torch.utils.data import DataLoader

def custom_collate(batch):
    # 按图像尺寸分组,batch为(图像, 标签)的列表
    size_groups = {}
    for im, label in batch:
        size_key = im.shape
        if size_key not in size_groups:
            size_groups[size_key] = {'imgs': [], 'labels': []}
        size_groups[size_key]['imgs'].append(im)
        size_groups[size_key]['labels'].append(label)
    
    # 每组转成tensor批量
    batch_imgs = []
    batch_labels = []
    for group in size_groups.values():
        batch_imgs.append(torch.stack(group['imgs']))
        batch_labels.append(torch.tensor(group['labels']))
    
    return batch_imgs, batch_labels

# 构建DataLoader
dataloader = DataLoader(dataset, batch_size=32, collate_fn=custom_collate)

# 训练循环
for imgs_groups, labels_groups in dataloader:
    optimizer.zero_grad()
    y_pred = []
    total_loss = 0.0
    # 对每组同尺寸图像批量计算
    for imgs, labels in zip(imgs_groups, labels_groups):
        pred = model(imgs)
        y_pred.append(pred)
        total_loss += loss_fn(pred, labels)
    
    total_loss.backward()
    optimizer.step()
    y_pred = torch.cat(y_pred)

3. 单进程下的性能优化

如果多进程实现复杂,可在单进程内做针对性优化:

  • 用torch.jit.trace或torch.jit.script将模型编译为TorchScript,提升单张图像的推理速度
  • 采用梯度累积:每处理N张图像后再执行一次反向传播,减少优化器更新的开销
# 梯度累积示例
accumulation_steps = 4
optimizer.zero_grad()

for idx, (im, label) in enumerate(zip(imgs, y)):
    pred = model(im)
    loss = loss_fn(pred, label)
    loss = loss / accumulation_steps  # 均分损失,避免梯度爆炸
    loss.backward()
    
    # 累积到指定步数后更新参数
    if (idx + 1) % accumulation_steps == 0:
        optimizer.step()
        optimizer.zero_grad()

# 处理剩余未累积到步数的样本
if len(imgs) % accumulation_steps != 0:
    optimizer.step()
    optimizer.zero_grad()

内容的提问来源于stack exchange,提问作者sdml

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.17 05:21:12