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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 15:05:33