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
相关产品推荐
相关产品推荐

