Skip to content

Commit 48a43af

Browse files
author
Your Name
committed
LowStorageRK: unify 2RP/3RP/4RP/5RP families on generic perform_step
1 parent e574b7c commit 48a43af

2 files changed

Lines changed: 373 additions & 740 deletions

File tree

lib/OrdinaryDiffEqLowStorageRK/src/generic_low_storage_perform_step.jl

Lines changed: 351 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -292,3 +292,354 @@ end
292292
end
293293
return nothing
294294
end
295+
296+
@muladd function _perform_step_oop!(integrator, tab::LowStorageRK2RPConstantCache)
297+
(; t, dt, u, uprev, f, fsalfirst, p) = integrator
298+
(; Aᵢ, Bₗ, B̂ₗ, Bᵢ, B̂ᵢ, Cᵢ) = tab
299+
300+
k = fsalfirst
301+
tmp = uprev
302+
integrator.opts.adaptive && (tmp = zero(uprev))
303+
304+
for i in eachindex(Aᵢ)
305+
integrator.opts.adaptive && (tmp = tmp + (Bᵢ[i] - B̂ᵢ[i]) * dt * k)
306+
gprev = u + Aᵢ[i] * dt * k
307+
u = u + Bᵢ[i] * dt * k
308+
k = f(gprev, p, t + Cᵢ[i] * dt)
309+
OrdinaryDiffEqCore.increment_nf!(integrator.stats, 1)
310+
end
311+
312+
integrator.opts.adaptive && (tmp = tmp + (Bₗ - B̂ₗ) * dt * k)
313+
u = u + Bₗ * dt * k
314+
315+
if integrator.opts.adaptive
316+
atmp = calculate_residuals(
317+
tmp, uprev, u, integrator.opts.abstol,
318+
integrator.opts.reltol, integrator.opts.internalnorm, t
319+
)
320+
OrdinaryDiffEqCore.set_EEst!(integrator, integrator.opts.internalnorm(atmp, t))
321+
end
322+
323+
integrator.k[1] = integrator.fsalfirst
324+
integrator.fsallast = f(u, p, t + dt)
325+
OrdinaryDiffEqCore.increment_nf!(integrator.stats, 1)
326+
integrator.u = u
327+
return nothing
328+
end
329+
330+
@muladd function _perform_step_iip!(integrator, cache, tab::LowStorageRK2RPConstantCache)
331+
(; t, dt, u, uprev, f, fsalfirst, p) = integrator
332+
(; k, gprev, tmp, atmp, stage_limiter!, step_limiter!, thread) = cache
333+
(; Aᵢ, Bₗ, B̂ₗ, Bᵢ, B̂ᵢ, Cᵢ) = tab
334+
335+
@.. broadcast = false thread = thread k = fsalfirst
336+
integrator.opts.adaptive && (@.. broadcast = false tmp = zero(uprev))
337+
338+
for i in eachindex(Aᵢ)
339+
integrator.opts.adaptive &&
340+
(@.. broadcast = false thread = thread tmp = tmp + (Bᵢ[i] - B̂ᵢ[i]) * dt * k)
341+
@.. broadcast = false thread = thread gprev = u + Aᵢ[i] * dt * k
342+
@.. broadcast = false thread = thread u = u + Bᵢ[i] * dt * k
343+
f(k, gprev, p, t + Cᵢ[i] * dt)
344+
OrdinaryDiffEqCore.increment_nf!(integrator.stats, 1)
345+
end
346+
347+
integrator.opts.adaptive &&
348+
(@.. broadcast = false thread = thread tmp = tmp + (Bₗ - B̂ₗ) * dt * k)
349+
@.. broadcast = false thread = thread u = u + Bₗ * dt * k
350+
351+
if integrator.opts.adaptive
352+
calculate_residuals!(
353+
atmp, tmp, uprev, u, integrator.opts.abstol,
354+
integrator.opts.reltol, integrator.opts.internalnorm, t,
355+
thread
356+
)
357+
OrdinaryDiffEqCore.set_EEst!(integrator, integrator.opts.internalnorm(atmp, t))
358+
end
359+
360+
step_limiter!(u, integrator, p, t + dt)
361+
f(k, u, p, t + dt)
362+
OrdinaryDiffEqCore.increment_nf!(integrator.stats, 1)
363+
return nothing
364+
end
365+
366+
@muladd function _perform_step_oop!(integrator, tab::LowStorageRK3RPConstantCache)
367+
(; t, dt, u, uprev, f, fsalfirst, p) = integrator
368+
(; Aᵢ₁, Aᵢ₂, Bₗ, B̂ₗ, Bᵢ, B̂ᵢ, Cᵢ) = tab
369+
370+
fᵢ₋₂ = zero(fsalfirst)
371+
k = fsalfirst
372+
uᵢ₋₁ = uprev
373+
uᵢ₋₂ = uprev
374+
tmp = uprev
375+
integrator.opts.adaptive && (tmp = zero(uprev))
376+
377+
for i in eachindex(Aᵢ₁)
378+
integrator.opts.adaptive && (tmp = tmp + (Bᵢ[i] - B̂ᵢ[i]) * dt * k)
379+
gprev = uᵢ₋₂ + (Aᵢ₁[i] * k + Aᵢ₂[i] * fᵢ₋₂) * dt
380+
u = u + Bᵢ[i] * dt * k
381+
fᵢ₋₂ = k
382+
uᵢ₋₂ = uᵢ₋₁
383+
uᵢ₋₁ = u
384+
k = f(gprev, p, t + Cᵢ[i] * dt)
385+
OrdinaryDiffEqCore.increment_nf!(integrator.stats, 1)
386+
end
387+
388+
integrator.opts.adaptive && (tmp = tmp + (Bₗ - B̂ₗ) * dt * k)
389+
u = u + Bₗ * dt * k
390+
391+
if integrator.opts.adaptive
392+
atmp = calculate_residuals(
393+
tmp, uprev, u, integrator.opts.abstol,
394+
integrator.opts.reltol, integrator.opts.internalnorm, t
395+
)
396+
OrdinaryDiffEqCore.set_EEst!(integrator, integrator.opts.internalnorm(atmp, t))
397+
end
398+
399+
integrator.k[1] = integrator.fsalfirst
400+
integrator.fsallast = f(u, p, t + dt)
401+
OrdinaryDiffEqCore.increment_nf!(integrator.stats, 1)
402+
integrator.u = u
403+
return nothing
404+
end
405+
406+
@muladd function _perform_step_iip!(integrator, cache, tab::LowStorageRK3RPConstantCache)
407+
(; t, dt, u, uprev, f, fsalfirst, p) = integrator
408+
(; k, uᵢ₋₁, uᵢ₋₂, gprev, fᵢ₋₂, tmp, atmp, stage_limiter!, step_limiter!, thread) = cache
409+
(; Aᵢ₁, Aᵢ₂, Bₗ, B̂ₗ, Bᵢ, B̂ᵢ, Cᵢ) = tab
410+
411+
@.. broadcast = false thread = thread fᵢ₋₂ = zero(fsalfirst)
412+
@.. broadcast = false thread = thread k = fsalfirst
413+
integrator.opts.adaptive && (@.. broadcast = false thread = thread tmp = zero(uprev))
414+
@.. broadcast = false thread = thread uᵢ₋₁ = uprev
415+
@.. broadcast = false thread = thread uᵢ₋₂ = uprev
416+
417+
for i in eachindex(Aᵢ₁)
418+
integrator.opts.adaptive &&
419+
(@.. broadcast = false thread = thread tmp = tmp + (Bᵢ[i] - B̂ᵢ[i]) * dt * k)
420+
@.. broadcast = false thread = thread gprev = uᵢ₋₂ + (Aᵢ₁[i] * k + Aᵢ₂[i] * fᵢ₋₂) * dt
421+
@.. broadcast = false thread = thread u = u + Bᵢ[i] * dt * k
422+
@.. broadcast = false thread = thread fᵢ₋₂ = k
423+
@.. broadcast = false thread = thread uᵢ₋₂ = uᵢ₋₁
424+
@.. broadcast = false thread = thread uᵢ₋₁ = u
425+
f(k, gprev, p, t + Cᵢ[i] * dt)
426+
OrdinaryDiffEqCore.increment_nf!(integrator.stats, 1)
427+
end
428+
429+
integrator.opts.adaptive &&
430+
(@.. broadcast = false thread = thread tmp = tmp + (Bₗ - B̂ₗ) * dt * k)
431+
@.. broadcast = false thread = thread u = u + Bₗ * dt * k
432+
433+
step_limiter!(u, integrator, p, t + dt)
434+
435+
if integrator.opts.adaptive
436+
calculate_residuals!(
437+
atmp, tmp, uprev, u, integrator.opts.abstol,
438+
integrator.opts.reltol, integrator.opts.internalnorm, t,
439+
thread
440+
)
441+
OrdinaryDiffEqCore.set_EEst!(integrator, integrator.opts.internalnorm(atmp, t))
442+
end
443+
444+
f(k, u, p, t + dt)
445+
OrdinaryDiffEqCore.increment_nf!(integrator.stats, 1)
446+
return nothing
447+
end
448+
449+
@muladd function _perform_step_oop!(integrator, tab::LowStorageRK4RPConstantCache)
450+
(; t, dt, u, uprev, f, fsalfirst, p) = integrator
451+
(; Aᵢ₁, Aᵢ₂, Aᵢ₃, Bₗ, B̂ₗ, Bᵢ, B̂ᵢ, Cᵢ) = tab
452+
453+
fᵢ₋₂ = zero(fsalfirst)
454+
fᵢ₋₃ = zero(fsalfirst)
455+
k = fsalfirst
456+
uᵢ₋₁ = uprev
457+
uᵢ₋₂ = uprev
458+
uᵢ₋₃ = uprev
459+
tmp = uprev
460+
integrator.opts.adaptive && (tmp = zero(uprev))
461+
462+
for i in eachindex(Aᵢ₁)
463+
integrator.opts.adaptive && (tmp = tmp + (Bᵢ[i] - B̂ᵢ[i]) * dt * k)
464+
gprev = uᵢ₋₃ + (Aᵢ₁[i] * k + Aᵢ₂[i] * fᵢ₋₂ + Aᵢ₃[i] * fᵢ₋₃) * dt
465+
u = u + Bᵢ[i] * dt * k
466+
fᵢ₋₃ = fᵢ₋₂
467+
fᵢ₋₂ = k
468+
uᵢ₋₃ = uᵢ₋₂
469+
uᵢ₋₂ = uᵢ₋₁
470+
uᵢ₋₁ = u
471+
k = f(gprev, p, t + Cᵢ[i] * dt)
472+
OrdinaryDiffEqCore.increment_nf!(integrator.stats, 1)
473+
end
474+
475+
integrator.opts.adaptive && (tmp = tmp + (Bₗ - B̂ₗ) * dt * k)
476+
u = u + Bₗ * dt * k
477+
478+
if integrator.opts.adaptive
479+
atmp = calculate_residuals(
480+
tmp, uprev, u, integrator.opts.abstol,
481+
integrator.opts.reltol, integrator.opts.internalnorm, t
482+
)
483+
OrdinaryDiffEqCore.set_EEst!(integrator, integrator.opts.internalnorm(atmp, t))
484+
end
485+
486+
integrator.k[1] = integrator.fsalfirst
487+
integrator.fsallast = f(u, p, t + dt)
488+
OrdinaryDiffEqCore.increment_nf!(integrator.stats, 1)
489+
integrator.u = u
490+
return nothing
491+
end
492+
493+
@muladd function _perform_step_iip!(integrator, cache, tab::LowStorageRK4RPConstantCache)
494+
(; t, dt, u, uprev, f, fsalfirst, p) = integrator
495+
(;
496+
k, uᵢ₋₁, uᵢ₋₂, uᵢ₋₃, gprev, fᵢ₋₂, fᵢ₋₃, tmp, atmp,
497+
stage_limiter!, step_limiter!, thread,
498+
) = cache
499+
(; Aᵢ₁, Aᵢ₂, Aᵢ₃, Bₗ, B̂ₗ, Bᵢ, B̂ᵢ, Cᵢ) = tab
500+
501+
@.. broadcast = false thread = thread fᵢ₋₂ = zero(fsalfirst)
502+
@.. broadcast = false thread = thread fᵢ₋₃ = zero(fsalfirst)
503+
@.. broadcast = false thread = thread k = fsalfirst
504+
integrator.opts.adaptive && (@.. broadcast = false thread = thread tmp = zero(uprev))
505+
@.. broadcast = false thread = thread uᵢ₋₁ = uprev
506+
@.. broadcast = false thread = thread uᵢ₋₂ = uprev
507+
@.. broadcast = false thread = thread uᵢ₋₃ = uprev
508+
509+
for i in eachindex(Aᵢ₁)
510+
integrator.opts.adaptive &&
511+
(@.. broadcast = false thread = thread tmp = tmp + (Bᵢ[i] - B̂ᵢ[i]) * dt * k)
512+
@.. broadcast = false thread = thread gprev = uᵢ₋₃ +
513+
(Aᵢ₁[i] * k + Aᵢ₂[i] * fᵢ₋₂ + Aᵢ₃[i] * fᵢ₋₃) * dt
514+
@.. broadcast = false thread = thread u = u + Bᵢ[i] * dt * k
515+
@.. broadcast = false thread = thread fᵢ₋₃ = fᵢ₋₂
516+
@.. broadcast = false thread = thread fᵢ₋₂ = k
517+
@.. broadcast = false thread = thread uᵢ₋₃ = uᵢ₋₂
518+
@.. broadcast = false thread = thread uᵢ₋₂ = uᵢ₋₁
519+
@.. broadcast = false thread = thread uᵢ₋₁ = u
520+
f(k, gprev, p, t + Cᵢ[i] * dt)
521+
OrdinaryDiffEqCore.increment_nf!(integrator.stats, 1)
522+
end
523+
524+
integrator.opts.adaptive &&
525+
(@.. broadcast = false thread = thread tmp = tmp + (Bₗ - B̂ₗ) * dt * k)
526+
@.. broadcast = false thread = thread u = u + Bₗ * dt * k
527+
528+
step_limiter!(u, integrator, p, t + dt)
529+
530+
if integrator.opts.adaptive
531+
calculate_residuals!(
532+
atmp, tmp, uprev, u, integrator.opts.abstol,
533+
integrator.opts.reltol, integrator.opts.internalnorm, t,
534+
thread
535+
)
536+
OrdinaryDiffEqCore.set_EEst!(integrator, integrator.opts.internalnorm(atmp, t))
537+
end
538+
539+
f(k, u, p, t + dt)
540+
OrdinaryDiffEqCore.increment_nf!(integrator.stats, 1)
541+
return nothing
542+
end
543+
544+
@muladd function _perform_step_oop!(integrator, tab::LowStorageRK5RPConstantCache)
545+
(; t, dt, u, uprev, f, fsalfirst, p) = integrator
546+
(; Aᵢ₁, Aᵢ₂, Aᵢ₃, Aᵢ₄, Bₗ, B̂ₗ, Bᵢ, B̂ᵢ, Cᵢ) = tab
547+
548+
fᵢ₋₂ = zero(fsalfirst)
549+
fᵢ₋₃ = zero(fsalfirst)
550+
fᵢ₋₄ = zero(fsalfirst)
551+
k = fsalfirst
552+
uᵢ₋₁ = uprev
553+
uᵢ₋₂ = uprev
554+
uᵢ₋₃ = uprev
555+
uᵢ₋₄ = uprev
556+
tmp = uprev
557+
integrator.opts.adaptive && (tmp = zero(uprev))
558+
559+
for i in eachindex(Aᵢ₁)
560+
integrator.opts.adaptive && (tmp = tmp + (Bᵢ[i] - B̂ᵢ[i]) * dt * k)
561+
gprev = uᵢ₋₄ + (Aᵢ₁[i] * k + Aᵢ₂[i] * fᵢ₋₂ + Aᵢ₃[i] * fᵢ₋₃ + Aᵢ₄[i] * fᵢ₋₄) * dt
562+
u = u + Bᵢ[i] * dt * k
563+
fᵢ₋₄ = fᵢ₋₃
564+
fᵢ₋₃ = fᵢ₋₂
565+
fᵢ₋₂ = k
566+
uᵢ₋₄ = uᵢ₋₃
567+
uᵢ₋₃ = uᵢ₋₂
568+
uᵢ₋₂ = uᵢ₋₁
569+
uᵢ₋₁ = u
570+
k = f(gprev, p, t + Cᵢ[i] * dt)
571+
OrdinaryDiffEqCore.increment_nf!(integrator.stats, 1)
572+
end
573+
574+
integrator.opts.adaptive && (tmp = tmp + (Bₗ - B̂ₗ) * dt * k)
575+
u = u + Bₗ * dt * k
576+
577+
if integrator.opts.adaptive
578+
atmp = calculate_residuals(
579+
tmp, uprev, u, integrator.opts.abstol,
580+
integrator.opts.reltol, integrator.opts.internalnorm, t
581+
)
582+
OrdinaryDiffEqCore.set_EEst!(integrator, integrator.opts.internalnorm(atmp, t))
583+
end
584+
585+
integrator.k[1] = integrator.fsalfirst
586+
integrator.fsallast = f(u, p, t + dt)
587+
OrdinaryDiffEqCore.increment_nf!(integrator.stats, 1)
588+
integrator.u = u
589+
return nothing
590+
end
591+
592+
@muladd function _perform_step_iip!(integrator, cache, tab::LowStorageRK5RPConstantCache)
593+
(; t, dt, u, uprev, f, fsalfirst, p) = integrator
594+
(;
595+
k, uᵢ₋₁, uᵢ₋₂, uᵢ₋₃, uᵢ₋₄, gprev, fᵢ₋₂, fᵢ₋₃, fᵢ₋₄, tmp,
596+
atmp, stage_limiter!, step_limiter!, thread,
597+
) = cache
598+
(; Aᵢ₁, Aᵢ₂, Aᵢ₃, Aᵢ₄, Bₗ, B̂ₗ, Bᵢ, B̂ᵢ, Cᵢ) = tab
599+
600+
@.. broadcast = false thread = thread fᵢ₋₂ = zero(fsalfirst)
601+
@.. broadcast = false thread = thread fᵢ₋₃ = zero(fsalfirst)
602+
@.. broadcast = false thread = thread fᵢ₋₄ = zero(fsalfirst)
603+
@.. broadcast = false thread = thread k = fsalfirst
604+
integrator.opts.adaptive && (@.. broadcast = false thread = thread tmp = zero(uprev))
605+
@.. broadcast = false thread = thread uᵢ₋₁ = uprev
606+
@.. broadcast = false thread = thread uᵢ₋₂ = uprev
607+
@.. broadcast = false thread = thread uᵢ₋₃ = uprev
608+
@.. broadcast = false thread = thread uᵢ₋₄ = uprev
609+
610+
for i in eachindex(Aᵢ₁)
611+
integrator.opts.adaptive &&
612+
(@.. broadcast = false thread = thread tmp = tmp + (Bᵢ[i] - B̂ᵢ[i]) * dt * k)
613+
@.. broadcast = false thread = thread gprev = uᵢ₋₄ +
614+
(Aᵢ₁[i] * k + Aᵢ₂[i] * fᵢ₋₂ + Aᵢ₃[i] * fᵢ₋₃ + Aᵢ₄[i] * fᵢ₋₄) * dt
615+
@.. broadcast = false thread = thread u = u + Bᵢ[i] * dt * k
616+
@.. broadcast = false thread = thread fᵢ₋₄ = fᵢ₋₃
617+
@.. broadcast = false thread = thread fᵢ₋₃ = fᵢ₋₂
618+
@.. broadcast = false thread = thread fᵢ₋₂ = k
619+
@.. broadcast = false thread = thread uᵢ₋₄ = uᵢ₋₃
620+
@.. broadcast = false thread = thread uᵢ₋₃ = uᵢ₋₂
621+
@.. broadcast = false thread = thread uᵢ₋₂ = uᵢ₋₁
622+
@.. broadcast = false thread = thread uᵢ₋₁ = u
623+
f(k, gprev, p, t + Cᵢ[i] * dt)
624+
OrdinaryDiffEqCore.increment_nf!(integrator.stats, 1)
625+
end
626+
627+
integrator.opts.adaptive &&
628+
(@.. broadcast = false thread = thread tmp = tmp + (Bₗ - B̂ₗ) * dt * k)
629+
@.. broadcast = false thread = thread u = u + Bₗ * dt * k
630+
631+
step_limiter!(u, integrator, p, t + dt)
632+
633+
if integrator.opts.adaptive
634+
calculate_residuals!(
635+
atmp, tmp, uprev, u, integrator.opts.abstol,
636+
integrator.opts.reltol, integrator.opts.internalnorm, t,
637+
thread
638+
)
639+
OrdinaryDiffEqCore.set_EEst!(integrator, integrator.opts.internalnorm(atmp, t))
640+
end
641+
642+
f(k, u, p, t + dt)
643+
OrdinaryDiffEqCore.increment_nf!(integrator.stats, 1)
644+
return nothing
645+
end

0 commit comments

Comments
 (0)