如何在R中借助Rcpp使用C++ ODE求解器?含R端ODE实现示例
在R中用Rcpp调用C++ ODE求解器的完整指南
针对你的SIR变体模型,我一步步给你拆解怎么做,还会帮你完成R和C++的速度对比。
1. 先搞定依赖包
首先得安装几个关键的包:Rcpp用来衔接R和C++,RcppODE提供C++原生的ODE求解器,还有你已经在用的deSolve用来对比R版的性能,再加个microbenchmark做速度测试。跑下面的代码安装:
install.packages(c("Rcpp", "RcppODE", "deSolve", "microbenchmark"))
2. 把你的R版ODE模型转成C++代码
你写的modelsir_cpp是R版本的,我们把它翻译成C++,这样才能让R调用。新建一个名为sir_model.cpp的文件,把下面的代码粘进去:
#include <Rcpp.h> using namespace Rcpp; // 这个标记告诉Rcpp要把这个函数暴露给R // [[Rcpp::export]] NumericVector sir_cpp(double t, NumericVector x, NumericVector parms) { // 提取状态变量 double S = x[0]; double I1 = x[1]; double I2 = x[2]; double N = S + I1 + I2; // 提取参数 double B = parms[0]; double mu = parms[1]; double beta = parms[2]; double lambda12 = parms[3]; // 计算各变量的导数 NumericVector res(3); res[0] = B*I1 - mu*S - beta*(S*(I1+I2)/N); res[1] = beta*(S*(I1+I2)/N) - B*I1 - lambda12*I1; res[2] = lambda12*I1; return res; }
然后回到R里,用sourceCpp加载这个C++函数:
sourceCpp("sir_model.cpp")
加载完之后,你就能在R里直接调用sir_cpp这个函数了,和调用普通R函数一样。
3. 两种调用C++ ODE求解器的方式
方式一:用deSolve调用C++函数
这种方式最直接,因为你已经熟悉deSolve的lsoda函数了。我们直接用lsoda调用刚才的C++版ODE函数,和R版的对比:
# 定义时间序列、参数和初始条件 times = seq(0, 10, by = 1/52) parm = c(B=0.01, mu=0.008, beta=10, lambda12=1) xstart = c(S=900, I1=100, I2=0) # C++版求解 out_cpp = as.data.frame(lsoda(xstart, times, sir_cpp, parm)) # 你的R版求解(补全了代码) modelsir_cpp = function(t,x){ S = x[1] I1 = x[2] I2 = x[3] N=S+I1+I2 with(as.list(parm), { dS=B*I1-mu*S-beta*(S*(I1+I2)/N) dI1=beta*(S*(I1+I2)/N)-B*I1-lambda12*I1 dI2=lambda12*I1 res=c(dS,dI1,dI2) return(res) }) } out_r = as.data.frame(lsoda(xstart, times, modelsir_cpp, parm)) # 验证结果是否一致(因为数值精度问题,设置容忍度) all.equal(out_r, out_cpp, tolerance = 1e-6)
如果返回TRUE,说明两者的结果是一致的,接下来就可以比速度了。
方式二:用RcppODE的原生C++求解器
如果想要更快的速度,可以用RcppODE封装的C求解器(比如CVODE),这种方式完全在C层面完成求解,减少了R和C++之间的来回调用开销。修改刚才的sir_model.cpp,加上下面的代码:
#include <RcppODE.h> using namespace ode; // 定义一个SIR系统类,继承自RcppODE的system类 class SIR : public ode::system { public: NumericVector parms; // 构造函数,传入参数 SIR(NumericVector p) : parms(p) {} // 导数计算函数,必须实现的接口 void derivs(double t, NumericVector& x, NumericVector& dxdt) { double S = x[0]; double I1 = x[1]; double I2 = x[2]; double N = S + I1 + I2; double B = parms[0]; double mu = parms[1]; double beta = parms[2]; double lambda12 = parms[3]; dxdt[0] = B*I1 - mu*S - beta*(S*(I1+I2)/N); dxdt[1] = beta*(S*(I1+I2)/N) - B*I1 - lambda12*I1; dxdt[2] = lambda12*I1; } }; // [[Rcpp::export]] DataFrame solve_sir_cpp(NumericVector xstart, NumericVector times, NumericVector parms) { // 初始化SIR系统 SIR sir(parms); // 使用CVODE求解器 CVODE solver(sir); solver.init(times[0], xstart); int ntimes = times.size(); NumericMatrix res(ntimes, 3); // 遍历所有时间点求解 for(int i=0; i<ntimes; i++){ solver.step(times[i]); res(i,_) = solver.state(); } // 把结果整理成DataFrame返回给R return DataFrame::create( _["time"] = times, _["S"] = res(_,0), _["I1"] = res(_,1), _["I2"] = res(_,2) ); }
然后重新加载C++代码,调用求解:
sourceCpp("sir_model.cpp") out_rcppode = solve_sir_cpp(xstart, times, parm)
4. 速度对比测试
用microbenchmark来跑多次测试,对比三种方式的速度:
library(microbenchmark) bench_result = microbenchmark( R原生版 = lsoda(xstart, times, modelsir_cpp, parm), deSolve调用C++版 = lsoda(xstart, times, sir_cpp, parm), RcppODE原生C++版 = solve_sir_cpp(xstart, times, parm), times = 100 # 跑100次取平均 ) print(bench_result)
你会看到RcppODE原生C++版的速度最快,deSolve调用C++版次之,而R原生版会慢很多——尤其是当你的ODE系统更复杂、求解时间更长的时候,这个差距会被放大好几倍。
几个需要注意的小细节
- 写C++代码时要注意变量类型,比如所有的浮点运算尽量用
double,避免整数除法的坑 - 用
// [[Rcpp::export]]标记的函数才能被R识别,不要漏掉 - 测试结果一致性时,因为数值计算的精度问题,用
all.equal时可以设置tolerance参数,比如1e-6 - 如果你的模型有大量参数,尽量把参数整理成向量传递,不要一个个单独传,这样更高效
内容的提问来源于stack exchange,提问作者Anthony Hauser
相关产品推荐
相关产品推荐

