Julia中使用Callbacks求解ODE时仅保留最后一步解的问题求助
Julia ODE求解内存占用过高问题排查与解决
问题场景与代码
我用Julia求解ODE,所用代码如下:
using DelimitedFiles using LinearAlgebra using Random using Distributions using DifferentialEquations using Plots N = 100 σ, a, J, K = 10.0, 1.0,1.0,1.0 ω = zeros(N); function heaviside(x::Real) return x >= 0.0 ? 1.0 : 0.0 end function rhs(du,u,p,t) u1 = @view u[1:N] du1 = @view du[1:N] u2 = @view u[N+1:2*N] du2 = @view du[N+1:2*N] σ,a,J,K=p for i in 1:N du1[i]=(1/N)*sum(j->((u1[j]-u1[i])*(1+J*cos(u2[j]-u2[i]))-sign(u1[j]-u1[i])),1:N) du2[i]=(ω[i]) + (K/N)* sum(j->(sin(u2[j]-u2[i])*(1-(u1[j]-u1[i])^2/σ^2)*heaviside(σ - abs((u1[j]-u1[i])))),1:N) end return du end ti=0 tf=500 tt=0.75*tf tspan = (ti, tf) dts=0.25 vector_t=tt:dts:tf p=[σ,a,J,K] Random.seed!(123) u0= [(rand() * 8.0 - 4.0, rand() * (2.0 * π) - π) for j in 1:N] u0=vcat([x[1] for x in u0], [x[2] for x in u0]); prob = ODEProblem(rhs,u0, tspan,p); # For making mean_u1 zero function condition(u,t,integrator) t == integrator.t end function affect!(integrator) u1_mean = mean(integrator.u[1:N]) integrator.u[1:N] .-= u1_mean end cbd = DiscreteCallback(condition, affect!) ################################## saved_values2= SavedValues(Float64,ComplexF64) function saver2(u,t,integrator) _pp=u[1:N] _qq=u[N+1:2*N] out2= mean((_pp).*exp.((_qq)*1im)) end cb2 = SavingCallback(saver2, saved_values2,saveat=tt:dts:tf) ################################### ################################### saved_values4= SavedValues(Float64,Float64) function saver4(u,t,integrator) duc=rhs(zeros(size(u)),u,integrator.p,t) _pp=duc[1:N] _qq=duc[N+1:2*N] out4= mean(sqrt.(((_pp).^2)+((_qq).^2))) end cb4 = SavingCallback(saver4, saved_values4,saveat=tt:dts:tf) ##################################### cbs = CallbackSet(cbd, cb2,cb4); @time sol= solve(prob, Tsit5(),reltol=1e-6,maxiters=1e20, callback = cbs,saveat=[tf]); sizeof(sol) saved_values4.saveval print(mean(saved_values4.saveval)); Z1=saved_values2.saveval; vel=saved_values4.saveval;
在上述代码中,我通过DiscreteCallback在每个积分步骤后将u1变量的均值置零,并使用SavingCallback在时间范围tt:dts:tf内存储两个指标Z1和vel,因此在SavingCallback中设置了saveat=tt:dts:tf。
我不需要求解器返回的完整sol结果,仅需要最后时刻tf的解,因此在solve函数中设置了saveat=[tf]。但即使如此,求解过程仍保存了所有步骤的解,导致sol占用大量内存、机器运行缓慢。请问这是什么原因,该如何解决?
问题原因
- 核心问题出在
DiscreteCallback的设置上:你的condition(u,t,integrator) = t == integrator.t条件会在每个积分步都触发回调,而DiscreteCallback默认会保存回调触发时的状态,这就导致求解器不断存储所有中间步骤的解,完全忽略了你设置的saveat=[tf]。 SavingCallback的存储是独立在saved_values对象中的,不会导致主解对象sol膨胀,无需针对它做修改。
解决方法
1. 阻止DiscreteCallback保存中间状态
给DiscreteCallback添加save_positions=(false,false)参数,明确告诉求解器不要保存回调触发时的状态:
cbd = DiscreteCallback(condition, affect!; save_positions=(false, false))
这会直接切断中间状态的存储,大幅降低sol的内存占用。
2. 优化回调触发条件(可选)
你的condition函数可以简化为直接返回true(因为你需要每个积分步都触发),逻辑更清晰:
function condition(u,t,integrator) return true end
3. 双重保证只保存最终状态
配合设置save_everystep=false,进一步确保求解器不会保存任何额外的中间步骤:
@time sol= solve(prob, Tsit5(), reltol=1e-6, maxiters=1e20, callback = cbs, saveat=[tf], save_everystep=false);
内容的提问来源于stack exchange,提问作者Sayeed
相关产品推荐
相关产品推荐

