【Bug已解决】How do I save a trained model in PyTorch? 解决方案
【Bug已解决】How do I save a trained model in PyTorch? 解决方案问题描述在 PyTorch 深度学习开发中模型训练往往需要消耗大量的时间和计算资源。一个中等规模的模型在 GPU 上训练可能需要数小时甚至数天。因此将训练好的模型保存到磁盘并在需要时重新加载是每个 PyTorch 开发者必须掌握的核心技能。然而很多初学者在保存和加载模型时会遇到各种令人困惑的问题保存了模型但加载后预测结果全错——这是因为只保存了模型结构而没有正确保存参数。加载模型时报错AttributeError或KeyError——保存方式与加载方式不匹配。跨设备加载失败——在 GPU 上训练保存的模型在 CPU 上加载报错。保存的模型文件过大——不知道如何只保存模型权重。恢复训练时优化器状态丢失——导致学习率、动量等状态重置训练不稳定。这些问题的根本原因在于 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: Cant get attribute SimpleNet on module __main__运行上述加载代码时你会看到如下报错AttributeError: Cant 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_locationtorch.device(cpu) to map your storages to the CPU.错误示例三恢复训练时丢失优化器状态# 第一阶段训练 model SimpleNet() optimizer torch.optim.SGD(model.parameters(), lr0.01, momentum0.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(), lr0.01, momentum0.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 优化器来说这意味着自适应学习率需要重新预热训练初期会出现 loss 突然跳变可能导致模型收敛到更差的局部最优解三、序列化的底层原理PyTorch 的torch.save底层使用的是一种自定义的序列化格式基于 ZIP 格式。当你调用torch.save时PyTorch 遍历所有张量将它们序列化为连续的内存块记录每个张量的元信息shape、dtype、device将这些信息打包成一个 ZIP 文件理解这一点很重要因为它解释了为什么torch.load需要知道map_location——加载时需要决定将张量映射到哪个设备上。解决方案方案一保存和加载模型权重推荐方式这是最通用、最安全的保存方式适用于绝大多数场景。import torch import torch.nn as nn class SimpleNet(nn.Module): def __init__(self, input_size784, hidden_size256, num_classes10): 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, devicecpu): 加载模型权重到指定设备 # map_location 确保可以在不同设备间迁移 state_dict torch.load(filepath, map_locationdevice) 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, devicecpu)方案二保存完整的训练检查点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, optimizerNone, devicecpu): 加载训练检查点恢复到中断前的完整状态 checkpoint torch.load(filepath, map_locationdevice) 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_locationcpu) model.load_state_dict(state_dict) model model.to(device) else: device torch.device(cpu) state_dict torch.load(filepath, map_locationcpu) 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_dim784, hidden_dims[512, 256], num_classes10, dropout0.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_rate0.001, devicecpu): self.model model.to(device) self.device device self.criterion nn.CrossEntropyLoss() self.optimizer optim.Adam(model.parameters(), lrlearning_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_infoNone): 保存完整的训练检查点 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_locationself.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_size32, shuffleTrue) # 初始化模型和训练器 model MLPClassifier(input_dim784, hidden_dims[256, 128], num_classes10) trainer ModelTrainer(model, learning_rate0.001, devicedevice) # 创建保存目录 os.makedirs(checkpoints, exist_okTrue) # 训练循环带自动检查点保存 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(fEpoch [{epoch1}/{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_dim784, hidden_dims[256, 128], num_classes10) state_dict torch.load(checkpoints/final_inference.pth, map_locationcpu) 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(dim1) 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 (epoch4) ... [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时会报键不匹配的错误。解决方案是使用strictFalse# 允许部分加载忽略新增或缺失的层 model.load_state_dict(torch.load(model.pth), strictFalse)陷阱三忘记切换到 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_locationcpu) 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格式可能不完全兼容。建议保存时记录 PyTorch 版本尽量使用相同版本加载如果必须跨版本使用torch.load(..., weights_onlyTrue)PyTorch 2.0checkpoint { model_state_dict: model.state_dict(), pytorch_version: torch.__version__, # ... }总结本文详细讲解了 PyTorch 中模型保存与加载的完整知识体系核心要点如下始终使用state_dict方式保存torch.save(model.state_dict(), path)避免保存整个模型对象带来的类路径依赖问题。恢复训练时保存完整检查点包括模型参数、优化器状态、epoch 等信息确保训练可以无缝恢复。注意设备兼容性使用map_location参数处理 GPU/CPU 之间的迁移。推理前切换到 eval 模式确保 Dropout 和 BatchNorm 行为正确。保存时记录元信息PyTorch 版本、模型结构参数等方便后续维护和调试。掌握这些知识后你就能够安全、高效地管理 PyTorch 模型的持久化无论是用于推理部署还是断点续训都能游刃有余。