ggplot中stat_function()返回错误结果的原因排查
ggplot2中stat_function无法正确绘制多分类回归概率曲线的原因及解决方法
问题重现
数据生成与回归模型
library(tidyverse) library(nnet) # 生成数据集 -------------------------- set.seed(100) helicopter <- rnorm(20, mean = 35, sd = 3) car <- rnorm(20, mean = 30, sd = 3) bus <- rnorm(20, mean = 25, sd = 3) bike <- rnorm(20, mean = 20, sd = 3) transportation_data <- data.frame(helicopter, car, bus, bike) %>% pivot_longer(cols = 1:4, values_to = "income", names_to = "mode") # 构建多分类回归模型 ------------------- transportation_regression <- multinom(mode~income, data = transportation_data)
出错的stat_function绘图代码
尝试用stat_function直接绘制各类别概率曲线,得到错误结果并触发警告Warning: longer object length is not a multiple of shorter object length:
ins <- coef(transportation_regression)[1:3] betas <- coef(transportation_regression)[4:6] transportation_data %>% ggplot(aes(x = income))+ stat_function(fun = function(x) { 1 / (1 + sum(exp(ins + betas * x))) }, aes(color = "bike"))+ stat_function(fun = function(x) { exp(ins[1] + betas[1] * x) / (1 + sum(exp(ins + betas * x))) }, aes(color = "bus"))+ stat_function(fun = function(x) { exp(ins[2] + betas[2] * x) / (1 + sum(exp(ins + betas * x))) }, aes(color = "car"))+ stat_function(fun = function(x) { exp(ins[3] + betas[3] * x) / (1 + sum(exp(ins + betas * x))) }, aes(color = "helicopter"))
正常的手动计算绘图代码
先逐个计算每个收入值对应的概率,再绘图,结果完全正确:
income <- seq(0,50,0.1) result <- matrix( , nrow = length(income), ncol = 4) i <- 1 for(x in income){ result[i,1] <- 1 / (1 + sum(exp(ins + betas * x))) # bike result[i,2] <- exp(ins[1] + betas[1] * x) / (1 + sum(exp(ins + betas * x))) # bus result[i,3] <- exp(ins[2] + betas[2] * x) / (1 + sum(exp(ins + betas * x))) # car result[i,4] <- exp(ins[3] + betas[3] * x) / (1 + sum(exp(ins + betas * x))) # helicopter i <- i + 1 } cbind(income, as.data.frame(result)) %>% pivot_longer(cols = V1:V4) %>% ggplot(aes(x = income, y = value, color = name))+ geom_line()
原因分析
核心问题是**stat_function传入的x是向量,而你的概率计算函数没有做向量化适配**:
- 当
x是向量时,ins(长度3)和betas(长度3)与x做运算会触发R的循环补齐机制,生成一个3×N的矩阵(N是x的长度)。 - 此时
sum(exp(ins + betas * x))会把矩阵中所有3×N个元素的exp值全部求和,得到一个单一数值,而非对每个x元素单独计算对应的3个exp值之和。 - 这就导致所有
x对应的概率分母都是同一个固定值,最终绘制出的曲线完全不符合预期,同时因循环补齐的长度不匹配触发警告。
而手动循环的方式是逐个处理单个x值,sum(exp(ins + betas * x))会对每个x对应的3个线性项的exp值求和,计算逻辑正确,所以结果正常。
解决方案
方案1:改造为向量化函数
重新编写支持向量输入的概率计算函数,对每个x元素独立计算分母:
# 定义向量化的概率计算函数 calc_probs <- function(x) { # 计算每个x对应的3个线性项,生成N×3的矩阵 linear_terms <- outer(x, betas, "*") + ins exp_terms <- exp(linear_terms) # 对每行(每个x)求和,得到每个x对应的分母 denominators <- 1 + rowSums(exp_terms) # 返回每个x对应的4类概率 cbind( bike = 1 / denominators, bus = exp_terms[,1] / denominators, car = exp_terms[,2] / denominators, helicopter = exp_terms[,3] / denominators ) } # 生成预测数据并绘图 income_seq <- seq(0, 50, 0.1) prob_data <- as.data.frame(calc_probs(income_seq)) %>% mutate(income = income_seq) %>% pivot_longer(cols = -income, names_to = "mode", values_to = "prob") ggplot(prob_data, aes(x = income, y = prob, color = mode)) + geom_line()
方案2:在stat_function中逐个处理x元素
用purrr::map_dbl对向量x的每个元素单独计算,适配stat_function的输入要求:
library(purrr) transportation_data %>% ggplot(aes(x = income))+ stat_function(fun = function(x) map_dbl(x, ~1/(1+sum(exp(ins + betas*.x)))), aes(color = "bike"))+ stat_function(fun = function(x) map_dbl(x, ~exp(ins[1]+betas[1]*.x)/(1+sum(exp(ins + betas*.x)))), aes(color = "bus"))+ stat_function(fun = function(x) map_dbl(x, ~exp(ins[2]+betas[2]*.x)/(1+sum(exp(ins + betas*.x)))), aes(color = "car"))+ stat_function(fun = function(x) map_dbl(x, ~exp(ins[3]+betas[3]*.x)/(1+sum(exp(ins + betas*.x)))), aes(color = "helicopter"))+ labs(color = "交通方式")
方案3:推荐用predict函数直接生成预测概率
更简洁的方式是用模型自带的predict函数生成概率,避免手动计算系数:
# 生成收入序列 income_seq <- tibble(income = seq(0, 50, 0.1)) # 预测各类别概率 prob_pred <- predict(transportation_regression, newdata = income_seq, type = "probs") %>% as.data.frame() %>% bind_cols(income_seq) %>% pivot_longer(cols = -income, names_to = "mode", values_to = "prob") # 绘图 ggplot(prob_pred, aes(x = income, y = prob, color = mode)) + geom_line()
内容的提问来源于stack exchange,提问作者Sebastian Geis
相关产品推荐
相关产品推荐

