如何获取仅Wide模式下Wide-Deep分类器的预测概率值?
获取Wide模式逻辑回归的预测概率
嘿,我刚好熟悉TensorFlow 1.x的Estimator架构,你遇到的问题其实很容易解决——官方的LinearClassifier(也就是你用的Wide模式模型)在predict方法的返回结果里,本来就包含了0-1区间的概率值,只是你没用到对应的字段而已!
你现在打印的是pred['classes'],这是模型给出的离散分类结果,但同一个预测字典里还有probabilities字段,它是一个长度为2的数组:
- 索引0对应类别0的概率
- 索引1对应类别1的概率
直接修改你的代码就能拿到置信度了:
pred_iter = model.predict(input_fn=lambda: input_fn(FLAGS.test_data, 1, False, 1)) for pred in pred_iter: # 打印两个类别的完整概率分布 print(f"类别0概率: {pred['probabilities'][0]:.4f}, 类别1概率: {pred['probabilities'][1]:.4f}") # 如果只需要正类(1类)的置信度,直接取索引1即可 print(f"正类置信度: {pred['probabilities'][1]:.4f}")
至于你提到的旧版prob_a函数无效,那是因为TensorFlow 1.x的Estimator架构已经完全替代了旧的低级API,官方的预定义分类器(比如LinearClassifier)已经把概率输出封装在predict的结果字典里了,不需要再调用那些过时的函数。
如果你好奇为什么这个字段存在,是因为逻辑回归模型本身就是通过计算sigmoid输出得到概率的,LinearClassifier在内部已经帮你完成了这一步,直接提取就行啦!
内容的提问来源于stack exchange,提问作者Caterpillaraoz
相关产品推荐
相关产品推荐

