Lasso回归坐标下降算法系数爆炸问题排查求助
Lasso坐标下降算法系数爆炸问题排查求助
我参考资料在R中实现了Lasso回归的坐标下降算法,按照教材将损失函数按1/n进行缩放(也尝试了未缩放版本,结果无差异)。当样本量N大于特征数p时一切正常,但当p>N或预测变量存在协方差(如模拟多元正态协方差数据)时,除了lambda足够大使所有系数为0的情况外,估计系数在几次迭代后会出现爆炸。我多次检查代码未发现错误,Lasso本应适用于此类场景,故求助排查问题。
复现问题的最小示例代码
# some libraries might be overkill library(dplyr) library(glue) library(tidyverse) library(MASS) library(rlist) RSS <- function(Y, X, beta){ n <- length(Y) return(sapply( 1:length(beta[,1]), function(x){1/n*sum((Y - beta[x,]%*%t(X))^2)} )) } set.seed(1234) n <- 5 p <- 10 lam <- 1 #get the covar matrix Sigma <- matrix(0.5, nrow = p, ncol = p) diag(Sigma) <- 1 mu <- rep(0, p) X <- mvrnorm(n, mu, Sigma) # scale X to omit beta0 X <- scale(X) # simulate data beta_true <- rep(c(1,-1,0), each = ceiling(p/3))[1:p] Y <- beta_true%*%t(X) beta_est <- Lasso(Y, X, rep(1, length(X[1,]))) rss <- RSS(Y, X, beta_est) penalty <- lam * rowSums(abs(beta_est)) ggplot(mapping = aes(x = 1:length(rss))) + geom_point(aes(y = rss + penalty, color = "Penalty + RSS")) + geom_point(aes(y = rss, color = "RSS")) + geom_point(aes(y = penalty, color = "Penalty")) + scale_color_manual(values = c("Penalty + RSS" = "black", "RSS" = "red", "Penalty" = "blue"), name = "Legend") + labs(y = NULL, x = "Iteration")
我的Lasso回归函数代码
soft_threshold <- function(theta, lambda) { return(sign(theta) * max(abs(theta) - lambda, 0)) } Lasso <- function(y, X, lam, starting_coef){ p <- length(starting_coef) n <- length(y) coef <- matrix(starting_coef, nrow = 1) z_j <- 1/n*colSums(X^2) # can calculate outside of loop t <- 1 # set time step # max number of iterations while(t<=50){ new_coef <- c() # initialize the new coefs # do it for all coefs cyclical for(j in 1:p){ roh_j <- 1/n * ( y - coef[t,-j] %*% t(X[, -j])) %*% X[,j] # the calculation of roh_j splitted up #r_ij <- y - coef[t, -j]%*%t(X[,-j]) #roh_j <- 1/n * sum( r_ij * X[,j] ) new_beta <- 1/z_j[j] * soft_threshold(roh_j, lam) new_coef <- c(new_coef, new_beta) } coef <- rbind(coef, new_coef) # check if significant change if(max(abs(coef[t+1,] - coef[t,])) < 10^-4){ break } t <- t+1 } return(coef) }
内容的提问来源于stack exchange,提问作者UeberQ
相关产品推荐
相关产品推荐

