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

Keras二分类模型最优AUC阈值获取及分类预测问询

Keras二分类模型:预测类别与最优阈值确定

我用Keras训练了一个二分类模型,模型定义如下:

model_binary = Sequential()
model_binary.add(layers.Dense(64, activation='relu',input_shape=(7,)))
model_binary.add(layers.Dropout(0.5))
model_binary.add(layers.Dense(32, activation='relu'))
model_binary.add(layers.Dropout(0.5))
model_binary.add(layers.Dense(16, activation='relu'))
model_binary.add(layers.Dropout(0.5))
model_binary.add(layers.Dense(1, activation='sigmoid')) 

训练代码及输出:

model_binary.compile(optimizer=opt, loss='binary_crossentropy', metrics=[tf.keras.metrics.AUC(name='auc')])
model_binary.fit(binary_train_data, binary_train_labels, batch_size=16, epochs=10, validation_split=0.1)

训练日志:

Epoch 1/10
507/507 [==============================] - 3s 4ms/step - loss: 0.4017 - auc: 0.5965 - val_loss: 0.2997 - val_auc: 0.8977
Epoch 2/10
507/507 [==============================] - 1s 2ms/step - loss: 0.3354 - auc: 0.7387 - val_loss: 0.2729 - val_auc: 0.9019
Epoch 3/10
507/507 [==============================] - 1s 3ms/step - loss: 0.3167 - auc: 0.7837 - val_loss: 0.2623 - val_auc: 0.9021
Epoch 4/10
507/507 [==============================] - 1s 2ms/step - loss: 0.3072 - auc: 0.8057 - val_loss: 0.2551 - val_auc: 0.9003
Epoch 5/10
507/507 [==============================] - 1s 2ms/step - loss: 0.2948 - auc: 0.8298 - val_loss: 0.2507 - val_auc: 0.9033
Epoch 6/10
507/507 [==============================] - 1s 2ms/step - loss: 0.2921 - auc: 0.8355 - val_loss: 0.2489 - val_auc: 0.9005
Epoch 7/10
507/507 [==============================] - 2s 4ms/step - loss: 0.2867 - auc: 0.8431 - val_loss: 0.2465 - val_auc: 0.9016
Epoch 8/10
507/507 [==============================] - 2s 4ms/step - loss: 0.2865 - auc: 0.8434 - val_loss: 0.2460 - val_auc: 0.9017
Epoch 9/10
507/507 [==============================] - 2s 4ms/step - loss: 0.2813 - auc: 0.8493 - val_loss: 0.2452 - val_auc: 0.9030
Epoch 10/10
507/507 [==============================] - 1s 3ms/step - loss: 0.2773 - auc: 0.8560 - val_loss: 0.2441 - val_auc: 0.9029

数据集存在类别倾斜,正样本占87%,负样本占13%。从val_auc来看模型表现尚可,但对训练数据预测时,最低输出得分约0.6,平衡数据集通常以0.5作为sigmoid分类阈值。现咨询两个问题:

  1. 给定数据x,如何获取模型对其的预测类别?
  2. 如何得到模型的最优分类阈值?

补充说明:train_labels为形状为N的0、1值ndarray。


一、获取单条/批量数据的预测类别

模型的sigmoid输出是样本属于正类(类别1)的概率,要得到预测类别,只需用分类阈值对概率进行判断:概率≥阈值则判定为1,否则为0。

代码实现

# 1. 得到模型预测概率(形状为(N,1)的数组)
pred_probs = model_binary.predict(x, verbose=0)
# 2. 用阈值转换为类别(这里先假设threshold为后续确定的最优值)
pred_classes = (pred_probs >= threshold).astype(int)
# 单条数据直接取标量结果
if len(pred_classes) == 1:
    pred_class = pred_classes[0][0]

注意:推理时Dropout会自动关闭,若需显式确认,可设置:

model_binary.trainable = False
pred_probs = model_binary.predict(x, verbose=0)

二、确定最优分类阈值

由于数据集类别严重倾斜,默认的0.5阈值会导致模型过度倾向于预测正类,无法有效识别负样本。最优阈值需要结合验证集的预测结果,基于业务关注的指标确定,以下是三种实用方法:

方法1:基于ROC曲线找最优阈值

ROC曲线的每个点对应一个阈值,选择使**Youden指数(灵敏度+特异度-1)**最大的阈值,该阈值兼顾正类和负类的识别能力。

代码实现

from sklearn.metrics import roc_curve

# 1. 获取验证集数据与预测概率
val_data, val_labels = model_binary.validation_data
val_probs = model_binary.predict(val_data, verbose=0).flatten()

# 2. 计算ROC曲线相关参数
fpr, tpr, thresholds = roc_curve(val_labels, val_probs)

# 3. 找到Youden指数最大的阈值
youden_index = tpr - fpr
best_idx = youden_index.argmax()
best_threshold_roc = thresholds[best_idx]
print(f"基于ROC曲线的最优阈值: {best_threshold_roc:.4f}")

方法2:基于精确率-召回率曲线找最优阈值

对于不平衡数据集,精确率-召回率(PR)曲线比ROC曲线更能反映模型性能。选择使F1-score最大的阈值,F1是精确率和召回率的调和平均,适合平衡两类样本的识别效果。

代码实现

from sklearn.metrics import precision_recall_curve

# 1. 计算精确率、召回率与对应阈值
precision, recall, thresholds = precision_recall_curve(val_labels, val_probs)

# 2. 计算每个阈值对应的F1-score(截断适配阈值长度)
f1_scores = 2 * (precision[:-1] * recall[:-1]) / (precision[:-1] + recall[:-1])

# 3. 找到F1-score最大的阈值
best_idx = f1_scores.argmax()
best_threshold_f1 = thresholds[best_idx]
print(f"基于F1-score的最优阈值: {best_threshold_f1:.4f}")
print(f"对应最大F1-score: {f1_scores[best_idx]:.4f}")

方法3:基于业务需求自定义阈值

如果业务对某类样本的识别有硬性要求(比如要求负样本召回率≥90%),可直接从PR曲线中找到满足该要求的最小阈值。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.15 00:55:20