基于Rcpp实现数据截断与协方差矩阵计算
Rcpp实现按行截断后计算样本协方差(与原R代码等效)
原R可复现代码
# 单向量截断函数 trunc_operator <- function(x, tau) { p <- length(x) ord <- order(abs(x), decreasing = TRUE) x[ord[(tau + 1):p]] <- 0 return(x) } # 主函数:按行截断后计算样本协方差 trunc_cov_R <- function(X, tau_vec) { stopifnot(nrow(X) == length(tau_vec)) X_trunc <- t(apply(X, 1, function(row) { tau <- tau_vec[rownames(X) == rownames(row)] trunc_operator(row, tau) })) cov_mat <- cov(X_trunc) return(cov_mat) } # 测试数据生成 set.seed(123) n <- 100 p <- 20 X <- matrix(rnorm(n*p), nrow = n, ncol = p) tau_vec <- sample(1:p, n, replace = TRUE) # 运行R版本得到基准结果 cov_R <- trunc_cov_R(X, tau_vec)
等效Rcpp实现
创建一个.cpp文件(或直接用sourceCpp内嵌代码),内容如下:
#include <Rcpp.h> #include <algorithm> using namespace Rcpp; // 对单个向量执行截断操作 NumericVector trunc_operator_cpp(NumericVector x, int tau) { int p = x.size(); IntegerVector idx = seq(0, p - 1); // 按元素绝对值降序排序索引 std::sort(idx.begin(), idx.end(), [&](int a, int b) { return std::abs(x[a]) > std::abs(x[b]); }); // 截断:tau之后的元素置0 for (int i = tau; i < p; ++i) { x[idx[i]] = 0.0; } return x; } // [[Rcpp::export]] NumericMatrix trunc_cov_cpp(NumericMatrix X, IntegerVector tau_vec) { int n = X.nrow(); int p = X.ncol(); // 参数校验 if (tau_vec.size() != n) { stop("tau_vec length must match number of rows in X"); } // 生成截断后的矩阵 NumericMatrix X_trunc(n, p); for (int i = 0; i < n; ++i) { X_trunc.row(i) = trunc_operator_cpp(X.row(i), tau_vec[i]); } // 计算列均值 NumericVector col_means(p, 0.0); for (int j = 0; j < p; ++j) { for (int i = 0; i < n; ++i) { col_means[j] += X_trunc(i, j); } col_means[j] /= n; } // 计算样本协方差矩阵(除以n-1) NumericMatrix cov_mat(p, p); for (int j1 = 0; j1 < p; ++j1) { for (int j2 = j1; j2 < p; ++j2) { double cov_val = 0.0; for (int i = 0; i < n; ++i) { cov_val += (X_trunc(i, j1) - col_means[j1]) * (X_trunc(i, j2) - col_means[j2]); } cov_val /= (n - 1); cov_mat(j1, j2) = cov_val; cov_mat(j2, j1) = cov_val; } } return cov_mat; }
结果验证
在R环境中执行以下代码,确认两个版本结果一致:
library(Rcpp) # 编译Rcpp函数(如果是文件,替换为sourceCpp("your_file.cpp")) sourceCpp(code = '上面的cpp代码内容') # 或者直接用文件路径 # 运行Rcpp版本 cov_cpp <- trunc_cov_cpp(X, tau_vec) # 检查结果一致性(允许微小浮点误差) all.equal(cov_R, cov_cpp, tolerance = 1e-10)
这个Rcpp实现通过避免R中apply的循环开销,直接在底层处理矩阵运算,能大幅提升大矩阵下的运行效率。
内容的提问来源于stack exchange,提问作者Dylan Dijk
相关产品推荐
相关产品推荐

