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

如何在Keras中自定义与F1逻辑类似的AUC评估函数?

问题根源

你的auc_metric无法正常运行的核心原因是调用了scikit-learn的roc_curve和auc函数,这两个函数基于NumPy数组运算,属于普通Python函数,无法接入Keras的静态计算图。自定义Keras指标的所有计算逻辑必须使用Keras后端(或对应框架如TensorFlow)的张量运算接口,和你实现的f1_metric全部采用K.xxx接口的逻辑保持一致。

正确的自定义AUC指标实现

以下给出和f1_metric风格统一、完全基于Keras后端运算的实现:

from keras import backend as K

def auc_metric(y_true, y_pred):
    # 对应你原来取第二列的逻辑,适配二分类输出为[负类概率,正类概率]的情况
    if y_pred.shape[-1] == 2:
        y_pred = y_pred[:, 1]
    # 排序预测值获取阈值对应索引
    sorted_indices = K.tf.argsort(y_pred, axis=0, direction='DESCENDING')
    y_true_sorted = K.gather(y_true, sorted_indices)
    # 累积计算真阳性、假阳性样本数
    tps = K.cumsum(y_true_sorted, axis=0)
    fps = K.cumsum(1 - y_true_sorted, axis=0)
    # 统计总正负样本数
    total_pos = K.sum(y_true)
    total_neg = K.sum(1 - y_true)
    # 样本分布异常时返回默认值避免除0
    if total_pos == 0 or total_neg == 0:
        return K.constant(0.0)
    # 归一化得到真阳性率、假阳性率
    tpr = tps / total_pos
    fpr = fps / total_neg
    # 梯形法计算ROC曲线下面积
    auc_val = K.tf.trapz(tpr, fpr)
    return auc_val

如果你的环境是TensorFlow 2.x+Keras,也可以直接复用TensorFlow封装好的AUC计算逻辑,稳定性和效率更高:

import tensorflow as tf

def auc_metric(y_true, y_pred):
    if y_pred.shape[-1] == 2:
        y_pred = y_pred[:, 1]
    auc_calculator = tf.keras.metrics.AUC()
    auc_calculator.update_state(y_true, y_pred)
    return auc_calculator.result()

训练代码修改

你需要在model.compile的metrics参数里加入自定义的指标,修改后的编译代码如下:

opt = SGD(lr=0.01,momentum=0.9) 
model.compile(loss='binary_crossentropy', optimizer=opt,metrics=['accuracy', f1_metric, auc_metric])
ca = SnapshotEnsemble(n_epochs, n_cycles, 0.01)

# fit model
history=model.fit(trainX, trainy, validation_data=(testX, testy), epochs=n_epochs, 
verbose='auto', callbacks=[ca],batch_size=32)

注意事项

  • 逐batch计算的AUC值和全局AUC值会有一定偏差,如果需要更准确的全局验证集AUC,可以通过回调函数在每个epoch结束后手动计算整个验证集的AUC输出。
  • 如果你使用的是二分类单输出(只输出正类概率),可以去掉代码中判断y_pred维度取第二列的逻辑。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 21:48:05