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

PyTorch GPU环境下调用Scikit-Learn r2_score报错如何解决?

问题解决指南

1. 修复Scikit-Learn设备不兼容错误

Scikit-Learn的指标函数仅支持NumPy数组,无法直接处理PyTorch张量。修改test函数中的评分代码,将张量转换为CPU上的NumPy数组:

print(r2_score(y.cpu().numpy(), y_pred.cpu().numpy()))

.cpu()确保张量处于CPU设备(即使模型跑在GPU上也能兼容),.numpy()完成张量到NumPy数组的转换。

2. 消除Torch模型加载的安全警告

按照警告提示,在加载模型时添加weights_only=True参数,这是PyTorch未来的默认安全设置:

model.load_state_dict(torch.load("model.h5", weights_only=True))

3. 真正启用GPU训练(可选)

你安装了CUDA版本的PyTorch,但当前代码全程在CPU运行,并未利用GPU。要启用GPU加速,需修改两处:

  • 初始化模型时移至GPU:
    model = MyMachine().to("cuda")
    
  • 生成数据集后将数据移至GPU:
    X, y = get_dataset()
    X = X.to("cuda")
    y = y.to("cuda")
    

测试阶段也要保持设备一致,将测试用的模型和数据同样移至GPU。

修改后的完整test函数示例

def test():
    model = MyMachine().to("cuda")  # 匹配训练时的设备
    model.load_state_dict(torch.load("model.h5", weights_only=True))
    model.eval()
    X, y = get_dataset()
    X = X.to("cuda")
    y = y.to("cuda")

    with torch.no_grad():
        y_pred = model(X)
        print(r2_score(y.cpu().numpy(), y_pred.cpu().numpy()))

内容的提问来源于stack exchange,提问作者infcs.ltd

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 15:36:04