R语言中stats::uniroot函数是否有运行速度更快的替代方案?
方案1:直接使用解析解(最优,速度提升1000倍以上)
你当前的求根场景存在闭式解析解,完全不需要调用数值求根函数:
你要求解的方程是 b * (t/b)^a = u,直接推导可得:t = b * (u / b) ^ (1/a)
直接在data.table中向量化计算即可,完全没有逐行调用的开销:
library(data.table) cumhaz <- function(t, a, b) b * (t/b)^a froot <- function(x, u, a, b) cumhaz(x, a, b) - u n <- 50000 u <- -log(runif(n)) a <- 1/2 b <- 1 dt = data.table(u = u, a = a, b = b) # 原方案耗时(8s左右) print(system.time( dt[, c := uniroot(froot, u=u, a=a, b=b, interval= c(0.01, 10), extendInt="yes")$root, by = u] )) # 解析解方案耗时(<0.01s) print(system.time( dt[, c_new := b * (u / b) ^ (1/a)] )) # 验证结果一致 all.equal(dt$c, dt$c_new)
运行验证可以看到结果完全一致,5万行运算耗时不到0.01秒,一百万行也仅需几十毫秒。
方案2:通用无解析解场景的提速方案
如果你实际使用的函数没有解析解,必须用数值求根,可以通过以下方式优化:
- 移除无意义的
by = u分组:你当前每个u都是唯一值,逐行分组会带来极大的额外开销 - 使用向量化的求根函数,替代逐行调用stats::uniroot
- 编译目标函数减少R层函数调用开销
示例代码:
library(data.table) library(nleqslv) library(compiler) # 编译目标函数为字节码,减少调用开销 cumhaz <- cmpfun(function(t, a, b) b * (t/b)^a) froot <- cmpfun(function(x, u, a, b) cumhaz(x, a, b) - u) n <- 50000 u <- -log(runif(n)) a <- 1/2 b <- 1 dt = data.table(u = u, a = a, b = b) # 向量化调用nleqslv求根 print(system.time({ res <- nleqslv(rep(1, nrow(dt)), froot, u = dt$u, a = dt$a, b = dt$b) dt$c <- res$x }))
该方案5万行耗时通常在0.5秒以内,比原方案提速15倍以上。如果还需要进一步提速,可以使用Rcpp实现自定义的二分法/牛顿法求根逻辑,速度还能再提升数倍。
内容的提问来源于stack exchange,提问作者Saurabh
相关产品推荐
相关产品推荐

