PyTorch CUDA设备不可用错误解决方案

PyTorch CUDA设备不可用错误解决方案

问题描述

在使用PyTorch加载模型时遇到以下错误:

复制代码
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 map_location=torch.device('cpu') to map your storages to the CPU.
已中止 (核心已转储)

解决方案

方案一:强制加载到CPU(快速解决)

python 复制代码
import torch

# 方法1:直接指定CPU设备
model = torch.load('model.pth', map_location=torch.device('cpu'))

# 方法2:使用字符串指定
model = torch.load('model.pth', map_location='cpu')

# 方法3:使用lambda函数
model = torch.load('model.pth', map_location=lambda storage, loc: storage.cpu())

方案二:检查并修复GPU环境(根本解决)

步骤1:检查显卡驱动
bash 复制代码
nvidia-smi
步骤2:验证PyTorch安装
python 复制代码
import torch
print(f"PyTorch版本: {torch.__version__}")
print(f"CUDA可用: {torch.cuda.is_available()}")
print(f"CUDA版本: {torch.version.cuda}")
步骤3:重新安装正确的PyTorch版本
bash 复制代码
# 安装支持CUDA的PyTorch
pip3 install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118

# 或使用conda
conda install pytorch torchvision torchaudio pytorch-cuda=11.8 -c pytorch -c nvidia

方案三:使用state_dict方式(最佳实践)

python 复制代码
# 保存模型时使用state_dict
torch.save(model.state_dict(), 'model_weights.pth')

# 加载模型时指定设备
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model.load_state_dict(torch.load('model_weights.pth', map_location=device))

特殊情况处理

AMD显卡用户

python 复制代码
# AMD显卡不支持CUDA,使用CPU
device = torch.device('cpu')

Apple M系列芯片

python 复制代码
# Apple M系列芯片使用MPS加速
if torch.backends.mps.is_available():
    device = torch.device('mps')
else:
    device = torch.device('cpu')

完整解决方案代码

python 复制代码
import torch
import os

def load_model_safely(model_path, model_class, **kwargs):
    """
    安全加载模型,自动处理设备兼容性问题
    """
    # 检测可用设备
    if torch.cuda.is_available():
        device = torch.device('cuda')
        print("✅ 检测到CUDA设备,使用GPU加速")
    elif hasattr(torch.backends, 'mps') and torch.backends.mps.is_available():
        device = torch.device('mps')
        print("✅ 检测到Apple M系列芯片,使用MPS加速")
    else:
        device = torch.device('cpu')
        print("💻 未检测到GPU设备,使用CPU")
    
    # 创建模型实例
    model = model_class(**kwargs)
    
    try:
        # 尝试加载state_dict
        state_dict = torch.load(model_path, map_location=device)
        
        # 处理checkpoint格式
        if isinstance(state_dict, dict) and 'model_state_dict' in state_dict:
            state_dict = state_dict['model_state_dict']
        
        # 处理多GPU训练的模型
        if list(state_dict.keys())[0].startswith('module.'):
            state_dict = {k.replace('module.', ''): v for k, v in state_dict.items()}
        
        model.load_state_dict(state_dict)
        model = model.to(device)
        model.eval()
        
        print(f"✅ 模型成功加载到 {device}")
        return model, device
        
    except Exception as e:
        print(f"❌ 模型加载失败: {str(e)}")
        raise

# 使用示例
if __name__ == "__main__":
    # 假设你有一个模型类MyModel
    # model, device = load_model_safely('model.pth', MyModel, num_classes=10)
    pass

预防措施

  1. 保存模型时使用state_dict:
python 复制代码
torch.save(model.state_dict(), 'model_weights.pth')
  1. 加载模型时指定设备:
python 复制代码
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model.load_state_dict(torch.load('model_weights.pth', map_location=device))
  1. 定期检查环境配置:
python 复制代码
print(f"PyTorch版本: {torch.__version__}")
print(f"CUDA可用: {torch.cuda.is_available()}")
print(f"GPU数量: {torch.cuda.device_count()}")
if torch.cuda.is_available():
    print(f"GPU名称: {torch.cuda.get_device_name(0)}")

常见问题排查

问题1:安装了GPU版PyTorch但CUDA不可用

  • 检查显卡驱动是否最新
  • 确认PyTorch版本与CUDA版本匹配
  • 重启Python环境

问题2:版本不兼容

  • 卸载当前PyTorch:pip uninstall torch torchvision torchaudio
  • 重新安装匹配版本:参考PyTorch官网安装命令

问题3:多环境冲突

  • 确认当前使用的Python环境
  • 检查虚拟环境中的PyTorch安装
  • 使用which python和which pip确认路径

注意:在生产环境中,建议始终使用state_dict方式保存和加载模型,以获得最佳的兼容性和灵活性。

相关推荐
回眸&啤酒鸭3 天前
【回眸】Minicart 电商购物车核心功能落地指南
人工智能
一隅论数智3 天前
给AI一张“业务概念地图“:本体如何从哲学走向企业智能
大数据·人工智能·经验分享·笔记·学习·学习方法·政务
默_笙3 天前
🍙 给每个请求过安检:FastAPI 是怎么把校验写进类型注解的
python
AI的探索之旅3 天前
97 个 OpenCV 实例(三十):双目立体,从标定到点云
人工智能·opencv·计算机视觉
AlbertZein3 天前
Step-5-Preview 上手实测:3D 游戏、金融分析、网页设计一次跑完
人工智能·aigc
LaughingZhu3 天前
Product Hunt 每日热榜 | 2026-09-19
人工智能·深度学习·神经网络·搜索引擎·百度
qq_426003963 天前
启动playwright录制codegen生成自动化测试脚本
python·自动化
虎头金猫3 天前
4K 视频总卡在公网带宽?用 N1 + OpenList 把网盘播放链路重新理顺
运维·服务器·网络·python·容器·beautifulsoup·pandas
美狐美颜SDK开放平台3 天前
开发直播APP时如何接入视频美颜SDK?开发流程与注意事项
android·人工智能·计算机视觉·音视频·直播美颜sdk