基于PyTorch的SRGAN植物叶片超分辨率模型CPU训练速度极慢——优化方案问询
针对CPU训练SRGAN的优化方案与代码调整建议
我明白在CPU上训练SRGAN这类模型有多头疼——尤其是处理植物叶片超分任务,计算量本来就不小,还得扛着GAN的双重训练负载。结合你的代码和需求,我整理了几个针对性的优化方向,从速度提升、epoch缩减到模型简化都有具体方案,你可以按需尝试:
一、直接提升CPU训练速度的优化
1. 替换无效的混合精度逻辑
你当前用的torch.cuda.amp是GPU专属工具,CPU下完全不生效,这也是混合精度没起作用的核心原因。换成PyTorch对CPU支持的自动混合精度方案,同时移除无用的梯度缩放代码:
# 开启CPU支持的bfloat16混合精度(PyTorch 1.10+支持) use_cpu_amp = True for lr_img, hr_img in train_dl: # 直接把数据转成bfloat16(CPU支持的话) lr_img = lr_img.to(DEVICE, dtype=torch.bfloat16 if use_cpu_amp else torch.float32) hr_img = hr_img.to(DEVICE, dtype=torch.bfloat16 if use_cpu_amp else torch.float32) # 替换原cuda autocast为CPU版本 with torch.autocast(device_type='cpu', dtype=torch.bfloat16, enabled=use_cpu_amp): # 判别器/生成器的计算逻辑 # ... # CPU下不需要scaler,直接反向传播+更新优化器 loss_D.backward() opt_D.step() # ... 生成器侧同理,删掉scaler相关代码
2. 优化数据加载环节(CPU训练的核心瓶颈)
CPU训练时,数据加载的耗时往往占比超过模型计算,建议:
- 调整
DataLoader参数:设置num_workers为CPU核心数的一半(比如4核CPU设为2),同时关闭pin_memory(CPU下不需要) - 提前离线预处理:把图像裁剪、归一化等操作提前完成,保存成
.npy文件,训练时直接加载numpy数组,减少实时计算 - 适当增大
batch_size:在CPU内存允许的前提下,把batch从4提升到8或16,提升CPU计算的利用率
3. 砍掉不必要的计算开销
- 感知损失简化:把
perc_net换成更轻量的模型(比如VGG11代替VGG19),并且只提取前3层特征,减少特征提取的计算量 - 标签平滑简化:如果你的判别器输出是全局平均池化后的单值,不需要用
expand_as扩展标签的空间尺寸,直接用torch.ones(B,1)和torch.zeros(B,1)即可
二、减少训练epoch数量的策略
1. 迁移学习初始化生成器
用自然图像预训练的SRGAN生成器权重初始化你的模型,然后在植物叶片数据集上微调,能大幅降低收敛所需的epoch数(比如从100epoch降到30-50epoch):
# 加载预训练权重(可以找公开的SRGAN预训练模型) pretrained_dict = torch.load('srgan_pretrained_generator.pth') model_dict = G.state_dict() # 只加载匹配的层(避免因输入通道、输出尺寸不同报错) pretrained_dict = {k: v for k, v in pretrained_dict.items() if k in model_dict and v.shape == model_dict[k].shape} model_dict.update(pretrained_dict) G.load_state_dict(model_dict)
2. 加入早停机制
设置合理的早停规则,当生成器损失连续5-10个epoch没有下降时,直接停止训练,避免无效迭代:
# 新增早停相关变量 patience = 10 early_stop_count = 0 best_loss = float('inf') for epoch in range(1, NUM_EPOCHS + 1): # ... 训练逻辑 ... mean_g = np.mean(g_losses) if mean_g < best_loss: best_loss = mean_g torch.save(G.state_dict(), 'best_generator.pth') early_stop_count = 0 # 重置计数 else: early_stop_count += 1 if early_stop_count >= patience: print(f"Early stopping at epoch {epoch} — no improvement for {patience} epochs") break
三、模型简化方案(牺牲少量精度换速度)
1. 轻量化生成器
把生成器中的残差块数量从16减少到8,或者用深度可分离卷积替换普通卷积,大幅减少参数量和计算量:
# 简化版残差块 class ResidualBlock(nn.Module): def __init__(self, channels): super().__init__() self.conv1 = nn.Conv2d(channels, channels, kernel_size=3, padding=1, bias=False) self.bn1 = nn.BatchNorm2d(channels) self.conv2 = nn.Conv2d(channels, channels, kernel_size=3, padding=1, bias=False) self.bn2 = nn.BatchNorm2d(channels) self.relu = nn.ReLU(inplace=True) # 生成器中残差块数量从16改为8 self.res_blocks = nn.Sequential(*[ResidualBlock(64) for _ in range(8)])
2. 简化判别器
减少判别器的卷积层数,同时降低通道数,比如从64→128→256→512改成64→128→128:
class Discriminator(nn.Module): def __init__(self): super().__init__() self.features = nn.Sequential( nn.Conv2d(3, 64, kernel_size=3, stride=1, padding=1), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(64, 128, kernel_size=3, stride=2, padding=1), nn.BatchNorm2d(128), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(128, 128, kernel_size=3, stride=1, padding=1), nn.BatchNorm2d(128), nn.LeakyReLU(0.2, inplace=True), nn.AdaptiveAvgPool2d(1), nn.Conv2d(128, 1, kernel_size=1) )
优化后的完整训练代码
best_loss = float('inf') # 新增早停参数 patience = 10 early_stop_count = 0 # CPU混合精度开关 use_cpu_amp = True for epoch in range(1, NUM_EPOCHS + 1): G.train(); D.train() g_losses, d_losses = [], [] for lr_img, hr_img in train_dl: # 数据转成CPU友好的精度 lr_img = lr_img.to(DEVICE, dtype=torch.bfloat16 if use_cpu_amp else torch.float32) hr_img = hr_img.to(DEVICE, dtype=torch.bfloat16 if use_cpu_amp else torch.float32) B = lr_img.size(0) real_label = torch.ones (B, 1, device=DEVICE) * 0.9 # 简化标签尺寸 fake_label = torch.zeros(B, 1, device=DEVICE) + 0.1 # -- 判别器训练 -- opt_D.zero_grad() with torch.autocast(device_type='cpu', dtype=torch.bfloat16, enabled=use_cpu_amp): fake_hr = G(lr_img).detach() d_real = D(hr_img).flatten() d_fake = D(fake_hr).flatten() loss_D = (criterion_adv(d_real, real_label) + criterion_adv(d_fake, fake_label)) * 0.5 loss_D.backward() opt_D.step() # -- 生成器训练 -- opt_G.zero_grad() with torch.autocast(device_type='cpu', dtype=torch.bfloat16, enabled=use_cpu_amp): fake_hr = G(lr_img) d_fake = D(fake_hr).flatten() loss_pix = criterion_pix(fake_hr, hr_img) loss_adv = criterion_adv(d_fake, real_label) 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) loss_G.backward() opt_G.step() 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') early_stop_count = 0 else: early_stop_count += 1 if early_stop_count >= patience: print(f"Early stopping at epoch {epoch} — no improvement for {patience} epochs") break 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)
内容的提问来源于stack exchange,提问作者AQSA ZAM ZAM MIRZA JOHAR BAIG
相关产品推荐
相关产品推荐

