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

基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.27 09:17:39