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

R中使用class_id调用metric_recall_at_precision编译模型报错求助

解决R Keras中metric_recall_at_precision/metric_precision_at_recall的class_id参数报错问题

在Windows 10 GPU环境下,使用R 3.6.3、keras 2.9.0、tensorflow 2.9.0(reticulate绑定Python 3.6.10)编译三分类模型时,调用metric_recall_at_precision或metric_precision_at_recall并传入class_id参数会触发如下报错:

Error in py_call_impl(callable, dots$args, dots$keywords) : 
  TypeError: __init__() got an unexpected keyword argument 'class_id'

Keras官方文档标注class_id是可选参数,但R的keras包装接口并未暴露该参数,导致直接调用R封装函数时出错。使用metric_sparse_categorical_accuracy或转为二分类(sigmoid输出)时模型可正常编译,报错的简化代码如下:

model <- keras_model_sequential() %>% 
     layer_conv_1d(filters = 64, kernel_size = 11, strides = 5, activation = "relu", input_shape = c(446,3)) %>% 
     layer_max_pooling_1d(pool_size = 5) 

model %>% 
    layer_dropout(rate = 0.1) %>%
    layer_flatten() %>% 
    layer_dense(units = 64, activation = "relu") %>%
    layer_dense(units = 3, activation = "softmax")  

model %>% compile(
    optimizer = "adam",
    loss = "sparse_categorical_crossentropy",   
    metrics =  metric_recall_at_precision(precision=precision, class_id=0))

方案1:通过reticulate调用Python原生指标类

R的keras包对部分TensorFlow/Keras指标的参数封装不全,直接用reticulate调用Python原生的RecallAtPrecision类,就能正常传入class_id参数:

# 获取Python的tensorflow.keras.metrics模块
tf_metrics <- import("tensorflow.keras.metrics")

# 初始化指定class_id的RecallAtPrecision指标(替换precision为你的目标值)
recall_at_precision <- tf_metrics$RecallAtPrecision(precision = 0.8, class_id = 0)

# 编译模型时使用该指标
model %>% compile(
    optimizer = "adam",
    loss = "sparse_categorical_crossentropy",   
    metrics = list(recall_at_precision)
)

方案2:自定义指标函数

如果不想依赖Python原生调用,可以手动编写自定义指标函数,针对指定类别计算召回率@精度:

# 自定义针对class_id=0的召回率@精度指标
custom_recall_at_precision <- function(y_true, y_pred) {
    # 转换为二分类任务的真实标签与预测概率
    y_true_class <- tf$cast(y_true == 0, tf$float32)
    y_pred_class <- y_pred[, 0]
    
    # 基于目标精度计算对应阈值,再计算召回率
    precision_obj <- tf$keras$metrics$Precision(class_id = 0)
    precision_obj$update_state(y_true, y_pred)
    target_precision <- precision_obj$result()
    
    # 获取满足精度要求的阈值,计算对应召回率
    pred_probs <- tf$boolean_mask(y_pred_class, y_true_class == 1)
    threshold <- tf$math$reduce_min(pred_probs)
    y_pred_thresholded <- tf$cast(y_pred_class >= threshold, tf$float32)
    recall <- tf$keras$metrics$recall(y_true_class, y_pred_thresholded)
    
    return(recall)
}

# 编译模型时使用自定义指标
model %>% compile(
    optimizer = "adam",
    loss = "sparse_categorical_crossentropy",   
    metrics = list(custom_recall_at_precision)
)

注意事项

  • 方案1是最可靠的解决方式,完全复用TensorFlow原生指标的实现,避免自定义函数的逻辑误差。
  • 确保reticulate已正确绑定目标Python环境(安装有tensorflow 2.9.0),可通过reticulate::py_config()检查配置。
  • 使用sparse_categorical_crossentropy损失时,class_id对应稀疏标签的0-based索引,需与输出层softmax的维度对应。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.21 16:18:18