如何复现使用caret构建的regLogistic模型(更新版)
问题分析与解决:caret中regLogistic模型标准化后手动计算预测概率不一致
我来帮你拆解下这个问题,你遇到的预测结果不一致,其实是两个关键的小错误导致的,咱们一步步说清楚:
为什么会出现不一致?
- 标准化参数搞反了:你代码里把
fit.log$preProcess$std(训练集的标准差)赋值给了mean_train,把fit.log$preProcess$mean(训练集的均值)赋值给了sd_train,这完全颠倒了标准化的计算逻辑!正确的标准化应该是(测试集特征 - 训练集均值)/训练集标准差,你这里用反了,特征缩放完全错误。 - Sigmoid公式符号错误:逻辑回归的预测概率公式是
p = 1/(1 + exp(-(截距 + β₁x₁ + β₂x₂ + ... + βₙxₙ))),但你手动计算时把公式里的负号漏掉了,直接用了exp((截距 + ...)),这导致最终的概率计算完全偏离。
当你关闭preProc参数时,模型用原始特征训练,你手动计算时的公式符号和系数匹配,且没有标准化错误,所以结果几乎一致;但开启标准化后,两个错误叠加,就出现了明显的结果差异。
如何正确复现预测结果?
只需要修正上面两个错误,就能得到和predict()完全一致的结果。下面是修正后的完整代码:
library(data.table) library(caret) library(LiblineaR) DT <- data.table(iris) DT$Species <- ifelse(DT$Species == "versicolor", 1,0) DT$Species <- factor(DT$Species, levels = c("1", "0"), labels = c("Yes", "No")) set.seed(42) trainIndex <- createDataPartition(DT$Species, p = .7, list = F) train <- DT[trainIndex[, 1], ] test <- DT[-trainIndex[, 1], ] # 带标准化的模型训练 set.seed(42) fit.log <- train(form = Species ~ ., data = train, method = "regLogistic", metric = "Accuracy", preProc = c("center", "scale")) # 模型预测概率 pred <- predict(fit.log, test, type = "prob") print("模型预测的Yes概率:") print(pred[,"Yes", drop = FALSE]) # 提取正确的均值和标准差 mean_train <- fit.log$preProcess$mean # 这里是均值,之前搞反了 sd_train <- fit.log$preProcess$std # 这里是标准差,之前搞反了 # 对测试集做正确的标准化 test[, Sepal.Length_STD:=(Sepal.Length - mean_train["Sepal.Length"])/sd_train["Sepal.Length"]] test[, Sepal.Width_STD:=(Sepal.Width - mean_train["Sepal.Width"])/sd_train["Sepal.Width"]] test[, Petal.Length_STD:=(Petal.Length - mean_train["Petal.Length"])/sd_train["Petal.Length"]] test[, Petal.Width_STD:=(Petal.Width - mean_train["Petal.Width"])/sd_train["Petal.Width"]] # 提取模型系数和截距 coefficients <- fit.log$finalModel$W[,1:4] # 特征系数 intercept <- fit.log$finalModel$W[,5] # 截距项(Bias) # 用正确的Sigmoid公式计算概率 test[, target:=1/(1+exp(-(intercept + Sepal.Length_STD * coefficients[1] + Sepal.Width_STD * coefficients[2] + Petal.Length_STD * coefficients[3] + Petal.Width_STD * coefficients[4])))] print("手动计算的Yes概率:") print(test[,.(target)]) # 验证一致性(差值几乎为0) print("两者差值:") print(pred[,"Yes"] - test$target)
运行这段代码后,你会发现模型预测的概率和手动计算的概率完全一致(差值几乎为0,是浮点数精度问题导致的微小差异)。
内容的提问来源于stack exchange,提问作者Anastasiia Dovgosheya
相关产品推荐
相关产品推荐

