Skip to content

Commit fd240f0

Browse files
author
Your Name
committed
fixes
1 parent a95ca36 commit fd240f0

3 files changed

Lines changed: 27 additions & 108 deletions

File tree

lib/OrdinaryDiffEqLowStorageRK/src/alg_utils.jl

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -43,6 +43,7 @@ alg_order(alg::CKLLSRK65_4M_4R) = 5
4343
alg_order(alg::SHLDDRK_2N) = 4
4444
alg_order(alg::SHLDDRK52) = 2
4545

46+
isfsal(alg::RK46NL) = false
4647
isfsal(alg::ORK256) = false
4748
isfsal(alg::CarpenterKennedy2N54) = false
4849
isfsal(alg::DGLDDRK84_F) = false

lib/OrdinaryDiffEqLowStorageRK/src/low_storage_rk_caches.jl

Lines changed: 17 additions & 39 deletions
Original file line numberDiff line numberDiff line change
@@ -47,6 +47,23 @@ function ORK256ConstantCache(::Type{T}, ::Type{T2}) where {T, T2}
4747
return LowStorageRK2NConstantCache(A2end, B1, B2end, c2end)
4848
end
4949

50+
function RK46NLConstantCache(::Type{T}, ::Type{T2}) where {T, T2}
51+
A2end = (
52+
T(-0.737101392796), T(-1.634740794343), T(-0.74473900378),
53+
T(-1.469897351522), T(-2.813971388035),
54+
)
55+
B1 = T(0.032918605146)
56+
B2end = (
57+
T(0.8232569982), T(0.3815309489), T(0.200092213184),
58+
T(1.718581042715), T(0.27),
59+
)
60+
c2end = (
61+
T2(0.032918605146), T2(0.249351723343), T2(0.466911705055),
62+
T2(0.582030414044), T2(0.847252983783),
63+
)
64+
return LowStorageRK2NConstantCache(A2end, B1, B2end, c2end)
65+
end
66+
5067
function alg_cache(
5168
alg::RK46NL, u, rate_prototype, ::Type{uEltypeNoUnits},
5269
::Type{uBottomEltypeNoUnits}, ::Type{tTypeNoUnits}, uprev, uprev2, f, t,
@@ -56,45 +73,6 @@ function alg_cache(
5673
return RK46NLConstantCache(constvalue(uBottomEltypeNoUnits), constvalue(tTypeNoUnits))
5774
end
5875

59-
struct RK46NLConstantCache{T, T2} <: OrdinaryDiffEqConstantCache
60-
α2::T
61-
α3::T
62-
α4::T
63-
α5::T
64-
α6::T
65-
β1::T
66-
β2::T
67-
β3::T
68-
β4::T
69-
β5::T
70-
β6::T
71-
c2::T2
72-
c3::T2
73-
c4::T2
74-
c5::T2
75-
c6::T2
76-
77-
function RK46NLConstantCache(::Type{T}, ::Type{T2}) where {T, T2}
78-
α2 = T(-0.737101392796)
79-
α3 = T(-1.634740794343)
80-
α4 = T(-0.74473900378)
81-
α5 = T(-1.469897351522)
82-
α6 = T(-2.813971388035)
83-
β1 = T(0.032918605146)
84-
β2 = T(0.8232569982)
85-
β3 = T(0.3815309489)
86-
β4 = T(0.200092213184)
87-
β5 = T(1.718581042715)
88-
β6 = T(0.27)
89-
c2 = T2(0.032918605146)
90-
c3 = T2(0.249351723343)
91-
c4 = T2(0.466911705055)
92-
c5 = T2(0.582030414044)
93-
c6 = T2(0.847252983783)
94-
return new{T, T2}(α2, α3, α4, α5, α6, β1, β2, β3, β4, β5, β6, c2, c3, c4, c5, c6)
95-
end
96-
end
97-
9876
function alg_cache(
9977
alg::RK46NL, u, rate_prototype, ::Type{uEltypeNoUnits},
10078
::Type{uBottomEltypeNoUnits}, ::Type{tTypeNoUnits}, uprev, uprev2, f, t,

lib/OrdinaryDiffEqLowStorageRK/src/low_storage_rk_perform_step.jl

Lines changed: 9 additions & 69 deletions
Original file line numberDiff line numberDiff line change
@@ -922,81 +922,21 @@ end
922922
@muladd function perform_step!(integrator, cache::RK46NLCache, repeat_step = false)
923923
(; t, dt, uprev, u, f, p) = integrator
924924
(; k, fsalfirst, tmp, stage_limiter!, step_limiter!, thread) = cache
925-
(; α2, α3, α4, α5, α6, β1, β2, β3, β4, β5, β6, c2, c3, c4, c5, c6) = cache.tab
925+
(; A2end, B1, B2end, c2end) = cache.tab
926926

927-
# u1
928927
@.. broadcast = false thread = thread tmp = dt * fsalfirst
929-
@.. broadcast = false thread = thread u = uprev + β1 * tmp
930-
stage_limiter!(u, integrator, p, t + c2 * dt)
931-
# u2
932-
f(k, u, p, t + c2 * dt)
933-
@.. broadcast = false thread = thread tmp = α2 * tmp + dt * k
934-
@.. broadcast = false thread = thread u = u + β2 * tmp
935-
stage_limiter!(u, integrator, p, t + c3 * dt)
936-
# u3
937-
f(k, u, p, t + c3 * dt)
938-
@.. broadcast = false thread = thread tmp = α3 * tmp + dt * k
939-
@.. broadcast = false thread = thread u = u + β3 * tmp
940-
stage_limiter!(u, integrator, p, t + c4 * dt)
941-
# u4
942-
f(k, u, p, t + c4 * dt)
943-
@.. broadcast = false thread = thread tmp = α4 * tmp + dt * k
944-
@.. broadcast = false thread = thread u = u + β4 * tmp
945-
stage_limiter!(u, integrator, p, t + c5 * dt)
946-
# u5 = u
947-
f(k, u, p, t + c5 * dt)
948-
@.. broadcast = false thread = thread tmp = α5 * tmp + dt * k
949-
@.. broadcast = false thread = thread u = u + β5 * tmp
950-
stage_limiter!(u, integrator, p, t + c6 * dt)
951-
952-
f(k, u, p, t + c6 * dt)
953-
@.. broadcast = false thread = thread tmp = α6 * tmp + dt * k
954-
@.. broadcast = false thread = thread u = u + β6 * tmp
928+
@.. broadcast = false thread = thread u = uprev + B1 * tmp
929+
for i in eachindex(A2end)
930+
stage_limiter!(u, integrator, p, t + c2end[i] * dt)
931+
f(k, u, p, t + c2end[i] * dt)
932+
OrdinaryDiffEqCore.increment_nf!(integrator.stats, 1)
933+
@.. broadcast = false thread = thread tmp = A2end[i] * tmp + dt * k
934+
@.. broadcast = false thread = thread u = u + B2end[i] * tmp
935+
end
955936
stage_limiter!(u, integrator, p, t + dt)
956937
step_limiter!(u, integrator, p, t + dt)
957-
958938
f(k, u, p, t + dt)
959-
OrdinaryDiffEqCore.increment_nf!(integrator.stats, 6)
960-
end
961-
962-
function initialize!(integrator, cache::RK46NLConstantCache)
963-
integrator.fsalfirst = integrator.f(integrator.uprev, integrator.p, integrator.t) # Pre-start fsal
964939
OrdinaryDiffEqCore.increment_nf!(integrator.stats, 1)
965-
integrator.kshortsize = 1
966-
integrator.k = typeof(integrator.k)(undef, integrator.kshortsize)
967-
968-
# Avoid undefined entries if k is an array of arrays
969-
integrator.fsallast = zero(integrator.fsalfirst)
970-
return integrator.k[1] = integrator.fsalfirst
971-
end
972-
973-
@muladd function perform_step!(integrator, cache::RK46NLConstantCache, repeat_step = false)
974-
(; t, dt, uprev, u, f, p) = integrator
975-
(; α2, α3, α4, α5, α6, β1, β2, β3, β4, β5, β6, c2, c3, c4, c5, c6) = cache
976-
977-
# u1
978-
tmp = dt * integrator.fsalfirst
979-
u = uprev + β1 * tmp
980-
# u2
981-
tmp = α2 * tmp + dt * f(u, p, t + c2 * dt)
982-
u = u + β2 * tmp
983-
# u3
984-
tmp = α3 * tmp + dt * f(u, p, t + c3 * dt)
985-
u = u + β3 * tmp
986-
# u4
987-
tmp = α4 * tmp + dt * f(u, p, t + c4 * dt)
988-
u = u + β4 * tmp
989-
# u5 = u
990-
tmp = α5 * tmp + dt * f(u, p, t + c5 * dt)
991-
u = u + β5 * tmp
992-
# u6
993-
tmp = α6 * tmp + dt * f(u, p, t + c6 * dt)
994-
u = u + β6 * tmp
995-
996-
integrator.fsallast = f(u, p, t + dt) # For interpolation, then FSAL'd
997-
OrdinaryDiffEqCore.increment_nf!(integrator.stats, 6)
998-
integrator.k[1] = integrator.fsalfirst
999-
integrator.u = u
1000940
end
1001941

1002942
function initialize!(integrator, cache::SHLDDRK52ConstantCache)

0 commit comments

Comments
 (0)