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

基于PyTorch的SRGAN植物叶片超分辨率模型CPU训练速度极慢:优化方案咨询

优化CPU端SRGAN训练效率的实用方案

我完全理解用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的对抗损失和感知损失会增加训练复杂度,你可以分两步走:

  1. 先只训练像素损失(把LAMBDA_ADV和LAMBDA_PERC设为0),让生成器快速收敛到基线水平
  2. 再加入对抗损失和感知损失微调,只需要少量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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.27 10:37:34