3.5 模型保存与加载 (Saving & Loading Models) PyTorch 模型保存与加载: 训练成果的持久化与再现 在深度学习模型的训练过程中,我们投入大量的时间和计算资源,最终目的是获得一个在特定任务上表现良好的模型。然而,训练过程并非一蹴而就,往往需要多次迭代和调整。此外,训练好的模型也需要在不同的场景中使用,例如模型部署、模型微调、模型共享等。因此,模型保存与加载成为了深度学习工作流程中至关重要的一环。 3.5 模型保存与加载 (Saving & Loading Models) 模型保存与加载主要解决以下几个核心问题: 持久化训练成果: 将训练好的模型参数和结构保存到磁盘,防止因程序意外中断或硬件故障导致训练成果丢失。
在深度学习模型的训练过程中,我们投入大量的时间和计算资源,最终目的是获得一个在特定任务上表现良好的模型。然而,训练过程并非一蹴而就,往往需要多次迭代和调整。此外,训练好的模型也需要在不同的场景中使用,例如模型部署、模型微调、模型共享等。因此,模型保存与加载成为了深度学习工作流程中至关重要的一环。
模型保存与加载主要解决以下几个核心问题:
持久化训练成果: 将训练好的模型参数和结构保存到磁盘,防止因程序意外中断或硬件故障导致训练成果丢失。
模型复用与部署: 加载已保存的模型,无需重新训练即可直接用于推理预测或部署到生产环境。
断点续训: 在训练过程中定期保存模型状态,以便在训练中断后从上次保存的状态继续训练,节省时间和计算资源。
模型迁移与共享: 将训练好的模型分享给他人,或在不同的环境和设备上加载模型。
PyTorch 提供了灵活而强大的机制来实现模型的保存与加载,主要涉及以下两种核心方法:
保存与加载模型状态字典 (State Dictionary): 这是推荐的保存和加载模型参数的方法,它仅保存模型的权重和偏置等可学习参数,不包含模型的结构信息。
保存与加载整个模型 (Entire Model): 这种方法保存了模型的完整结构和参数,包括模型的类定义、层结构以及权重参数。
接下来,我们将分别详细介绍这两种方法,并通过代码示例进行演示。
状态字典 (State Dictionary) 是 PyTorch 模型中一个核心概念。它本质上是一个 Python 字典 (OrderedDict),将模型中每一层的可学习参数 (例如卷积层的权重 weight 和偏置 bias,线性层的 weight 和 bias 等) 映射到它们的张量值。 状态字典不包含模型的结构信息,只关注模型的参数。
为什么推荐保存状态字典?
灵活性高: 状态字典只保存模型的参数,不依赖于模型的类定义。这意味着即使模型的代码发生改变,只要模型结构保持一致,我们仍然可以使用状态字典加载参数。
跨平台性好: 状态字典本质上是张量数据的序列化,可以方便地在不同的平台和设备之间迁移。
安全性高: 相比保存整个模型,保存状态字典可以避免潜在的代码执行风险,因为我们只需要加载参数数据,而不需要执行模型的代码。
节省存储空间: 状态字典通常比保存整个模型文件更小,因为它只包含参数数据,不包含模型的代码结构。
如何保存状态字典?
在 PyTorch 中,我们可以使用 torch.save() 函数来保存状态字典。 在保存之前,我们需要先获取模型的状态字典,可以通过 model.state_dict() 方法获得。
代码实践 3.5.1.1: 保存模型状态字典
import torch import torch.nn as nn # 定义一个简单的模型 class SimpleNet(nn.Module): def __init__(self, input_size, hidden_size, output_size): super(SimpleNet, self).__init__() self.fc1 = nn.Linear(input_size, hidden_size) self.relu = nn.ReLU() self.fc2 = nn.Linear(hidden_size, output_size) def forward(self, x): x = self.fc1(x) x = self.relu(x) x = self.fc2(x) return x # 创建模型实例 input_size = 784 hidden_size = 128 output_size = 10 model = SimpleNet(input_size, hidden_size, output_size) # 打印模型的状态字典 print("模型状态字典 (保存前):") print(model.state_dict().keys()) # 查看状态字典的键 # 指定保存的文件路径 save_path = 'model_state_dict.pth' # 使用 torch.save() 保存模型状态字典 torch.save(model.state_dict(), save_path) print(f"\n模型状态字典已保存到: {save_path}") # 清空模型的状态字典,模拟重新加载 model.load_state_dict({}) # 清空状态字典 print("\n模型状态字典 (清空后):") print(model.state_dict().keys()) # 加载保存的状态字典 loaded_state_dict = torch.load(save_path) # 将加载的状态字典加载到模型中 model.load_state_dict(loaded_state_dict) print("\n模型状态字典 (加载后):") print(model.state_dict().keys()) print("\n模型状态字典 (加载后,部分参数示例):") print(model.state_dict()['fc1.weight'][0][:5]) # 打印第一层线性层权重的部分值
代码详解 3.5.1.1:
定义模型 SimpleNet: 我们首先定义了一个简单的神经网络模型 SimpleNet,包含两个线性层和一个 ReLU 激活函数。
创建模型实例 model: 我们实例化了 SimpleNet 模型。
打印模型状态字典: model.state_dict() 方法返回一个 OrderedDict,其中包含了模型的可学习参数。我们打印了状态字典的键 (keys),可以看到类似 'fc1.weight', 'fc1.bias', 'fc2.weight', 'fc2.bias' 等键,分别对应第一层和第二层线性层的权重和偏置。
指定保存路径 save_path: 我们指定了保存状态字典的文件路径为 model_state_dict.pth。 .pth 或 .pt 是常用的 PyTorch 模型文件扩展名。
torch.save(model.state_dict(), save_path): 这是关键的保存步骤。torch.save() 函数将 model.state_dict() 返回的状态字典保存到指定的文件路径 save_path。
清空和加载状态字典 (模拟加载过程): 为了演示加载过程,我们首先使用 model.load_state_dict({}) 清空了模型的状态字典,然后使用 torch.load(save_path) 加载了之前保存的状态字典,并使用 model.load_state_dict(loaded_state_dict) 将加载的状态字典加载到模型中。
验证加载结果: 我们再次打印了模型的状态字典,并打印了第一层线性层权重的部分值,可以看到加载后的模型参数与保存前的参数一致,说明状态字典加载成功。
Graph TD 图 3.5.1.1: 保存模型状态字典流程
图示解释 3.5.1.1:
模型实例 (Model Instance): 表示已经训练好或需要保存的模型实例。
model.state_dict(): 调用模型的 state_dict() 方法,提取模型的参数信息。
状态字典 (State Dictionary): 提取出的参数信息以 Python 字典的形式存储,键是层名称和参数名称,值是参数张量。
torch.save(): 使用 torch.save() 函数将状态字典序列化并保存到文件中。
模型状态字典文件 (model_state_dict.pth): 最终保存的模型状态字典文件,通常使用 .pth 或 .pt 扩展名。
加载模型状态字典的过程是保存过程的逆向操作。我们需要先加载状态字典文件,然后将其加载到模型实例中。
如何加载状态字典?
创建模型实例: 首先需要创建一个与保存状态字典时结构相同的模型实例。注意:模型的结构必须与保存状态字典时完全一致,否则加载可能会出错或导致模型性能下降。
加载状态字典文件: 使用 torch.load() 函数从文件路径加载状态字典。
加载状态字典到模型: 使用 model.load_state_dict(loaded_state_dict) 方法将加载的状态字典应用到模型实例中。
代码实践 3.5.2.1: 加载模型状态字典
import torch import torch.nn as nn # 定义模型结构 (必须与保存模型时结构一致) class SimpleNet(nn.Module): def __init__(self, input_size, hidden_size, output_size): super(SimpleNet, self).__init__() self.fc1 = nn.Linear(input_size, hidden_size) self.relu = nn.ReLU() self.fc2 = nn.Linear(hidden_size, output_size) def forward(self, x): x = self.fc1(x) x = self.relu(x) x = self.fc2(x) return x # 创建新的模型实例 (结构相同) input_size = 784 hidden_size = 128 output_size = 10 loaded_model = SimpleNet(input_size, hidden_size, output_size) # 指定状态字典文件路径 (与保存时路径一致) load_path = 'model_state_dict.pth' # 加载状态字典文件 loaded_state_dict = torch.load(load_path) # 将加载的状态字典加载到新的模型实例中 loaded_model.load_state_dict(loaded_state_dict) print("加载后的模型状态字典 (部分参数示例):") print(loaded_model.state_dict()['fc1.weight'][0][:5]) # 打印第一层线性层权重的部分值 # 可以使用加载的模型进行推理 # 假设输入数据 input_tensor = torch.randn(1, input_size) # 模拟一个batch size为1的输入 # 设置模型为评估模式 (重要!) loaded_model.eval() # 进行推理 with torch.no_grad(): # 在推理阶段禁用梯度计算,提高效率 output = loaded_model(input_tensor) print("\n模型推理输出:") print(output)
代码详解 3.5.2.1:
定义模型结构 SimpleNet (与保存时一致): 我们再次定义了 SimpleNet 模型,确保其结构与保存状态字典的模型结构完全一致。
创建新的模型实例 loaded_model: 我们创建了一个新的 SimpleNet 模型实例 loaded_model,用于加载状态字典。
指定状态字典文件路径 load_path: 我们指定了之前保存的状态字典文件路径 model_state_dict.pth。
loaded_state_dict = torch.load(load_path): 使用 torch.load() 函数从 load_path 加载状态字典文件,返回加载的状态字典 loaded_state_dict。
loaded_model.load_state_dict(loaded_state_dict): 将加载的状态字典 loaded_state_dict 加载到新的模型实例 loaded_model 中。
验证加载结果: 我们打印了加载后的模型状态字典的部分参数,确认参数已成功加载。
模型推理示例: 我们展示了如何使用加载的模型进行推理。
loaded_model.eval(): 非常重要! 在进行推理或评估之前,需要将模型设置为评估模式 model.eval()。这会影响到某些层的行为,例如 Dropout 和 BatchNorm 层,在评估模式下它们会停止随机失活和参数更新,而使用训练阶段的统计量。
with torch.no_grad():: 在推理阶段,我们通常不需要计算梯度,使用 torch.no_grad() 上下文管理器可以禁用梯度计算,提高推理效率并减少内存占用。
output = loaded_model(input_tensor): 将输入数据 input_tensor 输入到加载的模型中进行前向传播,得到模型的输出 output。
Graph TD 图 3.5.2.1: 加载模型状态字典流程
图示解释 3.5.2.1:
模型状态字典文件 (model_state_dict.pth): 之前保存的模型状态字典文件。
torch.load(): 使用 torch.load() 函数从文件中加载状态字典。
加载的状态字典 (Loaded State Dictionary): 从文件加载得到的模型参数信息,以 Python 字典形式存储。
新的模型实例 (New Model Instance): 新创建的模型实例,结构必须与保存状态字典的模型结构一致。
loaded_model.load_state_dict(): 调用新模型实例的 load_state_dict() 方法,将加载的状态字典应用到模型中。
加载参数后的模型实例 (Model Instance with Loaded Parameters): 最终得到加载了参数的模型实例,可以用于推理或进一步训练。
除了保存状态字典,PyTorch 也允许我们保存整个模型对象。这种方法会保存模型的完整结构和参数,包括模型的类定义、层结构以及权重参数。
如何保存整个模型?
与保存状态字典类似,我们仍然使用 torch.save() 函数,但是这次我们直接将整个模型实例作为参数传递给 torch.save()。
代码实践 3.5.3.1: 保存整个模型
import torch import torch.nn as nn # 定义模型 (与之前相同) class SimpleNet(nn.Module): def __init__(self, input_size, hidden_size, output_size): super(SimpleNet, self).__init__() self.fc1 = nn.Linear(input_size, hidden_size) self.relu = nn.ReLU() self.fc2 = nn.Linear(hidden_size, output_size) def forward(self, x): x = self.fc1(x) x = self.relu(x) x = self.fc2(x) return x # 创建模型实例 input_size = 784 hidden_size = 128 output_size = 10 model = SimpleNet(input_size, hidden_size, output_size) # 指定保存文件路径 save_path = 'entire_model.pth' # 使用 torch.save() 保存整个模型 torch.save(model, save_path) print(f"整个模型已保存到: {save_path}")
代码详解 3.5.3.1:
torch.save() 的第一个参数直接是 model,而不是 model.state_dict()。这表示我们保存的是整个模型对象。Graph TD 图 3.5.3.1: 保存整个模型流程
图示解释 3.5.3.1:
模型实例 (Model Instance): 表示需要保存的整个模型实例。
torch.save(): 使用 torch.save() 函数将整个模型对象序列化并保存到文件中。
整个模型文件 (entire_model.pth): 最终保存的包含模型结构和参数的完整模型文件,通常使用 .pth 或 .pt 扩展名。
加载整个模型同样使用 torch.load() 函数,并直接将加载的模型文件赋值给一个新的模型变量。
代码实践 3.5.4.1: 加载整个模型
import torch # 指定模型文件路径 (与保存时路径一致) load_path = 'entire_model.pth' # 加载整个模型 loaded_model = torch.load(load_path) print("加载后的模型:") print(loaded_model) # 打印加载的模型结构 # 可以使用加载的模型进行推理 (与加载状态字典后推理步骤相同) input_size = 784 input_tensor = torch.randn(1, input_size) loaded_model.eval() with torch.no_grad(): output = loaded_model(input_tensor) print("\n模型推理输出:") print(output)
代码详解 3.5.4.1:
loaded_model = torch.load(load_path): 这是加载整个模型的关键步骤。torch.load() 函数直接从 load_path 加载整个模型对象,并将其赋值给 loaded_model 变量。
打印加载的模型结构: 我们打印了 loaded_model,可以看到加载的模型结构信息,包括模型的类定义和层结构。
模型推理示例: 推理步骤与加载状态字典后的推理步骤完全相同。
Graph TD 图 3.5.4.1: 加载整个模型流程
图示解释 3.5.4.1:
整个模型文件 (entire_model.pth): 之前保存的整个模型文件。
torch.load(): 使用 torch.load() 函数从文件中加载整个模型对象。
加载的整个模型实例 (Loaded Entire Model Instance): 从文件加载得到的完整模型实例,包含模型结构和参数,可以直接用于推理或进一步训练。
在训练过程中,除了模型参数,优化器 (Optimizer) 的状态 (例如学习率、动量等) 也需要保存。这对于断点续训非常重要。我们可以同时保存模型的状态字典和优化器的状态字典。
代码实践 3.5.5.1: 保存和加载模型与优化器状态
import torch import torch.nn as nn import torch.optim as optim # 定义模型 (与之前相同) class SimpleNet(nn.Module): # ... (模型定义与之前相同) ... pass # 创建模型实例 input_size = 784 hidden_size = 128 output_size = 10 model = SimpleNet(input_size, hidden_size, output_size) # 定义优化器 optimizer = optim.Adam(model.parameters(), lr=0.001) # 模拟训练过程 (简单示例) for epoch in range(2): # 训练2个epoch for i in range(10): # 每个epoch 训练10个batch (简化) input_data = torch.randn(32, input_size) # 模拟 batch 数据 target_data = torch.randint(0, output_size, (32,)) # 模拟 target 数据 criterion = nn.CrossEntropyLoss() # 交叉熵损失函数 optimizer.zero_grad() # 梯度清零 output = model(input_data) # 前向传播 loss = criterion(output, target_data) # 计算损失 loss.backward() # 反向传播 optimizer.step() # 更新参数 if (i+1) % 5 == 0: # 每5个batch 保存一次 checkpoint checkpoint_path = f'checkpoint_epoch_{epoch+1}_batch_{i+1}.pth' torch.save({ 'epoch': epoch + 1, 'batch': i + 1, 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), 'loss': loss.item(), }, checkpoint_path) print(f"保存 checkpoint 到: {checkpoint_path}, Epoch: {epoch+1}, Batch: {i+1}, Loss: {loss.item():.4f}") print("\n训练完成,模型和优化器状态已保存为 checkpoints.") # ---- 加载 checkpoint 并继续训练 ---- print("\n---- 加载 checkpoint 并继续训练 ----") # 创建新的模型实例和优化器实例 (结构和参数与之前相同) loaded_model = SimpleNet(input_size, hidden_size, output_size) loaded_optimizer = optim.Adam(loaded_model.parameters(), lr=0.001) # 指定 checkpoint 文件路径 (加载最新的 checkpoint) checkpoint = torch.load('checkpoint_epoch_2_batch_10.pth') # 加载模型状态字典和优化器状态字典 loaded_model.load_state_dict(checkpoint['model_state_dict']) loaded_optimizer.load_state_dict(checkpoint['optimizer_state_dict']) epoch = checkpoint['epoch'] batch = checkpoint['batch'] loss = checkpoint['loss'] print(f"已加载 checkpoint, Epoch: {epoch}, Batch: {batch}, Loss: {loss:.4f}") # 继续训练 (从加载的 checkpoint 开始) start_epoch = epoch start_batch = batch num_epochs = 3 # 再训练 1 个 epoch (总共 3 个 epoch) for epoch in range(start_epoch, num_epochs): for i in range(start_batch, 10): # 继续从上次 batch 开始训练 input_data = torch.randn(32, input_size) target_data = torch.randint(0, output_size, (32,)) criterion = nn.CrossEntropyLoss() loaded_optimizer.zero_grad() output = loaded_model(input_data) loss = criterion(output, target_data) loss.backward() loaded_optimizer.step() if (i+1) % 5 == 0: checkpoint_path = f'checkpoint_epoch_{epoch+1}_batch_{i+1}.pth' torch.save({ 'epoch': epoch + 1, 'batch': i + 1, 'model_state_dict': loaded_model.state_dict(), 'optimizer_state_dict': loaded_optimizer.state_dict(), 'loss': loss.item(), }, checkpoint_path) print(f"保存 checkpoint 到: {checkpoint_path}, Epoch: {epoch+1}, Batch: {i+1}, Loss: {loss.item():.4f}") start_batch = 1 # 下一个 epoch 从 batch 1 开始 print("\n继续训练完成,模型和优化器状态已更新.")
代码详解 3.5.5.1:
定义模型和优化器: 我们定义了模型 SimpleNet 和优化器 optim.Adam。
模拟训练过程: 我们模拟了一个简单的训练过程,训练 2 个 epoch,每个 epoch 10 个 batch。
保存 Checkpoint: 在训练过程中,我们每 5 个 batch 保存一次 checkpoint。
checkpoint_path = f'checkpoint_epoch_{epoch+1}_batch_{i+1}.pth': checkpoint 文件名包含 epoch 和 batch 信息,方便管理。
torch.save({ ... }, checkpoint_path): 我们将一个 Python 字典作为 checkpoint 保存,字典中包含了:
'epoch': 当前 epoch 数。
'batch': 当前 batch 数。
'model_state_dict': 模型状态字典。
'optimizer_state_dict': 优化器状态字典。
'loss': 当前 batch 的损失值。
加载 Checkpoint 并继续训练:
创建新的模型和优化器实例: 创建与之前训练时结构和参数相同的新的模型和优化器实例。
checkpoint = torch.load('checkpoint_epoch_2_batch_10.pth'): 加载最新的 checkpoint 文件。
加载状态字典: 分别使用 loaded_model.load_state_dict(checkpoint['model_state_dict']) 和 loaded_optimizer.load_state_dict(checkpoint['optimizer_state_dict']) 加载模型状态字典和优化器状态字典。
恢复训练状态: 从 checkpoint 中恢复 epoch、batch 和 loss 信息。
继续训练: 从加载的 checkpoint 状态继续训练,注意循环的起始 epoch 和 batch 需要根据加载的 checkpoint 信息调整。
Graph TD 图 3.5.5.1: 保存和加载 Checkpoint 流程
图示解释 3.5.5.1:
训练过程 (Training Process): 表示模型训练的流程。
模型实例 & 优化器 (Model & Optimizer): 训练的模型和优化器实例。
训练循环 (Training Loop): 模型的训练迭代过程。
保存 Checkpoint (Save Checkpoint): 在训练过程中定期保存模型和优化器状态。
Checkpoint 文件 (checkpoint_epoch_*.pth): 保存的 checkpoint 文件,包含模型和优化器状态。
断点续训 (Resuming Training): 表示从 checkpoint 恢复训练的流程。
Checkpoint 文件 (checkpoint_epoch_*.pth): 之前保存的 checkpoint 文件。
加载 Checkpoint (Load Checkpoint): 加载 checkpoint 文件,读取模型和优化器状态。
恢复模型 & 优化器状态 (Restore Model & Optimizer State): 将加载的状态应用到新的模型和优化器实例。
继续训练 (Continue Training): 从恢复的状态继续进行模型训练。
在实际应用中,我们可能会在不同的设备上训练和部署模型,例如在 GPU 上训练模型,然后在 CPU 上进行推理,或者在不同的 GPU 设备之间迁移模型。PyTorch 提供了 map_location 参数来处理设备映射问题。
map_location 参数的作用:
torch.load() 函数的 map_location 参数允许我们将模型加载到指定的设备上。它可以接受以下几种类型的参数:
None (默认值): 如果模型是在 GPU 上保存的,并且当前环境有 GPU 可用,则模型会被加载到 GPU 上。如果当前环境没有 GPU 可用,则会报错。
torch.device 对象: 例如 torch.device('cpu') 或 torch.device('cuda:0'),指定模型加载到的设备。
函数: 可以自定义一个函数来映射设备。
代码实践 3.5.6.1: 将 GPU 上训练的模型加载到 CPU 上