Skip to content

Commit 931262d

Browse files
addsteps optimization
1 parent 126e7cf commit 931262d

1 file changed

Lines changed: 11 additions & 4 deletions

File tree

lib/OrdinaryDiffEqCore/src/disco.jl

Lines changed: 11 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,7 @@ function find_discontinuity(u, uprev, integrator)
2222
breakpointθ = -one(dt)
2323
disco_probs = get_disco_probs(integrator.controller_cache)
2424
idx = 1
25+
addsteps_called = false
2526
for i in cb.continuous_callbacks
2627
if (!(i.maybe_discontinuity))
2728
continue
@@ -38,10 +39,13 @@ function find_discontinuity(u, uprev, integrator)
3839
len_cb = i.len
3940
i.condition(disco_zero.out_low, uprev, t, integrator)
4041
i.condition(disco_zero.out_high, u, t + dt, integrator)
41-
_ode_addsteps!(disco_zero.k, disco_zero.tprev, disco_zero.uprev, disco_zero.u,
42-
disco_zero.dt, disco_zero.f, disco_zero.p, disco_zero.cache, false, true, false)
4342
for j in 1:len_cb
4443
if (disco_zero.out_low[j] * disco_zero.out_high[j] < zero(disco_zero.out_low[j]))
44+
if (!addsteps_called)
45+
addsteps_called = true
46+
_ode_addsteps!(disco_zero.k, disco_zero.tprev, disco_zero.uprev, disco_zero.u,
47+
disco_zero.dt, disco_zero.f, disco_zero.p, disco_zero.cache, false, true, false)
48+
end
4549
disco_zero.ind = j
4650
sol = solve(disco_prob, tspan = (0.0, breakpointθ == -one(dt) ? 1.0 : breakpointθ); save_everystep = false)
4751
tmp = sol[]
@@ -56,8 +60,11 @@ function find_discontinuity(u, uprev, integrator)
5660
out_prev = i.condition(uprev, t, integrator)
5761
out_curr = i.condition(u, t + dt, integrator)
5862
if (out_prev * out_curr < zero(out_prev))
59-
_ode_addsteps!(disco_zero.k, disco_zero.tprev, disco_zero.uprev, disco_zero.u,
60-
disco_zero.dt, disco_zero.f, disco_zero.p, disco_zero.cache, false, true, false)
63+
if (!addsteps_called)
64+
addsteps_called = true
65+
_ode_addsteps!(disco_zero.k, disco_zero.tprev, disco_zero.uprev, disco_zero.u,
66+
disco_zero.dt, disco_zero.f, disco_zero.p, disco_zero.cache, false, true, false)
67+
end
6168
sol = solve(disco_prob, tspan = (0.0, breakpointθ == -one(dt) ? 1.0 : breakpointθ); save_everystep = false)
6269
tmp = sol[]
6370
if (!isnan(tmp) && (breakpointθ < zero(breakpointθ) || tmp < breakpointθ))

0 commit comments

Comments
 (0)