1. 为什么“保存与加载”是PyTorch实战的生死线
在PyTorch的日常开发中,我们常常沉迷于模型架构的精妙设计、损失函数的反复调优,或是数据增强的奇技淫巧。然而,一个看似基础却至关重要的环节——模型的保存与加载——却往往被新手甚至一些有经验的开发者所轻视。我见过太多令人扼腕的场景:一个训练了三天三夜的复杂模型,因为保存不当,在程序意外退出后一切归零;一个精心调优的模型权重,在部署到另一台机器时因为版本或环境差异而无法加载,排查过程痛苦不堪;更常见的是,在实验的中间阶段,因为没有保存中间状态,导致无法从某个检查点(Checkpoint)恢复,只能从头再来。
这不仅仅是“保存一个文件”那么简单。它关乎你工作的可复现性、实验的连续性以及模型交付的可靠性。torch.save和torch.load这两个函数,几乎出现在每一个PyTorch项目的关键路径上。理解它们背后的机制、掌握正确的使用模式、并规避那些隐藏的陷阱,是每一位PyTorch使用者从“玩具代码”走向“生产级应用”的必经之路。今天,我们就来彻底拆解PyTorch中保存与加载的方方面面,从最基础的张量操作,到复杂的模型部署,让你真正掌握这条“生死线”。
2. 核心基石:张量(Tensor)的保存与加载
在深入模型之前,我们必须先夯实基础。PyTorch中的所有数据,无论是中间计算结果、模型参数还是优化器状态,其本质都是张量。因此,理解张量的序列化是理解一切保存操作的前提。
2.1torch.save与torch.load的基本用法
PyTorch使用Python的pickle模块作为其序列化引擎,torch.save和torch.load是对pickle的高层封装,专门为PyTorch对象(尤其是张量)进行了优化。
import torch # 创建一个示例张量 x = torch.randn(3, 4, requires_grad=True) print(f"原始张量: \n{x}") print(f"requires_grad: {x.requires_grad}") # 保存张量到文件 torch.save(x, 'tensor.pt') # .pt 或 .pth 是PyTorch保存文件的常见扩展名 # 从文件加载张量 x_loaded = torch.load('tensor.pt') print(f"\n加载后的张量: \n{x_loaded}") print(f"requires_grad: {x_loaded.requires_grad}")注意:保存的文件扩展名
.pt或.pth没有强制规定,只是一种社区约定。torch.save生成的文件本质上是一个Python pickle文件,内部包含了重建对象所需的所有信息。
关键点解析:
- 保存了什么?不仅仅是张量的数据(
data),还包括其元数据,如形状(shape)、数据类型(dtype)、设备信息(device)以及是否计算梯度(requires_grad)。加载后,这些属性会完全恢复。 - 文件格式:虽然扩展名自定义,但内容结构是PyTorch定义的。你可以尝试用
pickle.load打开一个.pt文件,会发现里面是一个包含序列化数据的字典。
2.2 保存与加载张量字典和列表
实际项目中,我们很少只保存单个张量。更常见的场景是保存一组相关的张量,例如一个批次的输入输出、一组模型中间层的特征图等。
# 保存一个包含多个张量的字典 data_dict = { 'input_tensor': torch.randn(10, 3, 224, 224), 'label_tensor': torch.randint(0, 10, (10,)), 'metadata': {'batch_id': 5, 'epoch': 2} } torch.save(data_dict, 'batch_data.pt') # 加载字典 loaded_dict = torch.load('batch_data.pt') print(f"加载的输入形状: {loaded_dict['input_tensor'].shape}") print(f"元数据: {loaded_dict['metadata']}") # 保存一个张量列表 tensor_list = [torch.ones(2,2), torch.zeros(3,3)] torch.save(tensor_list, 'tensor_list.pt')为什么是字典?使用字典结构进行保存,键值对提供了明确的语义信息(如'input','target','model_state'),这比单纯依赖列表的索引顺序要可靠得多,尤其是在多人协作或长时间后回顾代码时。这是一种强烈推荐的最佳实践。
2.3 跨设备加载:CPU与GPU的兼容性问题
这是早期最容易踩的坑之一。如果你在GPU上创建了一个张量并保存,然后在没有GPU或不同GPU索引的环境下加载,会发生什么?
# 假设在GPU 0上保存 if torch.cuda.is_available(): tensor_gpu = torch.randn(5, 5).cuda() torch.save(tensor_gpu, 'tensor_gpu.pt') # 在另一个可能没有GPU的环境加载 try: loaded_tensor = torch.load('tensor_gpu.pt') print(f"加载成功,设备: {loaded_tensor.device}") except RuntimeError as e: print(f"加载失败: {e}")直接加载可能会失败,因为.cuda()保存的张量包含了指向特定GPU内存的引用。解决方案是使用map_location参数,它允许你重定向张量到指定的设备。
# 方法1:强制加载到CPU loaded_on_cpu = torch.load('tensor_gpu.pt', map_location=torch.device('cpu')) print(f"设备: {loaded_on_cpu.device}") # 输出: cpu # 方法2:加载到当前可用的GPU(例如GPU 0) loaded_on_gpu = torch.load('tensor_gpu.pt', map_location=torch.device('cuda:0')) print(f"设备: {loaded_on_gpu.device}") # 输出: cuda:0 # 方法3:通用写法,自动映射到当前设备 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') loaded_auto = torch.load('tensor_gpu.pt', map_location=device)经验之谈:为了最大程度的可移植性,一个常见的技巧是在保存前将模型或张量转移到CPU。这样保存的文件在任何环境下都能直接加载,无需关心map_location。
# 保存前的“净化”操作 tensor_to_save = tensor_gpu.cpu() if tensor_gpu.is_cuda else tensor_gpu torch.save(tensor_to_save, 'tensor_portable.pt') # 之后在任何地方都可以直接用 torch.load('tensor_portable.pt')3. 模型保存的三种核心模式与陷阱
保存模型比保存张量复杂,因为“模型”包含多个层面:仅参数、完整模型定义、训练状态等。PyTorch提供了灵活的保存方式,但用错了地方就会导致灾难。
3.1 模式一:仅保存模型参数(state_dict)——最推荐的做法
这是PyTorch社区最主流、最推荐的模型保存方式。state_dict是一个Python字典对象,它将模型每一层映射到其对应的参数张量(权重和偏置)。
import torch.nn as nn import torch.optim as optim # 定义一个简单模型 class SimpleNet(nn.Module): def __init__(self): super().__init__() self.fc1 = nn.Linear(10, 5) self.fc2 = nn.Linear(5, 2) def forward(self, x): return self.fc2(torch.relu(self.fc1(x))) model = SimpleNet() optimizer = optim.SGD(model.parameters(), lr=0.01) # 保存模型的 state_dict torch.save(model.state_dict(), 'model_weights.pth') print("模型state_dict的键:", list(model.state_dict().keys())) # 输出类似: ['fc1.weight', 'fc1.bias', 'fc2.weight', 'fc2.bias']为什么推荐只保存state_dict?
- 文件小巧:只保存参数,不保存模型类定义、前向传播逻辑等代码,文件体积最小。
- 灵活性高:加载时,你需要先实例化一个结构完全相同的模型类,然后调用
load_state_dict方法。这强制要求你的模型定义代码是可用的,保证了模型结构与参数的一致性。 - 安全:避免了通过pickle加载整个类定义可能带来的安全风险(pickle可以执行任意代码)。
对应的加载方式:
# 1. 必须重新实例化模型结构 loaded_model = SimpleNet() # 类定义必须在当前作用域可用 # 2. 加载参数 loaded_model.load_state_dict(torch.load('model_weights.pth')) # 3. 将模型设置为评估模式(如果用于推理) loaded_model.eval()关键陷阱:如果保存后修改了模型类的定义(例如增减了层、改了层名),那么加载state_dict时会因为键不匹配而报错KeyError。此时需要用到strict=False参数,并手动处理不匹配的键。
# 假设新模型比旧模型多了一个层 try: loaded_model.load_state_dict(torch.load('model_weights.pth'), strict=False) except Exception as e: print(f"加载出错: {e}") # 使用 strict=False 会忽略不匹配的键,只加载能匹配上的参数。 # 加载后,新加的层将保持随机初始化状态。3.2 模式二:保存整个模型对象——便捷但有风险
你可以直接把模型实例model保存下来。
torch.save(model, 'entire_model.pth') # 加载 loaded_entire_model = torch.load('entire_model.pth') loaded_entire_model.eval()优点:极其方便,一行代码保存,一行代码加载,连模型类定义都不需要。
致命缺点与风险:
- pickle依赖:这种方式使用Python的pickle来序列化整个模型对象。pickle的序列化结果与模型类定义的具体代码路径紧密绑定。如果你移动了模型类定义的文件,或者修改了类的名称、导入方式,加载时就可能失败。
- 安全风险:pickle在反序列化时会执行保存文件中的字节码。如果
.pth文件来自不可信的来源,加载它可能执行恶意代码。 - 框架版本兼容性:不同版本的PyTorch在内部实现上可能有细微差别,直接pickle整个模型对象可能导致跨版本加载失败。
结论:除非是快速原型验证或临时保存,绝不推荐在生产环境或需要长期维护的项目中使用此方法。state_dict是唯一正确的长期选择。
3.3 模式三:保存检查点(Checkpoint)——训练过程的救星
在长时间的训练任务中(如训练大语言模型、扩散模型),我们不仅需要保存模型参数,还需要保存优化器状态、当前的epoch数、损失记录等,以便在训练中断后能精准恢复。这就是检查点。
# 假设在某个训练循环中 epoch = 10 train_loss = 0.5 model = SimpleNet() optimizer = optim.SGD(model.parameters(), lr=0.01) scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=30, gamma=0.1) # ... 进行了一些训练 ... # 保存检查点 checkpoint = { 'epoch': epoch, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'scheduler_state_dict': scheduler.state_dict() if scheduler else None, 'loss': train_loss, # 可以保存任何其他你想记录的信息,如随机数生成器状态 'rng_state': torch.get_rng_state(), } torch.save(checkpoint, 'checkpoint_epoch_{}.pth'.format(epoch)) print("检查点已保存,包含键:", checkpoint.keys())对应的恢复训练流程:
# 恢复训练 def resume_training(checkpoint_path, model, optimizer, scheduler=None): checkpoint = torch.load(checkpoint_path) # 恢复模型 model.load_state_dict(checkpoint['model_state_dict']) # 恢复优化器(非常重要!优化器内部有动量等状态) optimizer.load_state_dict(checkpoint['optimizer_state_dict']) # 恢复学习率调度器 if scheduler and 'scheduler_state_dict' in checkpoint and checkpoint['scheduler_state_dict']: scheduler.load_state_dict(checkpoint['scheduler_state_dict']) # 恢复随机数状态,保证数据加载顺序一致(如果重要) torch.set_rng_state(checkpoint['rng_state']) start_epoch = checkpoint['epoch'] + 1 # 从下一轮开始 print(f"从第 {start_epoch} 轮恢复训练,上一轮损失: {checkpoint['loss']}") return start_epoch # 使用 loaded_model = SimpleNet() loaded_optimizer = optim.SGD(loaded_model.parameters(), lr=0.01) start_epoch = resume_training('checkpoint_epoch_10.pth', loaded_model, loaded_optimizer)检查点策略:
- 定期保存:例如每N个epoch保存一次。
- 保存最佳模型:在验证集上监控指标(如准确率),只保存指标最好的那个检查点。
- 滚动保存:只保留最新的K个检查点,避免磁盘被占满。
4. 实战中的高级场景与疑难杂症
掌握了基本模式后,我们来看看那些让开发者头疼的进阶问题。
4.1 多GPU训练(DataParallel/DistributedDataParallel)模型的保存与加载
当你使用nn.DataParallel或nn.parallel.DistributedDataParallel包装模型进行多GPU训练时,模型参数被分布到了多个GPU上。直接保存包装后的模型state_dict,键名会带有module.前缀。
import torch.nn as nn model = SimpleNet() if torch.cuda.device_count() > 1: model = nn.DataParallel(model) # 或者 nn.parallel.DistributedDataParallel(model) model.cuda() # 训练后保存 torch.save(model.state_dict(), 'dp_model.pth') # 查看保存的键 checkpoint = torch.load('dp_model.pth', map_location='cpu') print(list(checkpoint.keys())[:2]) # 输出类似: ['module.fc1.weight', 'module.fc1.bias']如果你试图用这个state_dict去加载一个没有被DataParallel包装的普通模型,会因为键名不匹配(缺少module.前缀)而失败。
解决方案1:保存前去除module.前缀(推荐)在保存之前,如果模型是DataParallel包装的,先获取其底层模块的state_dict。
# 保存时 if isinstance(model, nn.DataParallel) or isinstance(model, nn.parallel.DistributedDataParallel): state_dict = model.module.state_dict() # 关键!获取内部模块的state_dict else: state_dict = model.state_dict() torch.save(state_dict, 'model_correct.pth') # 现在保存的键是 ['fc1.weight', 'fc1.bias', ...],没有`module.`前缀解决方案2:加载时处理module.前缀如果拿到的是一个带module.前缀的检查点,而你的新模型不是并行化的,可以手动去除前缀。
def load_weights_for_single_gpu(model, checkpoint_path): checkpoint = torch.load(checkpoint_path, map_location='cpu') # 创建一个新的state_dict,去掉 `module.` 前缀 new_state_dict = {} for k, v in checkpoint.items(): name = k[7:] if k.startswith('module.') else k # 去掉 'module.' new_state_dict[name] = v # 加载处理后的state_dict model.load_state_dict(new_state_dict) return model4.2 自定义层与复杂对象的保存
如果你的模型包含了非nn.Module的自定义对象(如一个复杂的损失函数类、一个数据处理器),并且这个对象有自己的状态需要保存,你需要确保这个对象本身是可pickle的。通常,让这个类继承自nn.Module或使其所有属性都是Python基本类型/PyTorch张量,就能保证可序列化。
class CustomLayer: def __init__(self, param): self.param = param # 如果param是张量或基本类型,没问题 self.cache = [] # 如果cache是列表,且里面都是可pickle对象,也没问题 # 如果包含文件句柄、网络连接等不可pickle对象,就会出错 # 更安全的做法是继承nn.Module class CustomSafeLayer(nn.Module): def __init__(self, param): super().__init__() self.param = nn.Parameter(torch.tensor(param)) # 注册为参数 self.register_buffer('cache', torch.zeros(10)) # 注册为buffer,也会被保存register_buffer的妙用:有些张量是模型的一部分,需要被保存和加载,但又不是需要梯度更新的参数(例如BatchNorm中的running_mean)。这时应该使用self.register_buffer('name', tensor)将其注册为buffer,它会被包含在state_dict中,但不会被优化器更新。
4.3 模型版本控制与兼容性处理
随着项目迭代,模型结构会变化。如何加载旧版本的模型参数到新版本模型中?
键名映射:如果只是层名改了,但结构没变,可以创建一个键名映射字典。
old_to_new = {'old_fc.weight': 'new_fc.weight', 'old_fc.bias': 'new_fc.bias'} old_state_dict = torch.load('old_model.pth') new_state_dict = {} for old_key, new_key in old_to_new.items(): if old_key in old_state_dict: new_state_dict[new_key] = old_state_dict[old_key] # 然后加载 new_state_dict,并用 strict=False model.load_state_dict(new_state_dict, strict=False)参数形状不匹配:如果新层和旧层参数形状不同(例如全连接层输入输出维度变了),旧参数无法直接加载。通常的策略是初始化新层,然后尽可能加载能匹配的部分。对于新增的层,它们会保持随机初始化。
使用
strict=False并分析缺失/多余的键:这是最常用的调试手段。missing_keys, unexpected_keys = model.load_state_dict(torch.load('checkpoint.pth'), strict=False) print(f"缺失的键(新模型有,检查点没有): {missing_keys}") print(f"多余的键(检查点有,新模型没有): {unexpected_keys}")根据打印信息,你可以判断是版本不匹配,还是加载错了文件。
4.4 部署优化:保存为TorchScript或ONNX
对于生产部署,我们通常不直接使用PyTorch的.pth文件,而是将其转换为更高效、与Python解耦的格式。
TorchScript:PyTorch自带的序列化格式,可以将模型转换为静态图,提高推理速度,并能在C++等环境中运行。
# 追踪模式 (Tracing) - 适用于无控制流的模型 example_input = torch.randn(1, 10) traced_script_module = torch.jit.trace(model, example_input) traced_script_module.save("model_traced.pt") # 脚本模式 (Scripting) - 适用于包含控制流的模型 scripted_model = torch.jit.script(model) scripted_model.save("model_scripted.pt")ONNX:开放的神经网络交换格式,支持在不同框架(PyTorch, TensorFlow, MXNet等)之间转换模型。
torch.onnx.export(model, example_input, "model.onnx", input_names=["input"], output_names=["output"], dynamic_axes={'input': {0: 'batch_size'}, 'output': {0: 'batch_size'}})
这两种格式的保存和加载,是模型部署流水线中的关键一步,其复杂性和注意事项足以单独成文。
5. 一个完整的训练循环保存示例
让我们将所有知识点整合到一个简单的、健壮的训练循环中。
import os import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader # 假设已有 dataset, model, train_loop 等定义 def train_model(model, train_loader, val_loader, epochs, device, save_dir='checkpoints'): os.makedirs(save_dir, exist_ok=True) optimizer = optim.Adam(model.parameters()) scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min') criterion = nn.CrossEntropyLoss() model.to(device) best_val_loss = float('inf') start_epoch = 0 # 尝试从最新的检查点恢复 checkpoint_files = [f for f in os.listdir(save_dir) if f.endswith('.pth')] if checkpoint_files: latest_checkpoint = max([os.path.join(save_dir, f) for f in checkpoint_files], key=os.path.getctime) print(f"发现检查点: {latest_checkpoint},尝试恢复...") checkpoint = torch.load(latest_checkpoint, map_location=device) model.load_state_dict(checkpoint['model_state_dict']) optimizer.load_state_dict(checkpoint['optimizer_state_dict']) scheduler.load_state_dict(checkpoint['scheduler_state_dict']) start_epoch = checkpoint['epoch'] + 1 best_val_loss = checkpoint.get('best_val_loss', best_val_loss) print(f"从 epoch {start_epoch} 恢复训练。") for epoch in range(start_epoch, epochs): # 训练阶段 model.train() train_loss = 0.0 for batch in train_loader: # ... 训练步骤 ... pass # 实际训练代码 # 验证阶段 model.eval() val_loss = 0.0 with torch.no_grad(): for batch in val_loader: # ... 验证步骤 ... pass # 实际验证代码 scheduler.step(val_loss) # --- 保存逻辑 --- # 1. 定期保存检查点 if (epoch + 1) % 5 == 0: checkpoint = { 'epoch': epoch, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'scheduler_state_dict': scheduler.state_dict(), 'train_loss': train_loss, 'val_loss': val_loss, 'best_val_loss': best_val_loss, } torch.save(checkpoint, os.path.join(save_dir, f'checkpoint_epoch_{epoch+1:03d}.pth')) print(f"已保存周期检查点到 epoch_{epoch+1:03d}.pth") # 2. 保存最佳模型(仅参数) if val_loss < best_val_loss: best_val_loss = val_loss # 保存前,如果模型是DataParallel,获取其内部模块 if isinstance(model, nn.DataParallel): state_to_save = model.module.state_dict() else: state_to_save = model.state_dict() torch.save(state_to_save, os.path.join(save_dir, 'best_model_weights.pth')) print(f"*** 发现更优验证损失 {val_loss:.4f}, 已保存最佳模型参数。") # 3. 保存最后一个epoch的模型(完整状态,便于完整恢复) if epoch == epochs - 1: final_checkpoint = { 'epoch': epoch, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'scheduler_state_dict': scheduler.state_dict(), 'final_train_loss': train_loss, 'final_val_loss': val_loss, } torch.save(final_checkpoint, os.path.join(save_dir, 'final_checkpoint.pth')) # 训练结束后,加载最佳模型进行最终评估或导出 best_model = YourModelClass() # 重新实例化结构 best_model.load_state_dict(torch.load(os.path.join(save_dir, 'best_model_weights.pth'), map_location=device)) best_model.to(device) best_model.eval() print("训练完成,最佳模型已加载。") return best_model这个示例涵盖了从恢复训练、定期检查点、保存最佳模型到最终模型加载的完整生命周期管理,是大多数项目可以直接参考的模板。
6. 常见“坑”与排查清单
即使知道了所有原理,实际操作中依然会出错。下面是一个快速排查清单:
RuntimeError: [enforce fail at inline_container.cc:145] . PytorchStreamReader failed reading zip archive: failed finding central directory原因:文件损坏或不是有效的PyTorch保存文件。可能是下载不完整,或文件被错误地修改。解决:重新下载或生成文件。检查文件大小是否异常。KeyError: ‘xxx’或Missing key(s) in state_dict原因:模型结构定义与保存的state_dict不匹配。可能是层名更改、模型类定义错误、或加载了错误的检查点文件。解决:- 打印
model.state_dict().keys()和checkpoint.keys()进行对比。 - 使用
strict=False加载,并检查missing_keys和unexpected_keys。 - 确认你是否在加载一个
DataParallel模型的state_dict到一个普通模型上(需要去除module.前缀)。
- 打印
加载后模型性能骤降或输出全是乱码原因:
- 忘记调用
model.eval()。这会导致Dropout层仍然生效,BatchNorm层使用训练时的统计量,从而引入随机性。 - 加载了错误的权重文件(例如分类数不同的模型)。
- 数据预处理方式与训练时不一致。解决:推理前务必
model.eval();确认模型权重与任务匹配;统一数据预处理流程。
- 忘记调用
GPU内存不足(OOM) when loading原因:试图将一个巨大的模型直接加载到GPU上。解决:先加载到CPU,再转移到GPU。
checkpoint = torch.load('huge_model.pth', map_location='cpu') model.load_state_dict(checkpoint) model.to('cuda')跨PyTorch版本加载失败原因:PyTorch不同版本间内部API可能有变动。解决:尽量在相同的PyTorch版本环境下进行保存和加载。如果必须跨版本,优先使用
state_dict方式,并做好测试。对于非常重要的模型,可以考虑同时保存state_dict和导出为ONNX/TorchScript作为备份。
掌握PyTorch的保存与加载,就像为你的模型上了保险。它不能直接提升模型精度,但能保证你的心血不会因为一次断电、一次误操作或一次环境迁移而白费。花时间设计一个健壮的保存加载策略,是每个严肃的深度学习项目不可或缺的一部分。从今天起,别再只用torch.save(model, ‘model.pth’)了。