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
相关产品推荐
相关产品推荐

