You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.08 20:05:55