加速含大量矩阵逆运算的R语言非线性回归模拟循环
非线性回归模拟代码加速与结构优化需求
我正在开展非线性回归模拟工作,核心函数Main_fn包含两次矩阵逆运算,目前未实现向量化。中等样本量(如n=1000、p=100)下,100次重复模拟耗时约12小时。仅熟悉R语言,不懂C++,寻求代码加速方案及结构优化建议。
现有代码
数据生成代码
ndata = 40 n = 50 t = 20 N = n * t p = 10 # 实际场景中p通常大于100 res = array(NA, dim = c(N, 2*p*(n-1), ndata)) for (j in 1:ndata) { u = runif(N) gu = cbind(2 * sin(2*pi*u), u * (1-2*u), exp(-u + 0.5), matrix(1,nrow = 1000) %x% matrix(1:7,nrow = 1)) Mu = rnorm(n-1) Lambda = rnorm(t-1) D = t(cbind(matrix(-1, n-1, 1), diag(n-1))) %x% matrix(1,t,1) S = matrix(1,n,1) %x% t(cbind(matrix(-1, t-1, 1), diag(t-1))) x = MASS::mvrnorm(N, rep(0,p), 0.5**abs(outer(1:p,1:p,'-'))) eps = rnorm(N) y = apply(x * gu, 1, sum) + D %*% Mu + S %*% Lambda + eps res[,,j] = Main_fn(x, y, u) }
核函数与核心函数代码
核函数My_den_fn
My_den_fn <- function(x,h) { idx = 0.75 * (1 - (x/h)**2) / h # 修正原代码笔误:将全局变量t替换为输入参数x kernal = 0.50 * (abs(idx) + idx) return(kernal) }
原始核心函数Main_fn
Main_fn <- function (x, y, u) { N = dim(x)[1] p = dim(x)[2] h = sd(u) * N**(- 0.2) * 2 n = 50 t = 20 D = t(cbind(matrix(-1, n-1, 1), diag(n-1))) %x% matrix(1,t,1) S = matrix(1,n,1) %x% t(cbind(matrix(-1, t-1, 1), diag(t-1))) M = cbind(D, S) ab = array(NA, dim = c(N, 2 * p)) for(i in 1:N) { W = diag(My_den_fn(u - u[i], h)) # My_den_fn返回值>=0 Q = diag(N) - M %*% solve(t(M) %*% W %*% M + 1e-3 * diag(n+t-2)) %*% crossprod(M, W) # 投影矩阵Q W_new = crossprod(Q,W %*% Q) G = cbind(x, x*(u - u[i])) ab[i,] = solve(crossprod(G,W_new)%*% G + 1e-3 * diag(2*p)) %*% (crossprod(G,W_new)%*% y) } return(ab) }
尝试优化后的函数Main_fn_A
Main_fn_A <- function (x, y, u) { N = dim(x)[1] p = dim(x)[2] h = sd(u) * N**(- 0.2) * 2 n = 50 t = 20 D = t(cbind(matrix(-1, n-1, 1), diag(n-1))) %x% matrix(1,t,1) S = matrix(1,n,1) %x% t(cbind(matrix(-1, t-1, 1), diag(t-1))) M = cbind(D, S) U = outer(u,u,'-') W = sapply(X = 1:N, FUN = function(i){My_den_fn(U[i,],h)}, simplify = TRUE) # 值>=0 Q = lapply(X = 1:N, FUN = function(i){diag(N)-M %*% solve(crossprod(M,W[,i]*M)+1e-3*diag(n+t-2)) %*% t(M*W[,i])}) W_new = lapply(X = 1:N, FUN = function(i){crossprod(Q[[i]],W[,i]*Q[[i]])}) GAMMA = lapply(X = 1:N, FUN = function(i){cbind(x, x*U[,i])}) ab = array(NA, dim = c(N, 2 * p)) for(i in 1:N) { ab[i,] = solve(crossprod(GAMMA[[i]],W_new[[i]]) %*% GAMMA[[i]] + 1e-3 * diag(2*p)) %*% (crossprod(GAMMA[[i]],W_new[[i]]) %*% y) } return(ab) }
代码加速方案
1. 避免重复计算静态矩阵
D、S、M仅依赖n和t,可提前生成并传入函数,减少重复计算:
# 在数据生成循环外预计算 n = 50 t = 20 D_global = t(cbind(matrix(-1, n-1, 1), diag(n-1))) %x% matrix(1,t,1) S_global = matrix(1,n,1) %x% t(cbind(matrix(-1, t-1, 1), diag(t-1))) M_global = cbind(D_global, S_global)
之后在Main_fn中直接使用M_global即可。
2. 优化矩阵求逆运算
将solve(A) %*% b替换为solve(A, b),R会自动选择更高效的算法,避免显式求逆:
- 原
Q的计算改写为:
W_vec = W_vals[i, ] tM_W_M = crossprod(M_global, W_vec * M_global) inv_term = solve(tM_W_M + 1e-3 * diag(n+t-2), crossprod(M_global, W_vec)) Q = diag(N) - M_global %*% inv_term
- 最终系数求解改写为:
A = crossprod(G, W_new) %*% G + 1e-3 * diag(2*p) b = crossprod(G, W_new) %*% y ab[i,] = solve(A, b)
3. 向量化核函数计算
一次性计算所有u的差值,批量生成核函数值,避免循环内重复计算:
U = outer(u, u, "-") W_vals = My_den_fn(U, h)
同时用权重向量代替大对角矩阵,减少内存开销:比如t(M) %*% W %*% M等价于crossprod(M, W_vec * M),无需生成N×N的对角矩阵。
4. 并行化独立计算
每个i的计算相互独立,用parallel包并行化循环:
library(parallel) num_cores = detectCores() - 1 Main_fn_parallel <- function(x, y, u, M_global) { N = nrow(x) p = ncol(x) h = sd(u) * N^(-0.2) * 2 U = outer(u, u, "-") W_vals = My_den_fn(U, h) n_t = ncol(M_global) compute_i <- function(i) { W_vec = W_vals[i, ] tM_W_M = crossprod(M_global, W_vec * M_global) inv_term = solve(tM_W_M + 1e-3 * diag(n_t), crossprod(M_global, W_vec)) Q = diag(N) - M_global %*% inv_term W_new = crossprod(Q, W_vec * Q) G = cbind(x, x * U[i, ]) A = crossprod(G, W_new) %*% G + 1e-3 * diag(2*p) b = crossprod(G, W_new) %*% y solve(A, b) } ab_list = mclapply(1:N, compute_i, mc.cores = num_cores) ab = do.call(rbind, ab_list) return(ab) }
Windows系统可改用parLapply并提前创建集群。
5. 使用高效矩阵工具包
推荐使用Matrix包处理稀疏矩阵(若W大部分为0),或fastmatrix包加速矩阵运算,无需额外C++知识。
代码结构优化建议
- 参数传递优先:将静态矩阵、固定参数(如
n、t)作为函数参数传入,避免依赖全局变量,减少重复计算。 - 模块化拆分:将投影矩阵计算、权重更新、系数求解拆分为独立小函数,提升代码可读性和可维护性。
- 内存优化:避免生成不必要的大矩阵,用向量权重代替对角矩阵,降低内存占用。
- 补全注释:给关键步骤(如正则项作用、投影矩阵含义)添加注释,方便后续调试和维护。
内容的提问来源于stack exchange,提问作者MOHAMMED
相关产品推荐
相关产品推荐

