模型参数的读写

掌握PyTorch中模型参数的保存与加载技术,学习训练过程中的检查点管理

发布于 2025-08-05 更新于 2025-08-051,045 字 3 分钟
模型参数的读写的文章封面
展开文章目录
  1. 1. 模型参数读写的重要性
  2. 2. 模型参数的保存
  3. 2.1 保存模型参数
  4. 2.2 保存优化器状态
  5. 2.3 联合保存(推荐)
  6. 3. 模型参数的加载
  7. 3.1 设备兼容性处理
  8. 3.2 加载模型参数
  9. 3.3 加载优化器状态
  10. 3.4 联合加载
  11. 4. 高级应用
  12. 4.1 训练过程中的自动保存
  13. 4.2 模型版本管理
  14. 4.3 跨平台兼容性
  15. 5. 常见问题与解决方案
  16. 5.1 常见错误及解决方法
  17. 5.2 性能优化建议
  18. 5.3 实际应用示例
  19. 参考资料

1. 模型参数读写的重要性

以 PyTorch 为例,可以使用torch.load()torch.save()函数读写包括张量在内的任意 Python 对象,如:列表、字典、模型的结构与参数、优化器状态等。


2. 模型参数的保存

2.1 保存模型参数

基本保存方法

使用Module.state_dict()方法可获取模型Module当前状态 (state) 的全部参数字典,包含了权重和偏置:

import torch
import torch.nn as nn

# 定义一个简单的模型
model = nn.Sequential(
    nn.Linear(10, 5),
    nn.ReLU(),
    nn.Linear(5, 1)
)

# 保存模型参数
torch.save(model.state_dict(), 'model.pth')

保存完整模型(不推荐)

# 直接保存完整模型(不推荐)
torch.save(model, 'complete_model.pth')

2.2 保存优化器状态

可以同样使用state_dict()方法保存当前的优化器状态:

import torch.optim as optim

# 定义优化器
optimizer = optim.Adam(model.parameters(), lr=0.001)

# 保存优化器状态
torch.save(optimizer.state_dict(), 'optimizer.pth')

2.3 联合保存(推荐)

保存策略

基本联合保存:

由于torch.save()函数支持保存任意 Python 对象,可用字典组织优化器和参数等状态后联合保存:

# 联合保存多个状态
state = {
    'model_state_dict': model.state_dict(),
    'optimizer_state_dict': optimizer.state_dict(),
    'epoch': 10,
    'loss': 0.1
}
torch.save(state, 'checkpoint.pth')

完整检查点保存:

import time

def save_checkpoint(model, optimizer, epoch, loss, filepath):
    """保存训练检查点"""
    checkpoint = {
        'model_state_dict': model.state_dict(),
        'optimizer_state_dict': optimizer.state_dict(),
        'epoch': epoch,
        'loss': loss,
        'model_architecture': str(model),
        'timestamp': torch.tensor(time.time())
    }
    torch.save(checkpoint, filepath)
    print(f"检查点已保存到: {filepath}")

3. 模型参数的加载

3.1 设备兼容性处理

3.2 加载模型参数

标准加载方法

使用Module.load_state_dict()方法可从文件对象加载状态字典(反序列化):

# 创建模型实例(结构必须与保存时一致)
model = nn.Sequential(
    nn.Linear(10, 5),
    nn.ReLU(),
    nn.Linear(5, 1)
)

# 加载模型参数
model.load_state_dict(torch.load('model.pth'))
model.eval()  # 设置为评估模式

加载完整模型

# 加载完整模型(如果之前保存的是完整模型)
model = torch.load('complete_model.pth')

3.3 加载优化器状态

# 创建优化器实例
optimizer = optim.Adam(model.parameters(), lr=0.001)

# 加载优化器状态
optimizer.load_state_dict(torch.load('optimizer.pth'))

3.4 联合加载

加载策略

基本联合加载:

# 加载检查点
checkpoint = torch.load('checkpoint.pth')

# 恢复模型和优化器状态
model.load_state_dict(checkpoint['model_state_dict'])
optimizer.load_state_dict(checkpoint['optimizer_state_dict'])

# 恢复训练信息
epoch = checkpoint['epoch']
loss = checkpoint['loss']

print(f"恢复到第 {epoch} 轮训练,损失值:{loss}")

安全加载(带错误处理):

def load_checkpoint(model, optimizer, filepath):
    """安全加载训练检查点"""
    try:
        checkpoint = torch.load(filepath, map_location='cpu')
        
        # 检查必要的键是否存在
        required_keys = ['model_state_dict', 'optimizer_state_dict', 'epoch', 'loss']
        for key in required_keys:
            if key not in checkpoint:
                raise KeyError(f"检查点文件缺少必要的键: {key}")
        
        # 加载状态
        model.load_state_dict(checkpoint['model_state_dict'])
        optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
        
        epoch = checkpoint['epoch']
        loss = checkpoint['loss']
        
        print(f"成功加载检查点: epoch={epoch}, loss={loss:.4f}")
        return epoch, loss
        
    except FileNotFoundError:
        print(f"检查点文件不存在: {filepath}")
        return 0, float('inf')
    except Exception as e:
        print(f"加载检查点时出错: {e}")
        return 0, float('inf')

4. 高级应用

4.1 训练过程中的自动保存

import os
import time

class CheckpointManager:
    def __init__(self, save_dir='checkpoints', max_keep=5):
        self.save_dir = save_dir
        self.max_keep = max_keep
        os.makedirs(save_dir, exist_ok=True)
        
    def save(self, model, optimizer, epoch, loss, metrics=None):
        """保存检查点"""
        timestamp = int(time.time())
        filename = f'checkpoint_epoch_{epoch}_{timestamp}.pth'
        filepath = os.path.join(self.save_dir, filename)
        
        checkpoint = {
            'model_state_dict': model.state_dict(),
            'optimizer_state_dict': optimizer.state_dict(),
            'epoch': epoch,
            'loss': loss,
            'timestamp': timestamp
        }
        
        if metrics:
            checkpoint['metrics'] = metrics
            
        torch.save(checkpoint, filepath)
        self._cleanup_old_checkpoints()
        return filepath
    
    def _cleanup_old_checkpoints(self):
        """清理旧的检查点文件"""
        checkpoints = []
        for filename in os.listdir(self.save_dir):
            if filename.startswith('checkpoint_') and filename.endswith('.pth'):
                filepath = os.path.join(self.save_dir, filename)
                checkpoints.append((filepath, os.path.getctime(filepath)))
        
        # 按创建时间排序,保留最新的几个
        checkpoints.sort(key=lambda x: x[1], reverse=True)
        for filepath, _ in checkpoints[self.max_keep:]:
            os.remove(filepath)
            print(f"删除旧检查点: {filepath}")

# 使用示例
checkpoint_manager = CheckpointManager()

# 在训练循环中
for epoch in range(num_epochs):
    # ... 训练代码 ...
    
    if epoch % 10 == 0:  # 每10个epoch保存一次
        checkpoint_manager.save(model, optimizer, epoch, train_loss)

4.2 模型版本管理

class ModelVersionManager:
    def __init__(self, base_path='models'):
        self.base_path = base_path
        os.makedirs(base_path, exist_ok=True)
    
    def save_model(self, model, version, metadata=None):
        """保存模型版本"""
        version_dir = os.path.join(self.base_path, f'v{version}')
        os.makedirs(version_dir, exist_ok=True)
        
        # 保存模型参数
        model_path = os.path.join(version_dir, 'model.pth')
        torch.save(model.state_dict(), model_path)
        
        # 保存元数据
        if metadata:
            metadata_path = os.path.join(version_dir, 'metadata.json')
            import json
            with open(metadata_path, 'w') as f:
                json.dump(metadata, f, indent=2)
        
        print(f"模型版本 v{version} 已保存到 {version_dir}")
    
    def load_model(self, model, version):
        """加载指定版本的模型"""
        model_path = os.path.join(self.base_path, f'v{version}', 'model.pth')
        if os.path.exists(model_path):
            model.load_state_dict(torch.load(model_path, map_location='cpu'))
            print(f"成功加载模型版本 v{version}")
            return True
        else:
            print(f"模型版本 v{version} 不存在")
            return False

# 使用示例
version_manager = ModelVersionManager()

# 保存模型版本
metadata = {
    'accuracy': 0.95,
    'loss': 0.05,
    'training_time': '2 hours',
    'dataset': 'CIFAR-10'
}
version_manager.save_model(model, '1.0', metadata)

4.3 跨平台兼容性

def save_model_portable(model, filepath):
    """保存可跨平台使用的模型"""
    # 确保模型在CPU上
    model_cpu = model.cpu()
    
    # 保存时指定协议版本以确保兼容性
    torch.save(model_cpu.state_dict(), filepath, _use_new_zipfile_serialization=False)
    
    print(f"模型已保存为跨平台兼容格式: {filepath}")

def load_model_safe(model, filepath):
    """安全加载模型,处理各种异常情况"""
    try:
        # 尝试加载
        state_dict = torch.load(filepath, map_location='cpu')
        
        # 检查键的兼容性
        model_keys = set(model.state_dict().keys())
        loaded_keys = set(state_dict.keys())
        
        missing_keys = model_keys - loaded_keys
        unexpected_keys = loaded_keys - model_keys
        
        if missing_keys:
            print(f"警告: 模型中缺少以下键: {missing_keys}")
        if unexpected_keys:
            print(f"警告: 加载的状态字典中有未预期的键: {unexpected_keys}")
        
        # 加载兼容的参数
        model.load_state_dict(state_dict, strict=False)
        print("模型加载成功")
        return True
        
    except Exception as e:
        print(f"模型加载失败: {e}")
        return False

5. 常见问题与解决方案

5.1 常见错误及解决方法

常见问题

设备不匹配错误:

# 错误示例
model.load_state_dict(torch.load('model.pth'))  # 可能出现CUDA/CPU不匹配

# 正确做法
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model.load_state_dict(torch.load('model.pth', map_location=device))

模型结构不匹配:

# 部分加载,忽略不匹配的层
def load_partial_state_dict(model, state_dict):
    model_dict = model.state_dict()
    # 过滤掉不匹配的参数
    filtered_dict = {k: v for k, v in state_dict.items() 
                    if k in model_dict and v.size() == model_dict[k].size()}
    model_dict.update(filtered_dict)
    model.load_state_dict(model_dict)
    print(f"成功加载 {len(filtered_dict)} 个参数")

内存不足问题:

# 对于大模型,可以分块加载
def load_large_model(model, filepath):
    checkpoint = torch.load(filepath, map_location='cpu')
    
    # 逐层加载以节省内存
    for name, param in model.named_parameters():
        if name in checkpoint:
            param.data.copy_(checkpoint[name])
    
    print("大模型加载完成")

5.2 性能优化建议

# 使用压缩保存大模型
def save_compressed_model(model, filepath):
    """保存压缩的模型文件"""
    import gzip
    import pickle
    
    with gzip.open(filepath, 'wb') as f:
        pickle.dump(model.state_dict(), f)
    
    print(f"压缩模型已保存: {filepath}")

def load_compressed_model(model, filepath):
    """加载压缩的模型文件"""
    import gzip
    import pickle
    
    with gzip.open(filepath, 'rb') as f:
        state_dict = pickle.load(f)
    
    model.load_state_dict(state_dict)
    print(f"压缩模型已加载: {filepath}")

5.3 实际应用示例

def training_with_checkpoints():
    """带检查点的完整训练示例"""
    # 模型和优化器初始化
    model = nn.Sequential(
        nn.Linear(784, 128),
        nn.ReLU(),
        nn.Linear(128, 10)
    )
    optimizer = optim.Adam(model.parameters(), lr=0.001)
    criterion = nn.CrossEntropyLoss()
    
    # 检查点管理器
    checkpoint_manager = CheckpointManager()
    
    # 尝试加载之前的检查点
    start_epoch = 0
    best_loss = float('inf')
    
    checkpoint_files = [f for f in os.listdir('checkpoints') if f.endswith('.pth')]
    if checkpoint_files:
        latest_checkpoint = max(checkpoint_files, key=lambda x: os.path.getctime(os.path.join('checkpoints', x)))
        start_epoch, best_loss = load_checkpoint(model, optimizer, os.path.join('checkpoints', latest_checkpoint))
    
    # 训练循环
    num_epochs = 100
    for epoch in range(start_epoch, num_epochs):
        # 训练代码...
        train_loss = 0.5  # 假设的训练损失
        
        # 定期保存检查点
        if epoch % 10 == 0 or train_loss < best_loss:
            checkpoint_manager.save(model, optimizer, epoch, train_loss)
            if train_loss < best_loss:
                best_loss = train_loss
                # 保存最佳模型
                torch.save(model.state_dict(), 'best_model.pth')
        
        print(f"Epoch {epoch}, Loss: {train_loss:.4f}")

# 运行训练
if __name__ == '__main__':
    training_with_checkpoints()

参考资料

如果这篇记录对你有帮助,可以留下一句回应。

前往留言

DISCUSSION

讨论与回应

欢迎补充细节、指出问题,或分享与这篇文章有关的经验。

昵称与邮箱为必填项,邮箱仅用于头像和回复通知,不会公开。网址可以留空。

正在准备留言区…

演示赞赏界面

谢谢你愿意支持长期写作。

第一阶段不会发起支付或收集任何信息。接入真实赞赏渠道后,这里会展示清楚的金额、渠道和完成状态。

输入关键词开始搜索 · 按 Esc 关闭
打开完整搜索页 →