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

CNN自编码器单epoch训练耗时过长,寻求优化建议

CNN自编码器单epoch训练耗时过长,寻求优化建议

我正在用PyTorch训练一个针对64x64二值图像的CNN自编码器,训练集包含1,079,156个样本,用的batch size是128,当前单轮epoch要花大概3小时。我试过切换BatchNorm到LayerNorm/GroupNorm,但训练速度没变化;也尝试了早停,但这解决不了单epoch耗时的问题。

我的模型代码如下:

import torch
import torch.nn as nn
import torch.nn.functional as F

class CNNAutoencoder(nn.Module):
    """
    Input shape: (B, 1, 64, 64)
    """
    def __init__(self, grid_size=64):
        super().__init__()
        self.grid_size = grid_size

        # 1) Encoder layers
        self.enc_conv1 = nn.Conv2d(1, 16, kernel_size=3, padding=1)
        self.bn_enc1   = nn.BatchNorm2d(16) 
        self.enc_conv2 = nn.Conv2d(16, 32, kernel_size=3, padding=1)
        self.bn_enc2   = nn.BatchNorm2d(32)
        self.pool      = nn.MaxPool2d(kernel_size=2, stride=2)

        # After two pools => shape is (32, grid_size//4, grid_size//4)
        flat_dim    = 32 * (grid_size // 4) * (grid_size // 4)
        latent_dim  = 128

        # For linear layers, we can use BatchNorm1d:
        self.fc_enc    = nn.Linear(flat_dim, latent_dim)
        self.bn_fc_enc = nn.BatchNorm1d(latent_dim)

        # 2) Decoder layers
        self.fc_dec    = nn.Linear(latent_dim, flat_dim)
        self.bn_fc_dec = nn.BatchNorm1d(flat_dim)

        self.dec_tconv1 = nn.ConvTranspose2d(32, 16, kernel_size=2, stride=2)
        self.bn_dec1    = nn.BatchNorm2d(16)
        self.dec_tconv2 = nn.ConvTranspose2d(16, 1, kernel_size=2, stride=2)

    def encoder(self, x):
        # x => (B,1,64,64)
        x = self.enc_conv1(x)          # => (B,16,64,64)
        x = self.bn_enc1(x)           
        x = F.relu(x)
        x = self.pool(x)              # => (B,16,32,32)

        x = self.enc_conv2(x)         # => (B,32,32,32)
        x = self.bn_enc2(x)          
        x = F.relu(x)
        x = self.pool(x)              # => (B,32,16,16)

        # Flatten
        x = x.view(x.size(0), -1)     # => (B, flat_dim=32*16*16)
        x = self.fc_enc(x)            # => (B, latent_dim=128)
        x = self.bn_fc_enc(x)         
        x = F.relu(x)                 
        return x

    def decoder(self, z):
        # z => (B,128)
        x = self.fc_dec(z)            # => (B, flat_dim)
        x = self.bn_fc_dec(x)
        x = F.relu(x)

        # Reshape to (B, 32, 16, 16)
        B = x.size(0)
        x = x.view(B, 32, self.grid_size // 4, self.grid_size // 4)

        x = self.dec_tconv1(x)        # => (B,16,32,32)
        x = self.bn_dec1(x)
        x = F.relu(x)

        # Final upsample => (B,1,64,64)
        x = self.dec_tconv2(x)
        x = torch.sigmoid(x)
        return x

    def forward(self, x):
        z = self.encoder(x)
        return self.decoder(z)

训练代码如下:

train_loader = DataLoader(train_ds, batch_size=128, shuffle=True, num_workers=4, pin_memory=True)
val_loader   = DataLoader(val_ds, batch_size=batch_size, shuffle=False)

# Model
model = CNNAutoencoder(grid_size=64).to(device)
optimizer = optim.Adam(model.parameters(), lr=1e-3)
loss_fn = nn.BCELoss()

best_val_loss = float('inf')
model_fname = f"best_{rt}_fold{fold}.pth"

global_step = 0

for ep in range(num_epochs):
    if global_step >= max_steps:
        print(f"Reached {max_steps} total steps; stopping early.")
        break

    t0 = time.time()

    #############################
    # Train Loop (per epoch)
    #############################
    model.train()
    total_loss = 0.0
    for x_in, x_tgt in train_loader:
        x_in = x_in.to(device)
        x_tgt = x_tgt.to(device)

        optimizer.zero_grad()
        out = model(x_in)
        loss = loss_fn(out, x_tgt)
        loss.backward()
        optimizer.step()

        total_loss += loss.item() * x_in.size(0)

        global_step += 1  # increment step count
        if global_step >= max_steps:
            print(f"Reached {max_steps} total steps; stopping in mid-epoch.")
            break

    train_epoch_loss = total_loss / len(train_loader.dataset)

    # If we already hit max_steps, break out
    if global_step >= max_steps:
        break

    #############################
    # Validation Loop
    #############################
    model.eval()
    val_loss_sum = 0.0
    with torch.no_grad():
        for x_in, x_tgt in val_loader:
            x_in, x_tgt = x_in.to(device), x_tgt.to(device)
            out = model(x_in)
            loss = loss_fn(out, x_tgt)
            val_loss_sum += loss.item() * x_in.size(0)

    current_val_loss = val_loss_sum / len(val_loader.dataset)
    dt = time.time() - t0

    print(f"Epoch [{ep+1}/{num_epochs}] => "
          f"train_loss={train_epoch_loss:.4f}, val_loss={current_val_loss:.4f}, dt={dt:.2f}s, "
          f"global_step={global_step}")

    # Save best model
    if current_val_loss < best_val_loss:
        best_val_loss = current_val_loss
        torch.save(model.state_dict(), model_fname)

兄弟,我来给你捋几个能实打实提速的方向,都是自己训练大数据集踩过坑总结的:

  • 数据加载环节(最容易出效果的优化点)
    你目前的val_loader没加num_workers和pin_memory,训练集用了4个worker但100多万样本可能不够。建议根据你的CPU核心数调整,比如设num_workers=8或16(别超过核心数的一半,避免资源抢占),同时给val_loader补上num_workers=4, pin_memory=True。
    另外试试加persistent_workers=True(PyTorch1.7+支持),这个参数能让worker进程在epoch之间不销毁,省去重启的时间,对大数据集提升很明显。修改后的DataLoader:

    train_loader = DataLoader(train_ds, batch_size=128, shuffle=True, num_workers=8, pin_memory=True, persistent_workers=True)
    val_loader   = DataLoader(val_ds, batch_size=128, shuffle=False, num_workers=4, pin_memory=True, persistent_workers=True)
    

    还要检查你的数据集是不是每次迭代都在重复做预处理?如果是,把预处理逻辑提前到初始化阶段,或者做数据缓存,避免重复计算。

  • 模型与计算逻辑优化
    把最后一层的sigmoid去掉,换成BCEWithLogitsLoss替代原来的BCELoss——这个操作能把sigmoid和loss合并计算,不仅数值更稳定,还能省掉一次单独的sigmoid计算时间,小幅度但稳定提速。
    另外一定要试试混合精度训练,用PyTorch的torch.cuda.amp模块,能把部分计算转成半精度,GPU计算速度能提升30%-50%,对自编码器这种任务精度损失几乎可忽略。训练循环改造示例:

    from torch.cuda.amp import GradScaler, autocast
    
    scaler = GradScaler()
    loss_fn = nn.BCEWithLogitsLoss()  # 替换原来的BCELoss
    
    # 训练循环内修改
    model.train()
    total_loss = 0.0
    for x_in, x_tgt in train_loader:
        x_in = x_in.to(device)
        x_tgt = x_tgt.to(device)
    
        optimizer.zero_grad()
        with autocast():
            out = model(x_in)  # decoder最后去掉sigmoid,直接输出logits
            loss = loss_fn(out, x_tgt)
        
        scaler.scale(loss).backward()
        scaler.step(optimizer)
        scaler.update()
    
        total_loss += loss.item() * x_in.size(0)
        # 其余逻辑不变
    

    验证循环也要加上with autocast():。

  • 训练配置细节优化
    看看你的GPU显存剩余空间,如果还有不少,直接增大batch size!比如从128加到256,单epoch的迭代次数直接砍半,时间能大幅缩短。另外检查训练过程中有没有频繁的磁盘IO操作(比如每步存日志),改成每N步或每个epoch存一次,减少IO等待。

  • 硬件环境兜底检查
    如果数据集在机械硬盘上,赶紧移到SSD!机械硬盘的随机读取速度完全跟不上GPU,会导致GPU经常空等数据。另外确认下torch.cuda.is_available()返回True,确保真的在GPU上训练,别不小心跑CPU上了。

备注:内容来源于stack exchange,提问作者SadnessAndSorrow

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.13 19:12:59