R语言无循环低计算成本实现Matrix Computation矩阵构建
R语言向量化构建目标矩阵方案
问题背景
需要根据指定公式构建目标矩阵:
初始实现采用双层嵌套循环,且存在冗余计算,运行效率低,初始代码如下:
V <- function(x,z){ V <- matrix(NA,nrow(x),nrow(x)) for(i in 1:nrow(x)){ for(j in 1:nrow(x)){ d11 <- apply(z,1,FUN = function(y) ((z-x[i,]))) d12 <- apply(z,1,FUN = function(y) ((z-x[j,]))) d1 <- d11*d12 d2 <- apply(d1, 2 , prod) d3 <- sum(d2) V[i,j]= d3 } } return(V) }
注:初始代码的apply调用存在笔误,传入的行参数y未被使用,每次循环都重复计算全量差值矩阵,额外增加了不必要的开销。
优化思路
目标矩阵的元素计算逻辑为:V[i,j] = 对z的每一行k,累乘所有维度d上的(z[k,d]-x[i,d])*(z[k,d]-x[j,d]),再对所有k的结果求和
根据乘法交换律,维度乘积项可以拆分:prod_d (z[k,d] - x[i,d])*(z[k,d] - x[j,d]) = prod_d(z[k,d]-x[i,d]) * prod_d(z[k,d]-x[j,d])
提前构造中间矩阵U,其中U[i,k] = prod_d (z[k,d] - x[i,d]),则最终目标矩阵可以直接通过矩阵乘法得到:V = U %*% t(U),完全避免显式循环,计算效率提升显著。
实现代码
无第三方依赖版本(基础R)
兼容所有原生R环境,无需安装额外包:
V_fast <- function(x, z) { n <- nrow(x) m <- nrow(z) # 扩展矩阵批量计算所有z[k,d] - x[i,d] x_expand <- x[rep(seq_len(n), each = m), , drop = FALSE] z_expand <- z[rep(seq_len(m), n), , drop = FALSE] diff_mat <- z_expand - x_expand # 按行计算维度乘积 u_vec <- apply(diff_mat, 1, prod) # 构造中间矩阵U并计算最终结果,tcrossprod等价于U%*%t(U),效率更高 U <- matrix(u_vec, nrow = n, ncol = m, byrow = TRUE) tcrossprod(U) }
高性能版本(推荐)
使用matrixStats包的行乘积函数替代apply,大矩阵下速度可以再提升一个数量级:
# 首次使用需先安装:install.packages("matrixStats") V_faster <- function(x, z) { n <- nrow(x) m <- nrow(z) x_expand <- x[rep(seq_len(n), each = m), , drop = FALSE] z_expand <- z[rep(seq_len(m), n), , drop = FALSE] u_vec <- matrixStats::rowProds(z_expand - x_expand) U <- matrix(u_vec, nrow = n, ncol = m, byrow = TRUE) tcrossprod(U) }
结果校验
可以用测试数据验证优化版本和原始逻辑的结果一致性:
# 修正笔误后的原始循环版本,作为结果基准 V_original <- function(x,z){ V <- matrix(NA,nrow(x),nrow(x)) for(i in 1:nrow(x)){ for(j in 1:nrow(x)){ d11 <- apply(z,1,FUN = function(y) y - x[i,]) d12 <- apply(z,1,FUN = function(y) y - x[j,]) d1 <- d11*d12 d2 <- apply(d1, 2 , prod) V[i,j] <- sum(d2) } } return(V) } # 一致性测试 set.seed(123) x <- matrix(rnorm(20), nrow = 5, ncol = 4) z <- matrix(rnorm(40), nrow = 10, ncol = 4) all.equal(V_original(x,z), V_fast(x,z)) # 返回TRUE表示结果一致 all.equal(V_original(x,z), V_faster(x,z)) # 返回TRUE表示结果一致
当x的行数超过100时,原始双层循环版本的计算时间会达到数秒到数十秒,而优化版本仅需几毫秒即可完成计算。
内容的提问来源于stack exchange,提问作者siamak khanjari
相关产品推荐
相关产品推荐

