PyTorch训练循环首个Epoch后完全停滞,求排查原因
排查PyTorch训练UNet时后续Epoch无训练的问题
核心问题定位
你遇到的后续Epoch训练循环不执行、指标全为0的情况,最可能的原因是DataLoader迭代器提前耗尽。如果你的train_iterator和test_iterator是通过iter(train_loader)这类方式提前生成的单次迭代器,那么第一次Epoch遍历后,迭代器就会走到末尾,后续再调用训练/评估函数时,for循环不会执行任何批次计算,导致loss和acc始终为0。
解决方案
1. 替换迭代器为原始DataLoader
不要使用提前生成的train_iterator/test_iterator,直接传入原始的DataLoader对象(比如train_loader)。PyTorch的DataLoader每次被遍历都会自动重新生成数据迭代,不会出现耗尽问题。
修改训练循环代码:
EPOCHS = 10 print(len(train_loader)) # 改为打印原始DataLoader的长度 for epoch in range(1,EPOCHS+1): print("EPOCH " + str(epoch)) # 传入train_loader而非train_iterator train_loss,train_acc = train(model,train_loader,optimizer,criterion) # 传入test_loader而非test_iterator valid_loss,valid_acc = evaluate(model,test_loader,criterion) # 打印训练统计 print(f'\tTrain Loss: {train_loss:.3f} | Train Acc: {train_acc*100:.2f}%') print(f'\t Val. Loss: {valid_loss:.3f} | Val. Acc: {valid_acc*100:.2f}%')
2. 优化训练/评估函数细节
同时调整训练和评估函数,兼容DataLoader遍历,并修复一些潜在问题:
def train(model, loader, optimizer, criterion): running_loss = 0.0 epoch_acc = 0.0 model.train() # 直接遍历DataLoader,每次自动重置迭代 for (batch_idx,batch) in enumerate(loader): spad,ground_truth = batch # 用统一的device变量管理设备,兼容CPU/GPU环境 spad = spad.to(device) ground_truth = ground_truth.to(device) output = model(spad) loss = criterion(output,ground_truth) acc = ssim_accuracy(output,ground_truth) optimizer.zero_grad() loss.backward() optimizer.step() running_loss += loss.item() epoch_acc += acc.item() if(batch_idx % 10 == 9): # 优化打印逻辑,显示当前批次的平均精度 print(f"Batch {batch_idx+1}: Avg Acc = {epoch_acc/(batch_idx+1):.4f}") return running_loss / len(loader), epoch_acc / len(loader)
def evaluate(model, loader, criterion): epoch_loss = 0.0 epoch_acc = 0.0 with torch.no_grad(): model.eval() # 必须切换到评估模式,关闭BN/Dropout for (batch_idx,batch) in enumerate(loader): spad,ground_truth = batch spad = spad.to(device) ground_truth = ground_truth.to(device) predictions = model(spad) acc = ssim_accuracy(predictions, ground_truth) loss = criterion(predictions,ground_truth) epoch_loss += loss.item() epoch_acc += acc.item() return epoch_loss / len(loader), epoch_acc / len(loader)
额外注意事项
- 设备统一:避免直接使用
.cuda()硬编码设备,用device变量管理,保证代码在CPU环境下也能运行。 - 评估模式切换:评估时必须调用
model.eval(),否则BatchNorm、Dropout等层会继续更新统计量,导致验证结果失真。 - 命名规范:建议将原始数据加载对象命名为
train_loader/test_loader,避免与迭代器混淆。
内容的提问来源于stack exchange,提问作者Anand Idris
相关产品推荐
相关产品推荐

