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
相关产品推荐
相关产品推荐

