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

mlr3中glmnet模型未知分类水平预测异常及解决诉求

问题:mlr3 glmnet模型对新数据中未知栖息地类型的预测概率未按预期设为0

使用mlr3训练了以分类变量habitat_2(栖息地类型)为预测因子的glmnet二分类模型,用于物种存在/缺失预测。但在预测新数据时发现,新数据包含训练集未出现的栖息地类型,这些样本的预测概率却大于0.5,不符合预期。尝试用mlr3pipelines::po("imputeconstant", param_vals = list(constant = 0, affect_columns = selector_grep("habitat_2")))处理,但未生效。

失效原因

你的流水线顺序有误:imputeconstant被放在了编码步骤之后,而encodeimpact、encode等步骤会将原始的habitat_2因子列转换为多个数值编码列,此时selector_grep("habitat_2")无法匹配到原始因子列,自然无法处理未知类别。另外,fixfactors仅对齐训练与测试集的因子水平,不会处理超出训练集范围(OOR)的类别,这类未知类别会被默认编码为全0,导致模型仍输出非0概率。

解决方案

方法1:编码前处理未知类别,预测后重置概率

调整流水线顺序,在编码前用imputeoor标记未知类别,之后在预测结果中直接将这些样本的概率设为0:

# 重构预处理流水线
factor_encoding <- mlr3pipelines::po("fixfactors", id = "po_factor_alignment") %>%
  # 先将训练集没有的habitat_2类别替换为"unknown"
  mlr3pipelines::po("imputeoor", 
                    affect_columns = selector_grep("habitat_2"),
                    param_vals = list(replacement = "unknown")) %>%
  # 保留原有的编码步骤
  mlr3pipelines::po("encodeimpact", affect_columns = selector_cardinality_greater_than(10), id = "high_cardinality_encoding") %>%
  mlr3pipelines::po("encode", method = "one-hot", affect_columns = selector_cardinality_greater_than(2), id = "low_cardinality_encoding") %>%
  mlr3pipelines::po("encode", method = "treatment", affect_columns = selector_type("factor"), id = "binary_encoding")

# 后续训练代码保持不变...

# 预测后处理未知类别
predictions <- run_training$predict_newdata(newdata = new_data)
new_data$predictions <- predictions$data$prob[,c("1")]

# 获取训练集的栖息地类型水平
train_habitat_levels <- levels(data$habitat_2)
# 将新数据中不属于训练集的栖息地类型的预测概率设为0
new_data$predictions[!new_data$habitat_2 %in% train_habitat_levels] <- 0

方法2:直接在预测后修改结果(无需调整流水线)

如果不想改动训练流水线,可直接在预测完成后筛选未知类别并重置概率:

# 获取训练集的栖息地类型水平
train_habitat_levels <- levels(data$habitat_2)

# 执行预测
predictions <- run_training$predict_newdata(newdata = new_data)

# 提取概率并修改未知类别的值
pred_prob <- predictions$data$prob[,c("1")]
pred_prob[!new_data$habitat_2 %in% train_habitat_levels] <- 0
new_data$predictions <- pred_prob

# 后续可视化代码保持不变...

关键注意点

  • 确保new_data$habitat_2的类型与训练集一致(因子型),若为字符型则直接用字符串匹配即可
  • 方法1中imputeoor的replacement值需确保不在训练集的habitat_2水平中,否则无法区分未知类别

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 12:15:59