基于PyTorch的SRGAN植物叶片超分辨率模型CPU训练速度极慢:优化方案咨询
我完全理解用CPU训练SRGAN的痛苦——对抗网络加上感知损失的计算量对CPU来说确实是巨大的负担。结合你的代码和需求,我从提速训练、减少轮次、简化模型三个方向整理了可落地的优化方案:
问题背景
我正在用PyTorch参与植物叶片超分辨率挑战赛,构建了带Generator和Discriminator的SRGAN模型,目标是将低分辨率植物叶片图像生成高分辨率版本。但只能用CPU训练,导致训练耗时极长——生成器、判别器的计算加上感知损失的求解,每个epoch都慢得离谱。已经尝试了混合精度优化,但CPU性能还是没达到预期,希望能从提升训练速度、减少训练轮次、简化模型结构这几个方向得到具体的优化方案。
我的训练代码如下:
best_loss = float('inf') for epoch in range(1, NUM_EPOCHS + 1): G.train(); D.train() g_losses, d_losses = [], [] for lr_img, hr_img in train_dl: lr_img = lr_img.to(DEVICE) hr_img = hr_img.to(DEVICE) B = lr_img.size(0) real_label = torch.ones (B, 1, 1, 1, device=DEVICE) * 0.9 # label smoothing fake_label = torch.zeros(B, 1, 1, 1, device=DEVICE) + 0.1 # -- Discriminator -- opt_D.zero_grad() with torch.cuda.amp.autocast(enabled=torch.cuda.is_available()): fake_hr = G(lr_img).detach() d_real = D(hr_img) d_fake = D(fake_hr) # Match spatial size of labels to discriminator output rl = real_label.expand_as(d_real) fl = fake_label.expand_as(d_fake) loss_D = (criterion_adv(d_real, rl) + criterion_adv(d_fake, fl)) * 0.5 scaler.scale(loss_D).backward() scaler.step(opt_D) # -- Generator -- opt_G.zero_grad() with torch.cuda.amp.autocast(enabled=torch.cuda.is_available()): fake_hr = G(lr_img) d_fake = D(fake_hr) rl = real_label.expand_as(d_fake) loss_pix = criterion_pix(fake_hr, hr_img) loss_adv = criterion_adv(d_fake, rl) if perc_net is not None: with torch.no_grad(): feat_real = perc_net(hr_img) feat_fake = perc_net(fake_hr) loss_perc = F.l1_loss(feat_fake, feat_real.detach()) else: loss_perc = torch.tensor(0.0, device=DEVICE) loss_G = (LAMBDA_PIX * loss_pix + LAMBDA_ADV * loss_adv + LAMBDA_PERC * loss_perc) scaler.scale(loss_G).backward() scaler.step(opt_G) scaler.update() g_losses.append(loss_G.item()) d_losses.append(loss_D.item()) sched_G.step(); sched_D.step() mean_g = np.mean(g_losses) mean_d = np.mean(d_losses) if mean_g < best_loss: best_loss = mean_g torch.save(G.state_dict(), 'best_generator.pth') if epoch % 10 == 0 or epoch == 1: print(f'Epoch [{epoch:>3}/{NUM_EPOCHS}] G: {mean_g:.4f} D: {mean_d:.4f}') print('Training complete. Best G loss:', best_loss)
一、代码层面:直接提升CPU训练速度
1. 替换无效的混合精度为CPU专用优化
你当前用的torch.cuda.amp只对GPU生效,CPU上完全没用。改用PyTorch 1.10+支持的CPU自动混合精度,能降低计算量:
from torch.cpu.amp import autocast, GradScaler scaler = GradScaler() # 把代码里的torch.cuda.amp.autocast改成autocast,enabled=True即可 with autocast(enabled=True): # 判别器/生成器的计算逻辑
同时确保你的PyTorch是MKL/MKL-DNN编译版本(官方预编译包默认包含),这会让CPU张量运算速度提升30%以上。
2. 砍掉冗余计算(最立竿见影的优化)
你的代码里生成器在判别器训练时生成了一次fake_hr,训练生成器时又重新跑了一遍G(lr_img)——这完全是浪费!直接复用之前的结果:
# -- Discriminator -- opt_D.zero_grad() with autocast(enabled=True): fake_hr = G(lr_img).detach() # 这里生成一次 d_real = D(hr_img) d_fake = D(fake_hr) # ... 判别器损失计算 # -- Generator -- opt_G.zero_grad() with autocast(enabled=True): # 直接复用fake_hr,不用再跑生成器! d_fake = D(fake_hr) # ... 生成器损失计算
这能节省整整一半的生成器前向计算时间,对CPU来说提升非常明显。
3. 优化数据加载 pipeline
数据加载是CPU训练的常见瓶颈,调整这几点:
- 把
DataLoader的num_workers设为CPU核心数的1-2倍(比如4核CPU设为num_workers=4),同时关闭pin_memory(CPU不需要内存锁定) - 如果数据集不大,提前把所有图像加载到内存里,避免每轮都读磁盘:
class LeafDataset(Dataset): def __init__(self, lr_paths, hr_paths): # 初始化时就把所有图像加载到内存 self.lr_imgs = [self.load_and_transform(p) for p in lr_paths] self.hr_imgs = [self.load_and_transform(p) for p in hr_paths] - 用PyTorch内置的
torchvision.transforms替代PIL操作,前者的CPU运算效率更高。
二、策略层面:减少不必要的训练轮次
1. 用预训练模型初始化生成器
不用从零训练SRGAN,先找一个经典单图像超分模型(比如EDSR、RCAN)的预训练权重,调整通道数和输出尺寸后初始化你的Generator,再用SRGAN的对抗损失微调。预训练模型已经学习了基本的超分特征,能把收敛所需epoch数减少50%以上。
2. 分阶段训练,优先收敛像素损失
SRGAN的对抗损失和感知损失会增加训练复杂度,你可以分两步走:
- 先只训练像素损失(把
LAMBDA_ADV和LAMBDA_PERC设为0),让生成器快速收敛到基线水平 - 再加入对抗损失和感知损失微调,只需要少量epoch就能达到SRGAN的效果
3. 加入早停机制
不要固定训练NUM_EPOCHS轮,监控验证集的PSNR或SSIM指标,当连续10轮指标没有提升时就停止训练,避免无效训练:
best_psnr = 0.0 patience = 10 no_improve_epoch = 0 for epoch in range(1, NUM_EPOCHS + 1): # ... 训练逻辑 ... # 验证阶段 G.eval() total_psnr = 0.0 with torch.no_grad(): for lr_val, hr_val in val_dl: lr_val = lr_val.to(DEVICE) hr_val = hr_val.to(DEVICE) fake_val = G(lr_val) psnr = torchmetrics.functional.peak_signal_noise_ratio(fake_val, hr_val) total_psnr += psnr.item() mean_psnr = total_psnr / len(val_dl) if mean_psnr > best_psnr: best_psnr = mean_psnr no_improve_epoch = 0 torch.save(G.state_dict(), 'best_generator.pth') else: no_improve_epoch += 1 if no_improve_epoch >= patience: print(f'Early stopping at epoch {epoch}, best PSNR: {best_psnr:.2f}') break
三、模型层面:简化结构适配CPU算力
1. 缩小生成器的通道数和残差块数量
原SRGAN的生成器用了64通道+16个残差块,对CPU来说太heavy。你可以把通道数降到32,残差块减到8个:
class ResidualBlock(nn.Module): def __init__(self, in_channels=32): # 原先是64 super().__init__() self.conv1 = nn.Conv2d(in_channels, in_channels, kernel_size=3, padding=1) self.bn1 = nn.BatchNorm2d(in_channels) self.conv2 = nn.Conv2d(in_channels, in_channels, kernel_size=3, padding=1) self.bn2 = nn.BatchNorm2d(in_channels) def forward(self, x): residual = x out = F.relu(self.bn1(self.conv1(x))) out = self.bn2(self.conv2(out)) out += residual return F.relu(out) class Generator(nn.Module): def __init__(self, scale_factor=4, num_res_blocks=8): # 原先是16 super().__init__() self.conv1 = nn.Conv2d(3, 32, kernel_size=9, padding=4) self.res_blocks = nn.Sequential(*[ResidualBlock(32) for _ in range(num_res_blocks)]) self.conv2 = nn.Conv2d(32, 32, kernel_size=3, padding=1) self.bn2 = nn.BatchNorm2d(32) self.upsample = nn.Sequential( nn.Conv2d(32, 32*scale_factor**2, kernel_size=3, padding=1), nn.PixelShuffle(scale_factor), nn.Conv2d(32, 3, kernel_size=9, padding=4) )
模型参数会减少70%以上,CPU计算量大幅降低,而植物叶片纹理相对规律,简化后性能下降不会太明显。
2. 简化判别器结构
原SRGAN的判别器可以砍掉一半的卷积块,缩小通道数:
class Discriminator(nn.Module): def __init__(self): super().__init__() self.conv_blocks = nn.Sequential( nn.Conv2d(3, 16, kernel_size=3, padding=1), # 原先是64 nn.LeakyReLU(0.2), nn.Conv2d(16, 16, kernel_size=3, stride=2, padding=1), nn.BatchNorm2d(16), nn.LeakyReLU(0.2), nn.Conv2d(16, 32, kernel_size=3, padding=1), nn.BatchNorm2d(32), nn.LeakyReLU(0.2), nn.Conv2d(32, 32, kernel_size=3, stride=2, padding=1), nn.BatchNorm2d(32), nn.LeakyReLU(0.2), nn.AdaptiveAvgPool2d(1), nn.Conv2d(32, 64, kernel_size=1), nn.LeakyReLU(0.2), nn.Conv2d(64, 1, kernel_size=1) ) def forward(self, x): out = self.conv_blocks(x) return torch.sigmoid(out)
判别器只要能提供对抗监督即可,简化后依然能完成任务,计算量却大幅减少。
3. 替换轻量模型做感知损失
原SRGAN用VGG19做感知损失,CPU上计算极慢。换成MobileNetV2或EfficientNet-B0的中间层特征:
from torchvision.models import mobilenet_v2 # 取MobileNetV2的前10层作为感知特征提取器 perc_net = mobilenet_v2(pretrained=True).features[:10].to(DEVICE) perc_net.eval()
轻量模型的特征提取速度比VGG19快3-5倍,感知损失的效果差异极小。
其他小技巧
- 调整batch size:CPU内存有限,测试
batch_size=4或8,找到速度和梯度稳定性的平衡点 - 关闭不必要的操作:训练时除了早停不要计算验证指标,只保存最优生成器权重,避免磁盘IO浪费时间
- 合理利用CPU核心:不要让DataLoader的
num_workers超过CPU核心数,避免进程切换开销
内容的提问来源于stack exchange,提问作者AQSA ZAM ZAM MIRZA JOHAR BAIG

