You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch+NLTK聊天机器人RuntimeError:张量设备不匹配求助

解决PyTorch聊天机器人设备不匹配的RuntimeError问题

核心原因

错误根源是模型参数与输入张量不在同一计算设备上——模型可能部署在CUDA,但输入仍保留在CPU,或反之。你虽定义了device变量,但未将所有相关组件统一迁移到该设备。

具体修复步骤

1. 训练阶段统一设备(train.py)

  • 模型初始化后立即迁移到目标设备:
    model = NeuralNet(input_size, hidden_size, num_classes).to(device)
    
  • 训练循环中,将输入和标签张量同步到设备:
    for (words, labels) in train_loader:
        words = words.to(device)
        labels = labels.to(device)
        # 后续训练逻辑代码...
    

2. 推理阶段同步设备(chat.py)

  • 加载模型后,必须将模型迁移到目标设备并设置为评估模式:
    model.load_state_dict(torch.load('data.pth'))
    model = model.to(device)
    model.eval()
    
  • 处理输入文本生成的张量时,同步到目标设备:
    sentence = tokenize(sentence)
    X = bag_of_words(sentence, all_words)
    X = X.reshape(1, X.shape[0])
    X = torch.from_numpy(X).to(device)
    

3. 兼容跨环境模型加载

如果训练与推理环境的CUDA可用性不同,加载模型时指定map_location确保设备一致:

model.load_state_dict(torch.load('data.pth', map_location=device))

验证方法

  • 打印模型参数所在设备:print(next(model.parameters()).device),需与device一致
  • 打印输入张量所在设备:print(X.device),需与device一致

内容的提问来源于stack exchange,提问作者Kunal Patil

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.05 08:22:50