加载PyTorch模型后重训练损失持续上升问题求助
模型重训练损失持续上升的排查方案
问题背景
首次训练模型时,无论训练10、20还是30个epoch,准确率持续上升、损失不断下降,表现正常。但加载已保存的模型重训练时,损失每个epoch都持续上升。尝试同时保存并加载optimizer的state_dict后,问题仍未解决。当前使用Adam优化器与CrossEntropyLoss损失函数。
训练代码
def train( epochs: int, batch_size: int, net: torch.nn.Module, trainDataLoader: DataLoader, testDataLoader: DataLoader, device: str, lossF: torch.nn.modules.loss._WeightedLoss, optimizer: torch.optim.Optimizer, save_path: str, ): for epoch in range(1, epochs + 1): processBar = tqdm(trainDataLoader, unit="step") net.train(True) for step, (train_seq, train_labels) in enumerate(processBar): train_seq = train_seq.to(device) train_labels = train_labels.to(device) optimizer.zero_grad() outputs = net(train_seq) loss = lossF(outputs, train_labels) predictions = torch.argmax(outputs, dim=1) accuracy = torch.sum(predictions == train_labels) / train_labels.shape[0] loss.backward() optimizer.step() processBar.set_description( "[%d/%d] Loss: %.4f, Acc: %.4f" % (epoch, epochs, loss.item(), accuracy.item()) ) if step == len(processBar) - 1: correct, total_loss = 0, 0 net.train(False) with torch.no_grad(): for test_seq, test_labels in testDataLoader: test_seq = test_seq.to(device) test_labels = test_labels.to(device) test_out = net(test_seq) tloss = lossF(test_out, test_labels) predictions = torch.argmax(test_out, dim=1) total_loss += tloss correct += torch.sum(predictions == test_labels) test_acc = correct / (batch_size * len(testDataLoader)) test_loss = total_loss / len(testDataLoader) processBar.set_description( "[%d/%d] Loss: %.4f, Acc: %.4f, Test Loss: %.4f, Test Acc: %.4f" % ( epoch, epochs, loss.item(), accuracy.item(), test_loss.item(), test_acc.item(), ) ) model_save_path = os.path.join(save_path, "checkpoint.pt") with open(model_save_path, "wb") as f: torch.save(net.state_dict(), f) processBar.close() def main(): conf = config.AllConfig model_path = os.path.join(conf.save_path, "checkpoint.pt") model = TaxonClassifier.TaxonModel( vocab_size=conf.vocab_size, embedding_size=conf.embedding_size, hidden_size=conf.hidden_size, device=conf.device, max_len=conf.max_len, num_layers=conf.num_layers, num_class=conf.num_class, drop_out=conf.drop_prob, ) model = model.to(device=conf.device) optimizer = torch.optim.Adam(model.parameters(), lr=conf.lr) if os.path.exists(model_path) is True: print("Loading existing model state_dict......") checkpoint = torch.load(model_path, map_location=conf.device, weights_only=True) model.load_state_dict(checkpoint) else: print("No existing model state......") print("Loading Dict Files......") all_dict = Dataset.Dictionary(conf.KmerFilePath, conf.TaxonFilePath) print("Loading dataset......") all_dataset = Dataset.AllDataset(conf.DataPath, conf.max_len, all_dict, conf.kmer) train_dataloader = DataLoader( dataset=all_dataset.train_dataset, batch_size=conf.batch_size, shuffle=True, num_workers=4, ) test_dataloader = DataLoader( dataset=all_dataset.test_dataset, batch_size=conf.batch_size, shuffle=False, num_workers=4, ) lossF = torch.nn.CrossEntropyLoss() print("Start Training") train( epochs=conf.epoch, batch_size=conf.batch_size, net=model, trainDataLoader=train_dataloader, testDataLoader=test_dataloader, device=conf.device, lossF=lossF, optimizer=optimizer, save_path=conf.save_path, )
模型代码
import torch from torch import nn from torch.nn import functional as F from . import LSTMLayer, EmbeddingLayer class TaxonModel(nn.Module): def __init__( self, vocab_size: int, embedding_size: int, hidden_size: int, device: str, max_len: int, num_layers: int, num_class: int, drop_out: float = 0.5, ): super(TaxonModel, self).__init__() self.num_layers = num_layers self.num_class = num_class self.vocab_size = vocab_size self.embedding_size = embedding_size self.hidden_size = hidden_size self.device = device self.max_len = max_len self.drop_out = drop_out self.seq_encoder = LSTMLayer.SeqEncoder( embedding_size, hidden_size, num_layers, drop_out ) self.embedding = EmbeddingLayer.FullEmbedding( vocab_size, embedding_size, max_len, device, drop_out ) # attention相关 self.key_matrix = nn.Parameter( torch.Tensor(hidden_size * 2, hidden_size * 2), requires_grad=True ) self.query = nn.Parameter(torch.Tensor(hidden_size * 2), requires_grad=True) # 初始化矩阵参数 nn.init.uniform_(self.key_matrix, -0.1, 0.1) nn.init.uniform_(self.query, -0.1, 0.1) # 解码器,输出class self.decoder = nn.Sequential( nn.Linear(hidden_size * 2, hidden_size), nn.BatchNorm1d(hidden_size), nn.GELU(), nn.Dropout(drop_out), nn.Linear(hidden_size, num_class), ) def forward(self, x): x = self.embedding(x) # [batch_size,seq_len,emb_size] x = x.permute(1, 0, 2) # [seq_len,batch_size,emb_size] x = self.seq_encoder(x) # x: [seq_len,batch_size,hidden_size*2] x = x.permute(1, 0, 2) # [batch_size,seq_len,hidden_size*2] key = torch.tanh( torch.matmul(x, self.key_matrix) ) # [seq_len,batch_size,hidden_size*2] # torch.matmul(key,self.query)的结果为 [batch_size,seq_len]因为做的是内积 # 再对第1维做softmax score = F.softmax(torch.matmul(key, self.query), dim=1).unsqueeze( -1 ) # [batch_size,seq_len,1] x = x * score # [batch_size,seq_len,hidden_size*2] x = torch.sum(x, dim=1) # [batch_size,hidden_size*2] final_outputs = self.decoder(x) return final_outputs
排查方向
1. 优化器状态未正确同步
当前代码仅保存/加载模型state_dict,即使尝试过加载优化器状态,需确认逻辑是否正确:
- 保存时:需同时存储模型和优化器状态:
torch.save({ 'model_state_dict': net.state_dict(), 'optimizer_state_dict': optimizer.state_dict() }, model_save_path) - 加载时:初始化优化器后再恢复其状态:
checkpoint = torch.load(model_path, map_location=conf.device, weights_only=True) model.load_state_dict(checkpoint['model_state_dict']) optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
2. 学习率不匹配
重训练时若使用和首次训练相同的初始学习率,模型已接近收敛,过大的学习率会导致参数在最优值附近震荡发散:
- 重训练时将学习率降至原数值的1/10;
- 加入学习率调度器(如
torch.optim.lr_scheduler.ReduceLROnPlateau),根据验证损失自动调整。
3. 数据加载异常
检查重训练时数据集是否与首次训练一致:
- 确认
all_dataset.train_dataset是否重新划分,或预处理环节(如shuffle)的随机性是否导致数据分布变化; - 临时将
num_workers设为0,排查多线程加载数据的异常。
4. BatchNorm层状态未恢复
模型包含BatchNorm1d,首次训练时会累积均值和方差,仅加载模型参数会导致重训练时使用初始统计值,破坏收敛状态:
- 确认
net.state_dict()已包含BatchNorm的running_mean和running_var,加载模型时完整恢复即可。
5. 设备与模式设置
- 检查加载模型时
map_location是否正确映射到目标设备,确保模型和数据在同一设备; - 确认重训练时模型已设置为
train()模式(当前代码net.train(True)正确),避免Dropout、BatchNorm处于评估状态。
6. 损失计算验证
- 确认重训练时损失函数输入输出维度匹配,CrossEntropyLoss的输入是模型原始输出(未经过softmax),标签是类别索引;
- 检查训练时显示的是单batch损失,验证损失是否为所有测试batch的平均值,避免统计误差导致误判。
内容的提问来源于stack exchange,提问作者LOYINuts
相关产品推荐
相关产品推荐

