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

PyTorch模型断点续训:如何从保存的epoch启动训练

PyTorch 断点续训:从保存的epoch恢复训练

你已经实现了write_checkpoint和load_checkpoint函数用于保存/加载模型、优化器、调度器状态及epoch,但当前训练流程无法从保存的epoch位置继续执行,以下是具体修正思路和代码调整方案:

核心问题分析

原代码将模型加载逻辑放在训练循环内部,且循环固定从0开始迭代,同时scheduler加载方式存在错误,导致无法正确延续之前的训练进度。

具体修改步骤

1. 将模型加载逻辑移至训练循环前

把加载操作放在循环开始前完成,避免干扰epoch计数,同时增加路径输入的灵活性:

# 询问是否加载模型
load_model = input('Load a model? (y/n)')
start_epoch = 0
if load_model.lower() == 'y':
    checkpoint_path = input('Enter checkpoint path: ')
    model, optimizer, start_epoch, scheduler = load_checkpoint(model=model, scheduler=scheduler, optimizer=optimizer, filename=checkpoint_path)
    # 确保优化器状态迁移到当前设备
    for state in optimizer.state.values():
        for k, v in state.items():
            if isinstance(v, torch.Tensor):
                state[k] = v.to(device)

2. 修正训练循环的起始范围

将循环从start_epoch开始,而非固定从0启动,保证延续之前的训练进度:

# 从保存的epoch开始训练
for epoch in range(start_epoch, num_epochs):
    # 训练、评估逻辑...

3. 修复load_checkpoint中的scheduler加载错误

原代码直接替换scheduler对象的方式会导致状态不匹配,应使用load_state_dict方法加载:

def load_checkpoint(model, optimizer, scheduler, filename='/content/checkpoint.pth'):
    start_epoch = 0
    if os.path.isfile(filename):
        print("=> loading checkpoint '{}'".format(filename))
        checkpoint = torch.load(filename)
        start_epoch = checkpoint['epoch']
        model.load_state_dict(checkpoint['state_dict'])
        optimizer.load_state_dict(checkpoint['optimizer'])
        # 修正scheduler状态加载方式
        scheduler.load_state_dict(checkpoint['scheduler'])
        print("=> loaded checkpoint '{}' (epoch {})".format(filename, checkpoint['epoch']))
    else:
        print("=> no checkpoint found at '{}'".format(filename))
    return model, optimizer, start_epoch, scheduler

4. 调整checkpoint保存时机

原代码在循环内频繁保存并立即加载的逻辑不符合断点续训需求,建议改为每个epoch结束后保存:

# 在epoch训练+评估完成后保存当前状态
write_checkpoint(model=model, epoch=epoch, scheduler=scheduler, optimizer=optimizer)

修改后的完整train_model函数

def train_model(model, train_loader, test_loader, device, learning_rate=1e-1, num_epochs=200):
    criterion = nn.CrossEntropyLoss()
    model.to(device)

    optimizer = optim.SGD(model.parameters(), lr=learning_rate, momentum=0.9, weight_decay=1e-4)
    scheduler = torch.optim.lr_scheduler.MultiStepLR(optimizer, milestones=[65, 75], gamma=0.75, last_epoch=-1)

    # 初始评估
    model.eval()
    eval_loss, eval_accuracy = evaluate_model(model=model, test_loader=test_loader, device=device, criterion=criterion)
    print("Epoch: {:02d} Eval Loss: {:.3f} Eval Acc: {:.3f}".format(-1, eval_loss, eval_accuracy))

    # 加载模型逻辑
    load_model = input('Load a model? (y/n)')
    start_epoch = 0
    if load_model.lower() == 'y':
        checkpoint_path = input('Enter checkpoint path: ')
        model, optimizer, start_epoch, scheduler = load_checkpoint(model=model, scheduler=scheduler, optimizer=optimizer, filename=checkpoint_path)
        # 迁移优化器状态到设备
        for state in optimizer.state.values():
            for k, v in state.items():
                if isinstance(v, torch.Tensor):
                    state[k] = v.to(device)

    # 从保存的epoch开始训练
    for epoch in range(start_epoch, num_epochs):
        model.train()
        running_loss = 0
        running_corrects = 0

        for inputs, labels in train_loader:
            inputs = torch.FloatTensor(inputs).to(device)
            labels = labels.to(device)

            optimizer.zero_grad()
            outputs = model(inputs)
            _, preds = torch.max(outputs, 1)
            loss = criterion(outputs, labels)
            loss.backward()
            optimizer.step()

            running_loss += loss.item() * inputs.size(0)
            running_corrects += torch.sum(preds == labels.data)

        train_loss = running_loss / len(train_loader.dataset)
        train_accuracy = running_corrects / len(train_loader.dataset)

        # 评估
        model.eval()
        eval_loss, eval_accuracy = evaluate_model(model=model, test_loader=test_loader, device=device, criterion=criterion)

        scheduler.step()

        print("Epoch: {:03d} Train Loss: {:.3f} Train Acc: {:.3f} Eval Loss: {:.3f} Eval Acc: {:.3f}".format(epoch, train_loss, train_accuracy, eval_loss, eval_accuracy))

        # 保存当前epoch的checkpoint
        write_checkpoint(model=model, epoch=epoch, scheduler=scheduler, optimizer=optimizer)

    return model

内容的提问来源于stack exchange,提问作者Manu Dwivedi

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.15 12:15:47