如何基于两类训练数据设计Keras网络实现三类(含其他类)预测
用猫狗二分类数据实现三类(猫/狗/非猫狗)预测的方案
嘿,这个问题挺关键的——咱们先戳破核心点:你手里只有猫和狗的训练数据,现有ConvNet是纯二分类模型,没法直接通过修改compile()参数或者调整训练超参数让它原生输出三类概率。原因很简单:模型从来没见过“非猫狗”的样本,根本学不到这类样本的特征,强行改输出层或者损失函数只会让模型乱猜第三类。
不过咱们可以换个思路,利用现有二分类模型的置信度输出,通过阈值判断来划分出“非猫狗”类,具体方案如下:
1. 先搞懂现有模型的输出逻辑
假设你的模型训练时用了binary_crossentropy损失,输出层是sigmoid激活,那predict(x)输出的是样本属于正类(比如你定义的“猫”)的概率值:
- 输出接近1 → 模型高度确信是猫
- 输出接近0 → 模型高度确信是狗
- 输出在中间区间(比如0.3~0.7)→ 模型对这个样本的分类非常不确定,这部分样本就可以归为“非猫狗”
2. 用置信度阈值实现三类划分
不需要改动模型的训练配置(compile()或者超参数都不用改),只需要在预测阶段加一层逻辑判断,设定两个阈值来区分三类:
- 举个具体的规则:
- 如果
猫的概率 >= 0.8→ 判定为猫 - 如果
猫的概率 <= 0.2→ 判定为狗(因为狗的概率就是1 - 猫的概率) - 如果概率在0.2~0.8之间 → 判定为非猫狗
- 如果
阈值怎么选?
你可以用手里的猫狗验证集来调优:
- 让模型跑一遍验证集,统计不同阈值下的误判率:比如如果想尽量少把猫狗误判成非猫狗,就把阈值设得更极端(比如0.85和0.15);如果想尽量多识别出非猫狗,就把阈值设得宽松些(比如0.6和0.4)
给你个简单的实现代码:
def predict_three_classes(model, input_image, upper_thresh=0.8, lower_thresh=0.2): # 获取模型输出的猫类概率 cat_prob = model.predict(input_image, verbose=0)[0][0] if cat_prob >= upper_thresh: return ("猫", cat_prob) elif cat_prob <= lower_thresh: return ("狗", 1 - cat_prob) else: # 模型对猫狗的置信度都很低,归为非猫狗 return ("非猫狗", max(cat_prob, 1 - cat_prob))
3. 进阶:用分布外(OOD)检测提升准确率
如果想更精准地识别非猫狗,可以用OOD检测方法,比如计算模型预测的熵——熵越高,说明模型对样本的分类越不确定,大概率是非猫狗。
示例代码(计算二分类的预测熵):
import numpy as np def calculate_prediction_entropy(prob): # 二分类熵公式:-p*log(p) - (1-p)*log(1-p) return -prob * np.log(prob) - (1 - prob) * np.log(1 - prob) def predict_with_ood(model, input_image, entropy_thresh=0.6): cat_prob = model.predict(input_image, verbose=0)[0][0] entropy = calculate_prediction_entropy(cat_prob) if entropy > entropy_thresh: return ("非猫狗", entropy) elif cat_prob >= 0.5: return ("猫", cat_prob) else: return ("狗", 1 - cat_prob)
为什么不能直接改compile或超参数?
再强调一遍:你的训练数据里没有“非猫狗”样本,模型没有任何学习这类样本的依据。就算你把输出层改成3个神经元+softmax,把compile()的损失改成categorical_crossentropy,模型也只会把所有样本强行分到猫或狗里——因为没有第三类的标签信号,模型根本不知道“非猫狗”是什么。
内容的提问来源于stack exchange,提问作者pdubois
相关产品推荐
相关产品推荐

