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

PyTorch Lightning预测输出张量含义解析及查询方法咨询

PyTorch Lightning预测张量输出含义解析

问题描述

我无法理解PyTorch Lightning中trainer.predict()的输出含义,相关代码及输出如下:

预测代码

# Predicting
path = analysis.best_checkpoint + '/' + "ray_ckpt"

model = GraphLevelGNN.load_from_checkpoint(path)
model.eval() 

trainer = pl.Trainer()
test_result = trainer.test(model, graph_test_loader, verbose=False)

print(test_result)
## [{'test_acc': 0.65625, 'test_f1': 0.7678904428904428, 'test_precision': 1.0, 'test_recall': 0.65625}]

predictions = trainer.predict(model, graph_test_loader)

print(predictions)

预测输出

[(tensor(0.7582), tensor(0.5000), 0.6666666666666666, 1.0, 0.5), (tensor(0.4276), tensor(0.7500), 0.8571428571428571, 1.0, 0.75), (tensor(0.4436), tensor(0.7500), 0.8571428571428571, 1.0, 0.75), (tensor(0.2545), tensor(1.), 1.0, 1.0, 1.0), (tensor(1.0004), tensor(0.3750), 0.5454545454545454, 1.0, 0.375)]

解析方案

  • 核心依据:trainer.predict()的输出完全由你的GraphLevelGNN模型类中predict_step()方法的返回值决定,PyTorch Lightning不会自动生成预测输出,完全遵循你定义的返回内容。
  • 输出结构推测:从输出的五元组和test_result中的全局指标对应来看,大概率是单样本维度的返回值,比如可能是:模型输出的logits/预测概率、真实标签、该样本对应的F1值、精确率、召回率。具体对应关系必须看predict_step()的代码实现。
  • 验证方法:直接查看GraphLevelGNN类的predict_step方法,看它return的变量顺序和含义。比如如果方法里写的是return (logits, true_label, sample_f1, sample_precision, sample_recall),就能直接对应上输出的每个元素。
  • 优化建议:如果是你自己编写的模型,建议在predict_step中用字典形式返回结果(比如return {"logits": logits, "true_label": true_label, "f1": sample_f1}),这样输出的可读性会大幅提升,不需要额外猜测含义。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.21 15:15:34