如何大幅提升R语言smsurv函数的运行效率?
提升R语言smsurv函数运行效率的优化方案
我正在尝试大幅提升R语言中smsurv函数的运行效率,下面是可运行的示例代码。我已经试过用lapply和sapply替换原函数中的for循环,但运行速度反而变慢。请问还有哪些可行的优化方案?我知道Rcpp包,但不清楚如何用C代码替换循环,希望能得到相关建议。
原函数代码
smsurv <- function(Time,Status,X,beta,w,model){ death_point <- sort(unique(subset(Time, Status==1))) if(model=='ph') coxexp <- exp((beta)%*%t(X[,-1])) n <- length(death_point) lambda <- numeric(n) for(i in 1: n){ if(model=='ph') temp <- sum(as.numeric(Time>=death_point[i])*w*drop(coxexp)) if(model=='aft') temp <- sum(as.numeric(Time>=death_point[i])*w) lambda[i] <- sum(Status*as.numeric(Time==death_point[i]))/temp } HHazard <- numeric() for(i in 1:length(Time)){ HHazard[i] <- sum(as.numeric(Time[i]>=death_point)*lambda) if(Time[i]>max(death_point))HHazard[i] <- Inf if(Time[i]<min(death_point))HHazard[i] <- 0 } survival <- exp(-HHazard) list(survival=survival) } nr_obs = 50000 Time_input <- rnorm(nr_obs, mean = 100, sd = 36) Status_input <- sample(c(0,1), replace=TRUE, size=nr_obs) w_input <- Status_input # 假设存在9个变量(第一列为截距项) n_variables <- 9 X_input <- matrix(rnorm(nr_obs*n_variables),nr_obs) X_input <- cbind(Intercept = rep(1, nrow(X_input)), X_input) beta_input <- runif(n_variables, min = -1, max = 1) model_input <- "ph" output <- smsurv(Time_input,Status_input,X_input,beta_input,w_input,model_input)
尝试优化后的函数代码
smsurv2 <- function(Time,Status,X,beta,w,model){ death_point <- sort(unique(subset(Time, Status==1))) if(model=='ph') coxexp <- exp((beta)%*%t(X[,-1])) if(model_input=='ph') lambda =unlist(lapply(death_point, function(z) sum(Status_input*as.numeric(Time_input==z))/ sum(as.numeric(Time_input>=z)*w_input*drop(coxexp)))) if(model=='aft') lambda =unlist( lapply(death_point, function(z) sum(Status_input*as.numeric(Time_input==z))/ sum(as.numeric(Time_input>=z)*w_input))) HHazard <- unlist(lapply(Time, function(t) {sum(as.numeric(t>=death_point)*lambda)})) HHazard[Time > max(death_point)] <- Inf HHazard[Time < min(death_point)] <- 0 survival <- exp(-HHazard) list(survival=survival) } smsurv3 <- function(Time, Status, X, beta, w, model){ death_point <- sort(unique(subset(Time, Status==1))) if(model=='ph') coxexp <- exp((beta)%*%t(X[,-1])) lambda <- sapply(death_point, function(dp) {return(sum(Status*as.numeric(Time==dp))/sum(as.numeric(Time>=dp)*w*drop(coxexp)))}) HHazard <- sapply(Time, function(t){return(sum(as.numeric(t>=death_point)*lambda))}) HHazard[Time > max(death_point)] <- Inf HHazard[Time < min(death_point)] <- 0 survival <- exp(-HHazard) list(survival=survival) }
优化方案
1. 纯R内的向量化与预计算优化
lapply/sapply本质还是循环,且存在函数调用开销,反而不如优化原生for循环或用完全向量化操作:
- 预计算重复值:
max(death_point)、min(death_point)只计算一次,避免循环内重复计算;coxexp计算后直接转为向量,省去每次drop()的开销。 - 矩阵运算批量计算lambda:将
Time >= death_point转为逻辑矩阵,用矩阵乘法批量计算分子(事件数)和分母(风险集权重和),替代逐个循环。 - 快速计算HHazard:利用
death_point已排序的特性,用findInterval定位每个Time对应的位置,结合cumsum(lambda)直接获取累加值,比逐个计算sum(as.numeric(t>=death_point)*lambda)高效数倍。
优化后的纯R函数示例:
smsurv_opt <- function(Time, Status, X, beta, w, model){ death_point <- sort(unique(Time[Status == 1])) n_dp <- length(death_point) if(n_dp == 0) stop("No death events") # 预计算重复值 min_dp <- min(death_point) max_dp <- max(death_point) # 处理权重:ph模型加入coxexp if(model == 'ph'){ coxexp <- exp(drop(X[,-1] %*% beta)) # 调整矩阵乘法顺序,减少内存占用 weights <- w * coxexp } else { weights <- w } # 向量化计算lambda的分子和分母 numerator <- sapply(death_point, function(dp) sum(Status[Time == dp])) denominator <- colSums(outer(Time, death_point, `>=`) * weights) lambda <- numerator / denominator # 计算累积lambda cum_lambda <- cumsum(lambda) # 快速生成HHazard idx <- findInterval(Time, death_point) HHazard <- cum_lambda[idx] HHazard[Time < min_dp] <- 0 HHazard[Time > max_dp] <- Inf survival <- exp(-HHazard) list(survival = survival) }
2. Rcpp实现方案
如果纯R优化仍达不到需求,Rcpp可直接操作底层数据,彻底消除R循环的开销,以下是核心实现思路:
步骤1:编写Rcpp代码
创建smsurv_rcpp.cpp文件,内容如下:
#include <Rcpp.h> #include <algorithm> using namespace Rcpp; // [[Rcpp::export]] NumericVector smsurv_rcpp(NumericVector Time, IntegerVector Status, NumericMatrix X, NumericVector beta, NumericVector w, std::string model) { // 提取并排序去重死亡时间点 NumericVector death_point; for(int i = 0; i < Time.size(); ++i){ if(Status[i] == 1){ death_point.push_back(Time[i]); } } std::sort(death_point.begin(), death_point.end()); death_point.erase(std::unique(death_point.begin(), death_point.end()), death_point.end()); int n_dp = death_point.size(); if(n_dp == 0){ stop("No death events"); } double min_dp = death_point[0]; double max_dp = death_point[n_dp - 1]; // 计算权重:ph模型加入coxexp NumericVector weights = clone(w); if(model == "ph"){ NumericVector X_beta = X(_, Range(1, X.ncol()-1)) * beta; NumericVector coxexp = exp(X_beta); weights = weights * coxexp; } // 计算lambda和累积lambda NumericVector lambda(n_dp); NumericVector cum_lambda(n_dp); for(int i = 0; i < n_dp; ++i){ double dp = death_point[i]; int numerator = 0; double denominator = 0.0; for(int j = 0; j < Time.size(); ++j){ if(Time[j] == dp && Status[j] == 1) numerator++; if(Time[j] >= dp) denominator += weights[j]; } lambda[i] = (double)numerator / denominator; cum_lambda[i] = (i == 0) ? lambda[i] : cum_lambda[i-1] + lambda[i]; } // 计算HHazard NumericVector HHazard(Time.size()); for(int i = 0; i < Time.size(); ++i){ double t = Time[i]; if(t < min_dp){ HHazard[i] = 0.0; } else if(t > max_dp){ HHazard[i] = R_PosInf; } else { // 找到第一个大于t的死亡点索引,取前一个的累积值 int idx = std::upper_bound(death_point.begin(), death_point.end(), t) - death_point.begin(); HHazard[i] = (idx == 0) ? 0.0 : cum_lambda[idx - 1]; } } return exp(-HHazard); }
步骤2:编译并使用
在R中运行以下代码编译调用:
library(Rcpp) sourceCpp("smsurv_rcpp.cpp") # 使用示例 output_rcpp <- smsurv_rcpp(Time_input, Status_input, X_input, beta_input, w_input, model_input)
3. 其他小优化点
- 避免循环内重复调用
as.numeric():R中逻辑值可直接参与运算,无需转换。 - 调整矩阵乘法顺序:将
beta %*% t(X[,-1])改为X[,-1] %*% beta,后者生成向量而非大矩阵,内存占用和计算速度更优。 - 减少不必要的列表转换:比如
unlist(lapply(...))的开销可通过向量化操作避免。
内容的提问来源于stack exchange,提问作者John
相关产品推荐
相关产品推荐

