优化R语言for循环提速:15万行数据集函数性能优化求助
R函数execution_strategy_export性能优化方案
你当前处理15万行数据集时,已使用future_lapply多核处理和cmpfun编译,但仍需提速,且因逻辑复杂难以用data.table重构,以下是几种可行的优化方向:
1. 替换数据框逐行操作为向量访问
原代码中data[i, "col"]这种数据框索引方式效率极低,尤其是在循环内重复调用。建议提前将所有需要的列提取为向量,循环内直接通过向量索引访问,同时预分配结果向量避免动态修改数据框:
execution_strategy_export_optimized <- function(data) { # 预提取所有列向量 close_vec <- data$close loss_interval_vec <- data$loss_interval ref_rate_vec <- data$ref_rate high_vec <- data$high adjustment_interval_vec <- data$adjustment_interval hour_vec <- data$hour period_one_vec <- data$period_one time_vec <- data$time open_vec <- data$open low_vec <- data$low # 预分配结果向量 n <- nrow(data) outside_business_hours <- rep(NA_character_, n) trigger_period <- rep(NA_integer_, n) trigger_time <- rep(NA, n) post_period_two <- rep(0, n) result <- rep(NA, n) adjusted_count <- rep(NA_integer_, n) for (i in 1:n) { today_period_one_close <- close_vec[i] today_loss_interval <- loss_interval_vec[i] today_ref_rate <- ref_rate_vec[i] today_period_one_high <- high_vec[i] today_adjusted_interval <- adjustment_interval_vec[i] today_adjusted_count <- floor((today_period_one_high - today_ref_rate) / today_adjusted_interval) pd_max_period_one <- today_ref_rate + (today_adjusted_interval * today_adjusted_count) - today_loss_interval IF1 <- 0 if (hour_vec[i] < 8 || hour_vec[i] > 17) { outside_business_hours[i] <- "Excl" } else { period_one <- period_one_vec[i] if (period_one > 0) { trigger_period[i] <- 1 trigger_time[i] <- time_vec[i] } else { j <- i + 1 current_adjusted_count <- today_adjusted_count while (j <= n) { today_open_j <- open_vec[j] today_high_j <- high_vec[j] today_low_j <- low_vec[j] today_close_j <- close_vec[j] input_dbl <- (today_period_one_close - today_ref_rate) / today_adjusted_interval # 替换向量化ifelse为标量if/else,提升单值判断效率 if (today_period_one_close > today_ref_rate) { close_ref_rate_adjusted_count <- floor(input_dbl) } else { close_ref_rate_adjusted_count <- ceiling(input_dbl) } prior_period_close_SL <- today_ref_rate + (today_adjusted_interval * close_ref_rate_adjusted_count) - today_loss_interval input_dbl_up <- (today_high_j - today_ref_rate) / today_adjusted_interval current_adjusted_count <- max(current_adjusted_count, floor(input_dbl_up)) today_max <- today_ref_rate + (today_adjusted_interval * current_adjusted_count) - today_loss_interval if (today_open_j < pd_max_period_one || (today_low_j < today_max || today_close_j < today_max)) { if (today_open_j < pd_max_period_one) { IF1 <- min(prior_period_close_SL, pd_max_period_one) } else { IF1 <- today_max } } trigger_period[i] <- j - i + 1 if (trigger_period[i] > 1 && (i + trigger_period[i]) <= n) { trigger_time[i] <- time_vec[i + trigger_period[i] - 1] } else { trigger_time[i] <- time_vec[i] } if (IF1 > 0) { break } j <- j + 1 } } } post_period_two[i] <- IF1 result[i] <- period_one_vec[i] + IF1 adjusted_count[i] <- current_adjusted_count } # 最后一次性将结果向量赋值回数据框 data$outside_business_hours <- outside_business_hours data$trigger_period <- trigger_period data$trigger_time <- trigger_time data$post_period_two <- post_period_two data$result <- result data$adjusted_count <- adjusted_count return(data) }
2. 用Rcpp重写核心嵌套循环
R的原生循环性能有限,外层for+内层while的嵌套逻辑是15万行数据的主要性能瓶颈。用Rcpp将核心逻辑转为C++代码,可获得10-100倍的速度提升:
步骤1:编写Rcpp代码
将以下代码保存为trigger_logic.cpp:
#include <Rcpp.h> #include <algorithm> using namespace Rcpp; // [[Rcpp::export]] List findTrigger(NumericVector close_vec, NumericVector loss_interval_vec, NumericVector ref_rate_vec, NumericVector high_vec, NumericVector adjustment_interval_vec, NumericVector open_vec, NumericVector low_vec, IntegerVector hour_vec, IntegerVector period_one_vec, CharacterVector time_vec) { int n = close_vec.size(); CharacterVector outside_business_hours(n, NA_STRING); IntegerVector trigger_period(n, NA_INTEGER); CharacterVector trigger_time(n, NA_STRING); NumericVector post_period_two(n, 0.0); NumericVector result(n, NA_REAL); IntegerVector adjusted_count(n, NA_INTEGER); for (int i = 0; i < n; ++i) { double today_period_one_close = close_vec[i]; double today_loss_interval = loss_interval_vec[i]; double today_ref_rate = ref_rate_vec[i]; double today_period_one_high = high_vec[i]; double today_adjusted_interval = adjustment_interval_vec[i]; int today_adjusted_count = floor((today_period_one_high - today_ref_rate) / today_adjusted_interval); double pd_max_period_one = today_ref_rate + (today_adjusted_interval * today_adjusted_count) - today_loss_interval; double IF1 = 0.0; if (hour_vec[i] < 8 || hour_vec[i] > 17) { outside_business_hours[i] = "Excl"; } else { int period_one = period_one_vec[i]; if (period_one > 0) { trigger_period[i] = 1; trigger_time[i] = time_vec[i]; } else { int j = i + 1; int current_adjusted_count = today_adjusted_count; while (j < n) { double today_open_j = open_vec[j]; double today_high_j = high_vec[j]; double today_low_j = low_vec[j]; double today_close_j = close_vec[j]; double input_dbl = (today_period_one_close - today_ref_rate) / today_adjusted_interval; int close_ref_rate_adjusted_count; if (today_period_one_close > today_ref_rate) { close_ref_rate_adjusted_count = floor(input_dbl); } else { close_ref_rate_adjusted_count = ceiling(input_dbl); } double prior_period_close_SL = today_ref_rate + (today_adjusted_interval * close_ref_rate_adjusted_count) - today_loss_interval; double input_dbl_up = (today_high_j - today_ref_rate) / today_adjusted_interval; current_adjusted_count = std::max(current_adjusted_count, (int)floor(input_dbl_up)); double today_max = today_ref_rate + (today_adjusted_interval * current_adjusted_count) - today_loss_interval; if (today_open_j < pd_max_period_one || (today_low_j < today_max || today_close_j < today_max)) { if (today_open_j < pd_max_period_one) { IF1 = std::min(prior_period_close_SL, pd_max_period_one); } else { IF1 = today_max; } } trigger_period[i] = j - i + 1; if (trigger_period[i] > 1 && (i + trigger_period[i]) < n) { trigger_time[i] = time_vec[i + trigger_period[i] - 1]; } else { trigger_time[i] = time_vec[i]; } if (IF1 > 0) { break; } j++; } adjusted_count[i] = current_adjusted_count; } } post_period_two[i] = IF1; result[i] = period_one_vec[i] + IF1; adjusted_count[i] = current_adjusted_count; } return List::create( _["outside_business_hours"] = outside_business_hours, _["trigger_period"] = trigger_period, _["trigger_time"] = trigger_time, _["post_period_two"] = post_period_two, _["result"] = result, _["adjusted_count"] = adjusted_count ); }
步骤2:编译并调用
在R中执行以下代码编译并使用:
library(Rcpp) sourceCpp("trigger_logic.cpp") execution_strategy_export_rcpp <- function(data) { # 调用Rcpp函数计算结果 res_list <- findTrigger(data$close, data$loss_interval, data$ref_rate, data$high, data$adjustment_interval, data$open, data$low, data$hour, data$period_one, data$time) # 将结果合并到原数据框 data$outside_business_hours <- res_list$outside_business_hours data$trigger_period <- res_list$trigger_period data$trigger_time <- res_list$trigger_time data$post_period_two <- res_list$post_period_two data$result <- res_list$result data$adjusted_count <- res_list$adjusted_count return(data) }
3. 其他细节优化
- 移除冗余赋值:原代码中
adjusted_count[i]被多次赋值,需确认逻辑后保留最终有效的赋值(示例中已修正)。 - 预计算常量:如果
loss_interval等列存在大量重复值,可提前缓存计算结果避免重复运算。 - 内存优化:确保数据框列类型正确(比如数值列用
numeric而非character),减少内存占用和类型转换开销。
内容的提问来源于stack exchange,提问作者Abri De Beer
相关产品推荐
相关产品推荐

