PyTorch Geometric多任务模型从test_loader提取全量测试集真值与预测值
PyTorch Geometric多回归任务测试集数据采集方案
1 提取完整测试集所有真值
你只需要遍历test_loader的所有批次,将单批次的两个真值tensor拼接后汇总即可,具体实现代码如下:
import torch # 初始化空列表存储所有批次的真值 all_y = [] for data in test_loader: # 将长度为2的真值列表按第二维堆叠,得到形状为[128, 2]的单批次真值 batch_y = torch.stack(data.y, dim=1) all_y.append(batch_y) # 沿样本维度拼接所有批次结果,得到形状为[测试集总样本数, 2]的全量真值 full_test_y = torch.cat(all_y, dim=0)
2 获取全量测试集预测结果
推理前先将模型切换为评估模式、关闭梯度计算,再遍历测试集所有批次执行推理,汇总预测结果即可,代码如下:
model.eval() all_pred = [] # 关闭梯度计算,降低显存占用、提升推理速度 with torch.no_grad(): for data in test_loader: batch_pred = model(data) all_pred.append(batch_pred) # 沿样本维度拼接所有批次预测结果 full_test_pred = torch.cat(all_pred, dim=0)
拿到full_test_y和full_test_pred后,就可以直接调用各类指标计算工具完成模型评估:
- 若使用PyTorch生态的指标工具,可直接传入两个tensor计算
- 若使用sklearn等CPU端的指标工具,先调用
.cpu().numpy()方法将tensor转为numpy数组即可
内容的提问来源于stack exchange,提问作者James Arten
相关产品推荐
相关产品推荐

