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

如何用ggplot重绘LASSO回归图?解决绘图不一致及次X轴问题

LASSO回归ggplot重绘图与原生plot不一致的问题及次X轴添加方法

问题描述

构建LASSO回归后,用plot(fit)可绘制标准LASSO图,但用ggplot重绘后出现X轴区间不同、线条不完全匹配的问题;同时希望在ggplot图中添加显示变量数量的次X轴,实现与原生plot一致的效果。

用户代码如下:

原生绘图代码

library(glmnet)
library(dplyr)
library(ggplot2)
x <- matrix(rnorm(100 * 20), 100, 20)
y <- sample(1:2, 100, replace = TRUE)
fit <- glmnet(x, y, family = "binomial")

# 原生绘图
plot(fit, xvar = "lambda", label = T)

ggplot重绘代码

tidied <- broom::tidy(fit) %>% filter(term!= "(Intercept)")

# ggplot重绘
ggplot(tidied, aes(lambda, estimate, group = term, color = term)) +
  geom_line() +
  scale_x_log10()

问题原因

  1. X轴区间差异:glmnet原生plot()默认只展示系数发生变化的lambda区间,自动截断无意义的端点;而broom::tidy()会返回所有lambda值(包括系数未变化的端点),导致ggplot的X轴范围更宽。
  2. 线条匹配问题:原生plot仅保留系数发生变化的lambda节点并连线,而ggplot直接用所有lambda值连线,同时tidy数据包含所有系数(包括始终为0的项),导致线条显示与原生图不一致。

解决方法

1. 修正ggplot图与原生plot对齐

从fit对象中提取和原生plot逻辑一致的有效数据,过滤系数未变化的冗余点:

# 从fit中提取关键数据:lambda、非截距项系数
lasso_data <- tibble(
  lambda = fit$lambda,
  # 转换系数矩阵为长格式
  t(coef(fit)) %>% as_tibble() %>% select(-`(Intercept)`)
) %>%
  pivot_longer(cols = -lambda, names_to = "term", values_to = "estimate") %>%
  # 仅保留系数发生变化的lambda点
  filter(estimate!= lag(estimate, default = 0) | estimate!= lead(estimate, default = 0))

# 绘制对齐后的基础ggplot图
p <- ggplot(lasso_data, aes(lambda, estimate, group = term, color = term)) +
  geom_line() +
  scale_x_log10(limits = range(fit$lambda[fit$df > 0])) # 截断到有非零系数的lambda区间

2. 添加显示变量数量的次X轴

原生plot的次X轴展示对应lambda下的非零系数数量(fit$df),通过映射lambda与df的对应关系添加次轴:

# 准备次X轴刻度数据:选取关键lambda对应的非零变量数
axis_df <- tibble(
  lambda = fit$lambda[fit$df %in% unique(fit$df)[seq(1, length(unique(fit$df)), by = 2)]],
  df = fit$df[fit$df %in% unique(fit$df)[seq(1, length(unique(fit$df)), by = 2)]]
)

# 为基础图添加次X轴并美化
p +
  annotation_logticks(sides = "b") +
  scale_x_continuous(
    trans = "log10",
    limits = range(fit$lambda[fit$df > 0]),
    # 主X轴:lambda刻度
    breaks = axis_df$lambda,
    labels = round(axis_df$lambda, 3),
    # 次X轴:非零变量数量刻度
    sec.axis = sec_axis(
      ~.,
      breaks = axis_df$lambda,
      labels = axis_df$df,
      name = "非零变量数量"
    )
  ) +
  labs(x = "Lambda (log10)", y = "系数估计值") +
  theme_bw()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 21:53:10