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

在Keras神经网络中计算各epoch的Top k Precision@n与Recall@n

在Keras中计算每个Epoch的Top k Precision@n和Recall@n

嘿,我懂你想要在Keras训练时,每个epoch都能追踪Top k的Precision和Recall指标的需求——这在多分类任务里简直太实用了!你之前找到的基础precision/recall函数是针对二分类或全局预测场景的,咱们得把它改成适配Top k逻辑的版本,直接看代码和说明吧:

自定义Top k指标函数

以下函数用Keras后端(K)实现,专门适配多分类场景(要求y_true是one-hot编码格式):

Precision@k

import keras.backend as K

def precision_at_k(y_true, y_pred, k=5):
    # 对每个样本,取出预测概率最高的k个类别索引
    top_k_indices = K.argsort(y_pred)[:, -k:]
    # 提取真实标签中这些索引对应的正例数量
    true_positives = K.sum(K.gather(y_true, top_k_indices, axis=1))
    # 预测的正例数固定为k(因为我们取了top k个候选)
    predicted_positives = k
    # 计算精度,加epsilon避免除零错误
    precision = true_positives / (predicted_positives + K.epsilon())
    return precision

Recall@k

def recall_at_k(y_true, y_pred, k=5):
    # 计算每个样本的真实正例总数(多分类中每个样本一般仅1个正例)
    total_positives = K.sum(y_true, axis=1)
    # 取出预测概率最高的k个类别索引
    top_k_indices = K.argsort(y_pred)[:, -k:]
    # 提取真实标签中这些索引对应的正例数量
    true_positives = K.sum(K.gather(y_true, top_k_indices, axis=1))
    # 计算召回率:对真实正例数为0的样本直接跳过(赋值0),避免无效计算
    recall = K.mean(K.switch(total_positives == 0, K.constant(0.0), true_positives / (total_positives + K.epsilon())))
    return recall

如何使用这些指标

在编译模型时,把自定义指标加入metrics参数即可,还能根据你的需求灵活调整k值:

model.compile(
    optimizer='adam',
    loss='categorical_crossentropy',  # 多分类任务常用损失函数
    metrics=[
        lambda y_true, y_pred: precision_at_k(y_true, y_pred, k=3),  # 自定义Precision@3
        lambda y_true, y_pred: recall_at_k(y_true, y_pred, k=3)     # 自定义Recall@3
    ]
)

这样训练时,每个epoch结束后就会输出你指定的Top k Precision和Recall啦!如果是多标签分类场景,只需要微调真实正例数的计算逻辑就能适配~

内容的提问来源于stack exchange,提问作者Muhammad shahrukh khan

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.22 09:24:24