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

加载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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.14 01:38:14