如何在DifferentialEquations.jl的EnsembleProblem中变分tspan并GPU并行?
带tspan变分的GPU并行ODE集成求解方案
问题背景
当前需要对一类由sys!定义的常微分方程(ODE)进行多参数变分的集成求解,已实现w、plateu_cycle、px、py的并行变分,但还需对tspan(重点是tend)进行变分。由于EnsembleProblem不支持直接变分tspan,希望借助DiffEqGPU.jl在GPU上完成所有集成问题的并行执行。
解决方案
核心思路是将tspan的信息融入参数集合,通过EnsembleProblem的prob_func动态生成每个子问题的时间区间,从而实现tspan与其他参数的联合变分,同时适配GPU并行。
步骤说明
- 构造包含
tend的参数列表:把每个子问题需要的w、plateu_cycle、px、py和tend打包成参数元组,生成所有待求解的参数组合。 - 自定义
prob_func生成子问题:在EnsembleProblem中,通过回调函数根据每个子问题的参数动态设置对应的tspan,生成新的ODEProblem实例。 - 配置GPU并行求解器:使用
DiffEqGPU提供的EnsembleGPUArray(或EnsembleGPUKernel)作为集成执行器,搭配适合GPU的ODE求解器。
完整代码示例
using DifferentialEquations, PhysicalConstants.CODATA2018, DiffEqGPU, CUDA # 原ODE定义部分保持不变 e::Float64 = ElementaryCharge.val function E_0(w::Float64) return w/e end function E(px::Float64, py::Float64) return sqrt(1+px^2+py^2) end function Vectorpotential(t, w::Float64, plateu_cycle::Float64) return -E_0(w)/w * sin.(t) * sineSquaredWindow(t,w,plateu_cycle) end function sineSquaredWindow(t, w::Float64, plateu_cycle::Float64) T = 2pi/w if t <= 0 return 0 elseif t/w >= 0 && (t/w <= T/2) return sin(t/2)^2 elseif (t / w >= T / 2) && (t / w <= T * (plateu_cycle + 1 / 2)) return 1 elseif (t / w >= T * (plateu_cycle + 1 / 2)) & (t / w <= T * (plateu_cycle + 1)) return 1-sin(w/2 * (t/w - T*(plateu_cycle+1/2)))^2 end return 0 end function sys!(du, u, params, t) w, plateu_cycle, px, py = params[1:4] # 前4个为原参数 Energy = E(px,py) function nu(t, plateu_cycle::Float64, w::Float64, px::Float64, py::Float64) return -1im * e * Vectorpotential(w*t,w,plateu_cycle)*exp(2im*Energy*t) * ( px*py /(Energy*(Energy+1))+1im*(1-py^2 / (Energy*(Energy+1)))) end function kappa(t, plateu_cycle::Float64, w::Float64, py::Float64) return 1im*e*Vectorpotential(w*t,w,plateu_cycle)*py/Energy end du[1] = kappa(t,plateu_cycle,w,py)*u[1] + nu(t, plateu_cycle,w, px,py) *u[2] du[2] = -conj(nu(t,plateu_cycle,w,px,py))*u[1] + conj(kappa(t,plateu_cycle,w, py))*u[2] return nothing end # 1. 构造包含tend的多组参数组合示例 param_list = [ (0.84, 8.0, 0.0, 0.0, 2pi/0.84*(8.0+1)+1), (0.84, 9.0, 0.0, 0.0, 2pi/0.84*(9.0+1)+1), (0.90, 8.0, 0.1, 0.0, 2pi/0.90*(8.0+1)+1), # 可根据需求添加更多参数组合 ] # 2. 创建基础问题(tspan为占位符,后续在prob_func中替换) f0=complex(0.0,0.0) g0=complex(1.0,0.0) init_cond = [f0, g0] base_prob = ODEProblem(sys!, init_cond, (0.0, 1.0), param_list[1]) # 3. 自定义prob_func:动态生成对应tspan的子问题 function prob_func(prob, i, repeat) params = param_list[i] tstart = 0.0 tend = params[5] new_tspan = (tstart, tend) # 重写问题:替换tspan、参数和saveat remake(prob, tspan=new_tspan, p=params[1:4], saveat=collect(range(tstart, tend, length=1000))) end # 4. 创建EnsembleProblem并配置GPU并行 ensemble_prob = EnsembleProblem(base_prob, prob_func=prob_func) # 5. GPU上并行求解 sol = solve(ensemble_prob, Tsit5(), EnsembleGPUArray(CUDA.CUDABackend()), trajectories=length(param_list), batch_size=length(param_list)) # 访问结果:sol[i]对应第i个参数组合的求解结果 println("第1个问题的时间区间:", sol[1].tspan)
关键说明
prob_func是核心:它允许为每个集成轨迹动态修改tspan、参数、初始条件等,完美适配tspan变分需求。- GPU适配:
EnsembleGPUArray会自动将任务分发到GPU核心并行执行,需确保CUDA环境配置正确。 - 求解器选择:大部分自适应求解器(如
Tsit5()、Vern7())均兼容DiffEqGPU,追求更高性能可尝试GPUODESolver系列。
内容的提问来源于stack exchange,提问作者derdotte
相关产品推荐
相关产品推荐

