【Bug已解决】How do I save a trained model in PyTorch? 解决方案
2026/8/24 1:52:52 网站建设 项目流程

【Bug已解决】How do I save a trained model in PyTorch? 解决方案

问题描述

在 PyTorch 深度学习开发中,模型训练往往需要消耗大量的时间和计算资源。一个中等规模的模型在 GPU 上训练可能需要数小时甚至数天。因此,将训练好的模型保存到磁盘并在需要时重新加载,是每个 PyTorch 开发者必须掌握的核心技能。

然而,很多初学者在保存和加载模型时会遇到各种令人困惑的问题:

  1. 保存了模型但加载后预测结果全错——这是因为只保存了模型结构而没有正确保存参数。
  2. 加载模型时报错AttributeErrorKeyError——保存方式与加载方式不匹配。
  3. 跨设备加载失败——在 GPU 上训练保存的模型,在 CPU 上加载报错。
  4. 保存的模型文件过大——不知道如何只保存模型权重。
  5. 恢复训练时优化器状态丢失——导致学习率、动量等状态重置,训练不稳定。

这些问题的根本原因在于 PyTorch 提供了两种不同的模型保存方式,开发者如果没有理解它们的区别,就很容易踩坑。

本文将深入剖析 PyTorch 模型保存与加载的完整知识体系,从原理到实践,帮助你彻底掌握这一技能。

错误复现

错误示例一:直接保存整个模型对象

import torch import torch.nn as nn # 定义一个简单的神经网络 class SimpleNet(nn.Module): def __init__(self): super(SimpleNet, self).__init__() self.fc1 = nn.Linear(784, 256) self.fc2 = nn.Linear(256, 10) self.relu = nn.ReLU() def forward(self, x): x = self.relu(self.fc1(x)) x = self.fc2(x) return x # 训练模型 model = SimpleNet() # ... 假设这里进行了大量训练 ... # 错误的保存方式:直接保存整个模型对象 torch.save(model, 'model_complete.pth') print("模型已保存") # 在另一个脚本中尝试加载 # ========== 另一个文件 load_model.py ========== # import torch # model = torch.load('model_complete.pth') # # 报错信息: # AttributeError: Can't get attribute 'SimpleNet' on <module '__main__'>

运行上述加载代码时,你会看到如下报错:

AttributeError: Can't get attribute 'SimpleNet' on <module '__main__'>

这个错误的原因是:torch.save(model, ...)使用了 Python 的pickle序列化机制,它保存的是类的路径字符串而非类定义本身。当你在另一个文件中加载时,Python 找不到SimpleNet类的定义,就会报错。

错误示例二:保存与加载的 map_location 问题

# 在 GPU 上训练并保存 model = SimpleNet().cuda() torch.save(model.state_dict(), 'model_weights.pth') # 在没有 GPU 的机器上加载 model = SimpleNet() model.load_state_dict(torch.load('model_weights.pth')) # 报错信息: # RuntimeError: Attempting to deserialize object on a CUDA device, # but torch.cuda.is_available() is False.

报错输出:

RuntimeError: Attempting to deserialize object on a CUDA device, but torch.cuda.is_available() is False. If you are running on a CPU-only machine, please use torch.load with a map_location=torch.device('cpu') to map your storages to the CPU.

错误示例三:恢复训练时丢失优化器状态

# 第一阶段训练 model = SimpleNet() optimizer = torch.optim.SGD(model.parameters(), lr=0.01, momentum=0.9) # 训练 50 个 epoch 后保存 torch.save(model.state_dict(), 'checkpoint.pth') # 第二阶段恢复训练 model = SimpleNet() model.load_state_dict(torch.load('checkpoint.pth')) optimizer = torch.optim.SGD(model.parameters(), lr=0.01, momentum=0.9) # 问题:优化器的动量累积全部丢失! # 这会导致恢复训练后初期 loss 突然跳变

根因分析

一、PyTorch 的两种保存机制

PyTorch 提供了两种保存模型的方式,理解它们的底层差异是解决所有问题的关键。

方式一:保存整个模型(torch.save(model, path)

这种方式使用 Python 的pickle模块将整个模型对象序列化。pickle在反序列化时需要能够访问到原始类的定义。它保存的内容包括:

  • 模型的类名和模块路径(如__main__.SimpleNet
  • 模型的所有参数(state_dict
  • 模型的所有子模块和缓冲区

致命缺陷pickle不保存类的源代码,只保存类的引用路径。这意味着加载时,你的代码环境中必须存在完全相同的类定义。一旦你重构了代码、改变了文件结构,或者将模型分享给他人,这种方式就会彻底失效。

方式二:保存状态字典(torch.save(model.state_dict(), path)

state_dict是一个 Python 字典,映射了每一层的参数名称到参数张量。例如:

# 打印 state_dict 的键 for key, value in model.state_dict().items(): print(f"{key}: {value.shape}")

输出:

fc1.weight: torch.Size([256, 784]) fc1.bias: torch.Size([256]) fc2.weight: torch.Size([10, 256]) fc2.bias: torch.Size([10])

这种方式只保存纯数据(张量),不依赖任何类定义。加载时,你只需要先创建一个相同结构的模型实例,然后将参数灌入即可。这是 PyTorch 官方推荐的方式。

二、为什么需要保存优化器状态

在训练过程中,优化器(如 SGD with momentum、Adam 等)会维护内部状态变量:

  • SGD with momentum:为每个参数维护一个动量缓冲区
  • Adam:为每个参数维护一阶矩估计和二阶矩估计

如果你只保存模型参数而不保存优化器状态,恢复训练时这些累积的统计量会全部归零。对于 Adam 优化器来说,这意味着:

  1. 自适应学习率需要重新预热
  2. 训练初期会出现 loss 突然跳变
  3. 可能导致模型收敛到更差的局部最优解

三、序列化的底层原理

PyTorch 的torch.save底层使用的是一种自定义的序列化格式(基于 ZIP 格式)。当你调用torch.save时:

  1. PyTorch 遍历所有张量,将它们序列化为连续的内存块
  2. 记录每个张量的元信息(shape、dtype、device)
  3. 将这些信息打包成一个 ZIP 文件

理解这一点很重要,因为它解释了为什么torch.load需要知道map_location——加载时需要决定将张量映射到哪个设备上。

解决方案

方案一:保存和加载模型权重(推荐方式)

这是最通用、最安全的保存方式,适用于绝大多数场景。

import torch import torch.nn as nn class SimpleNet(nn.Module): def __init__(self, input_size=784, hidden_size=256, num_classes=10): super(SimpleNet, self).__init__() self.fc1 = nn.Linear(input_size, hidden_size) self.relu = nn.ReLU() self.fc2 = nn.Linear(hidden_size, num_classes) def forward(self, x): out = self.fc1(x) out = self.relu(out) out = self.fc2(out) return out # ==================== 保存模型权重 ==================== def save_model_weights(model, filepath): """ 只保存模型的 state_dict(推荐方式) 优点:不依赖类定义的路径,跨文件/跨项目安全 """ torch.save(model.state_dict(), filepath) print(f"模型权重已保存到 {filepath}") # ==================== 加载模型权重 ==================== def load_model_weights(model, filepath, device='cpu'): """ 加载模型权重到指定设备 """ # map_location 确保可以在不同设备间迁移 state_dict = torch.load(filepath, map_location=device) model.load_state_dict(state_dict) model.to(device) print(f"模型权重已从 {filepath} 加载") return model # 使用示例 model = SimpleNet() save_model_weights(model, 'model_weights.pth') # 在任何地方加载 new_model = SimpleNet() new_model = load_model_weights(new_model, 'model_weights.pth', device='cpu')

方案二:保存完整的训练检查点(Checkpoint)

当你需要中断并恢复训练时,需要保存更多信息。

def save_checkpoint(epoch, model, optimizer, loss, filepath): """ 保存完整的训练检查点 包含:epoch、模型参数、优化器状态、损失值 """ checkpoint = { 'epoch': epoch, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'loss': loss, } torch.save(checkpoint, filepath) print(f"检查点已保存:epoch={epoch}, loss={loss:.4f}") def load_checkpoint(filepath, model, optimizer=None, device='cpu'): """ 加载训练检查点,恢复到中断前的完整状态 """ checkpoint = torch.load(filepath, map_location=device) model.load_state_dict(checkpoint['model_state_dict']) if optimizer is not None: optimizer.load_state_dict(checkpoint['optimizer_state_dict']) epoch = checkpoint['epoch'] loss = checkpoint['loss'] print(f"检查点已加载:epoch={epoch}, loss={loss:.4f}") return epoch, loss

方案三:处理设备兼容性

def load_model_any_device(model, filepath): """ 自动处理设备兼容性的模型加载函数 无论模型在什么设备上训练保存的,都能正确加载 """ # 检查当前环境是否有 GPU if torch.cuda.is_available(): device = torch.device('cuda') # 先加载到 CPU,再移动到 GPU,避免直接加载到 GPU 时的内存问题 state_dict = torch.load(filepath, map_location='cpu') model.load_state_dict(state_dict) model = model.to(device) else: device = torch.device('cpu') state_dict = torch.load(filepath, map_location='cpu') model.load_state_dict(state_dict) print(f"模型已加载到 {device}") return model, device

完整修复代码

下面是一个完整的、可运行的训练-保存-加载-推理流程:

import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader, TensorDataset import os # ==================== 模型定义 ==================== class MLPClassifier(nn.Module): """多层感知机分类器""" def __init__(self, input_dim=784, hidden_dims=[512, 256], num_classes=10, dropout=0.3): super(MLPClassifier, self).__init__() layers = [] prev_dim = input_dim for hidden_dim in hidden_dims: layers.append(nn.Linear(prev_dim, hidden_dim)) layers.append(nn.BatchNorm1d(hidden_dim)) layers.append(nn.ReLU()) layers.append(nn.Dropout(dropout)) prev_dim = hidden_dim layers.append(nn.Linear(prev_dim, num_classes)) self.network = nn.Sequential(*layers) def forward(self, x): return self.network(x) # ==================== 训练器类 ==================== class ModelTrainer: """完整的模型训练、保存、加载管理器""" def __init__(self, model, learning_rate=0.001, device='cpu'): self.model = model.to(device) self.device = device self.criterion = nn.CrossEntropyLoss() self.optimizer = optim.Adam(model.parameters(), lr=learning_rate) self.train_losses = [] self.val_accuracies = [] def train_epoch(self, train_loader): """训练一个 epoch""" self.model.train() running_loss = 0.0 correct = 0 total = 0 for batch_idx, (data, target) in enumerate(train_loader): data, target = data.to(self.device), target.to(self.device) data = data.view(data.size(0), -1) # 展平 # 前向传播 self.optimizer.zero_grad() output = self.model(data) loss = self.criterion(output, target) # 反向传播 loss.backward() self.optimizer.step() # 统计 running_loss += loss.item() _, predicted = output.max(1) total += target.size(0) correct += predicted.eq(target).sum().item() epoch_loss = running_loss / len(train_loader) epoch_acc = 100. * correct / total return epoch_loss, epoch_acc def save_checkpoint(self, epoch, filepath, extra_info=None): """保存完整的训练检查点""" checkpoint = { 'epoch': epoch, 'model_state_dict': self.model.state_dict(), 'optimizer_state_dict': self.optimizer.state_dict(), 'train_losses': self.train_losses, 'val_accuracies': self.val_accuracies, } if extra_info: checkpoint.update(extra_info) torch.save(checkpoint, filepath) print(f"[Checkpoint] 已保存到 {filepath} (epoch={epoch})") def load_checkpoint(self, filepath): """加载训练检查点,恢复完整训练状态""" checkpoint = torch.load(filepath, map_location=self.device) self.model.load_state_dict(checkpoint['model_state_dict']) self.optimizer.load_state_dict(checkpoint['optimizer_state_dict']) self.train_losses = checkpoint.get('train_losses', []) self.val_accuracies = checkpoint.get('val_accuracies', []) start_epoch = checkpoint['epoch'] + 1 print(f"[Checkpoint] 已从 {filepath} 恢复 (从 epoch {start_epoch} 继续)") return start_epoch def save_inference_model(self, filepath): """只保存模型权重,用于推理部署""" torch.save(self.model.state_dict(), filepath) print(f"[Inference] 推理模型已保存到 {filepath}") # ==================== 完整使用示例 ==================== def main(): # 设置设备 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') print(f"使用设备: {device}") # 创建模拟数据集 num_samples = 1000 X = torch.randn(num_samples, 784) y = torch.randint(0, 10, (num_samples,)) dataset = TensorDataset(X, y) train_loader = DataLoader(dataset, batch_size=32, shuffle=True) # 初始化模型和训练器 model = MLPClassifier(input_dim=784, hidden_dims=[256, 128], num_classes=10) trainer = ModelTrainer(model, learning_rate=0.001, device=device) # 创建保存目录 os.makedirs('checkpoints', exist_ok=True) # 训练循环(带自动检查点保存) num_epochs = 10 checkpoint_path = 'checkpoints/best_model.pth' # 如果存在之前的检查点,则恢复 if os.path.exists(checkpoint_path): start_epoch = trainer.load_checkpoint(checkpoint_path) else: start_epoch = 0 for epoch in range(start_epoch, num_epochs): loss, acc = trainer.train_epoch(train_loader) trainer.train_losses.append(loss) trainer.val_accuracies.append(acc) print(f"Epoch [{epoch+1}/{num_epochs}] Loss: {loss:.4f}, Acc: {acc:.2f}%") # 每 5 个 epoch 保存一次检查点 if (epoch + 1) % 5 == 0: trainer.save_checkpoint(epoch, checkpoint_path, extra_info={'model_arch': 'MLPClassifier'}) # 保存最终推理模型 trainer.save_inference_model('checkpoints/final_inference.pth') # ==================== 模拟在新环境中加载推理 ==================== print("\n===== 模拟推理环境 =====") inference_model = MLPClassifier(input_dim=784, hidden_dims=[256, 128], num_classes=10) state_dict = torch.load('checkpoints/final_inference.pth', map_location='cpu') inference_model.load_state_dict(state_dict) inference_model.eval() # 推理 with torch.no_grad(): test_input = torch.randn(5, 784) predictions = inference_model(test_input) predicted_classes = predictions.argmax(dim=1) print(f"预测结果: {predicted_classes.tolist()}") if __name__ == '__main__': main()

运行结果示例:

使用设备: cpu Epoch [1/10] Loss: 2.3145, Acc: 12.30% Epoch [2/10] Loss: 2.2145, Acc: 18.50% ... [Checkpoint] 已保存到 checkpoints/best_model.pth (epoch=4) ... [Inference] 推理模型已保存到 checkpoints/final_inference.pth ===== 模拟推理环境 ===== 预测结果: [3, 7, 1, 9, 5]

常见陷阱与注意事项

陷阱一:混淆两种保存方式

# 错误:用 state_dict 方式保存,却用完整模型方式加载 torch.save(model.state_dict(), 'model.pth') loaded = torch.load('model.pth') # 这会返回一个 dict,不是模型对象! loaded(...) # 报错:dict 不可调用 # 正确做法 model = SimpleNet() model.load_state_dict(torch.load('model.pth'))

陷阱二:保存和加载时模型结构不匹配

如果你修改了模型结构(比如增加了层数),加载旧的state_dict时会报键不匹配的错误。解决方案是使用strict=False

# 允许部分加载(忽略新增或缺失的层) model.load_state_dict(torch.load('model.pth'), strict=False)

陷阱三:忘记切换到 eval 模式

加载模型用于推理时,必须调用model.eval(),否则 Dropout 和 BatchNorm 的行为不正确:

model.load_state_dict(torch.load('model.pth')) model.eval() # 必须调用! # 或者使用上下文管理器 with torch.no_grad(): output = model(input)

陷阱四:GPU 到 CPU 的设备迁移

# 模型在 GPU 上训练保存,在 CPU 上加载 # 错误方式 model.load_state_dict(torch.load('model.pth')) # RuntimeError: CUDA device not available # 正确方式 state_dict = torch.load('model.pth', map_location='cpu') model.load_state_dict(state_dict)

陷阱五:保存路径的跨平台问题

# 使用 os.path.join 而不是手动拼接路径 import os filepath = os.path.join('checkpoints', 'model.pth') # 跨平台安全 # 而不是 filepath = 'checkpoints/model.pth' # 在 Windows 上可能有问题

陷阱六:版本兼容性

PyTorch 不同版本之间的state_dict格式可能不完全兼容。建议:

  1. 保存时记录 PyTorch 版本
  2. 尽量使用相同版本加载
  3. 如果必须跨版本,使用torch.load(..., weights_only=True)(PyTorch 2.0+)
checkpoint = { 'model_state_dict': model.state_dict(), 'pytorch_version': torch.__version__, # ... }

总结

本文详细讲解了 PyTorch 中模型保存与加载的完整知识体系,核心要点如下:

  1. 始终使用state_dict方式保存torch.save(model.state_dict(), path)),避免保存整个模型对象带来的类路径依赖问题。

  2. 恢复训练时保存完整检查点,包括模型参数、优化器状态、epoch 等信息,确保训练可以无缝恢复。

  3. 注意设备兼容性,使用map_location参数处理 GPU/CPU 之间的迁移。

  4. 推理前切换到 eval 模式,确保 Dropout 和 BatchNorm 行为正确。

  5. 保存时记录元信息(PyTorch 版本、模型结构参数等),方便后续维护和调试。

掌握这些知识后,你就能够安全、高效地管理 PyTorch 模型的持久化,无论是用于推理部署还是断点续训,都能游刃有余。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询