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

