You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何大幅提升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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.09 17:47:02