Skip to content

Commit cb4a3f9

Browse files
Split Julia 1 MIRK AD shards by backend
Co-Authored-By: Chris Rackauckas <accounts@chrisrackauckas.com>
1 parent c2b7485 commit cb4a3f9

3 files changed

Lines changed: 228 additions & 46 deletions

File tree

lib/BoundaryValueDiffEqMIRK/test/AD/ad_tests.jl

Lines changed: 88 additions & 43 deletions
Original file line numberDiff line numberDiff line change
@@ -3,10 +3,11 @@ using Test
33

44
mirk_ad_group() = get(ENV, "BOUNDARYVALUEDIFFEQ_MIRK_AD_GROUP", "ALL")
55
run_mirk_ad_group(group) = mirk_ad_group() in ("ALL", group)
6+
mirk_ad_backend() = get(ENV, "BOUNDARYVALUEDIFFEQ_MIRK_AD_BACKEND", "ALL")
7+
run_mirk_ad_backend(backend) = mirk_ad_backend() in ("ALL", backend)
68

79
@testset "Different AD compatibility" begin
810
using BoundaryValueDiffEqMIRK
9-
using ForwardDiff, Enzyme, Mooncake
1011

1112
if run_mirk_ad_group("MULTIPOINT_GRID")
1213
@testset "Test different AD on multipoint BVP" begin
@@ -23,22 +24,38 @@ run_mirk_ad_group(group) = mirk_ad_group() in ("ALL", group)
2324
u0 = [pi / 2, pi / 2]
2425
tspan = (0.0, pi / 2)
2526
prob = BVProblem(simplependulum!, bc!, u0, tspan)
26-
jac_alg_forwarddiff = BVPJacobianAlgorithm(
27-
bc_diffmode = AutoSparse(AutoForwardDiff()), nonbc_diffmode = AutoForwardDiff()
28-
)
29-
jac_alg_enzyme = BVPJacobianAlgorithm(
30-
bc_diffmode = AutoSparse(
31-
AutoEnzyme(
32-
mode = Enzyme.Reverse, function_annotation = Enzyme.Duplicated
27+
28+
if run_mirk_ad_backend("FORWARDDIFF")
29+
using ForwardDiff
30+
jac_alg = BVPJacobianAlgorithm(
31+
bc_diffmode = AutoSparse(AutoForwardDiff()), nonbc_diffmode = AutoForwardDiff()
32+
)
33+
sol = solve(prob, MIRK4(; jac_alg = jac_alg), dt = 0.05)
34+
@test SciMLBase.successful_retcode(sol)
35+
end
36+
if run_mirk_ad_backend("ENZYME")
37+
using Enzyme
38+
jac_alg = BVPJacobianAlgorithm(
39+
bc_diffmode = AutoSparse(
40+
AutoEnzyme(
41+
mode = Enzyme.Reverse, function_annotation = Enzyme.Duplicated
42+
)
43+
),
44+
nonbc_diffmode = AutoEnzyme(
45+
mode = Enzyme.Forward, function_annotation = Enzyme.Duplicated
3346
)
34-
),
35-
nonbc_diffmode = AutoEnzyme(mode = Enzyme.Forward, function_annotation = Enzyme.Duplicated)
36-
)
37-
jac_alg_mooncake = BVPJacobianAlgorithm(
38-
bc_diffmode = AutoSparse(AutoMooncake(; config = nothing)),
39-
nonbc_diffmode = AutoEnzyme(mode = Enzyme.Forward, function_annotation = Enzyme.Duplicated)
40-
)
41-
for jac_alg in [jac_alg_forwarddiff, jac_alg_enzyme, jac_alg_mooncake]
47+
)
48+
sol = solve(prob, MIRK4(; jac_alg = jac_alg), dt = 0.05)
49+
@test SciMLBase.successful_retcode(sol)
50+
end
51+
if run_mirk_ad_backend("MOONCAKE")
52+
using Enzyme, Mooncake
53+
jac_alg = BVPJacobianAlgorithm(
54+
bc_diffmode = AutoSparse(AutoMooncake(; config = nothing)),
55+
nonbc_diffmode = AutoEnzyme(
56+
mode = Enzyme.Forward, function_annotation = Enzyme.Duplicated
57+
)
58+
)
4259
sol = solve(prob, MIRK4(; jac_alg = jac_alg), dt = 0.05)
4360
@test SciMLBase.successful_retcode(sol)
4461
end
@@ -60,22 +77,38 @@ run_mirk_ad_group(group) = mirk_ad_group() in ("ALL", group)
6077
u0 = [pi / 2, pi / 2]
6178
tspan = (0.0, pi / 2)
6279
prob = BVProblem(simplependulum!, bc!, u0, tspan)
63-
jac_alg_forwarddiff = BVPJacobianAlgorithm(
64-
bc_diffmode = AutoSparse(AutoForwardDiff()), nonbc_diffmode = AutoForwardDiff()
65-
)
66-
jac_alg_enzyme = BVPJacobianAlgorithm(
67-
bc_diffmode = AutoSparse(
68-
AutoEnzyme(
69-
mode = Enzyme.Reverse, function_annotation = Enzyme.Duplicated
80+
81+
if run_mirk_ad_backend("FORWARDDIFF")
82+
using ForwardDiff
83+
jac_alg = BVPJacobianAlgorithm(
84+
bc_diffmode = AutoSparse(AutoForwardDiff()), nonbc_diffmode = AutoForwardDiff()
85+
)
86+
sol = solve(prob, MIRK4(; jac_alg = jac_alg), dt = 0.05)
87+
@test SciMLBase.successful_retcode(sol)
88+
end
89+
if run_mirk_ad_backend("ENZYME")
90+
using Enzyme
91+
jac_alg = BVPJacobianAlgorithm(
92+
bc_diffmode = AutoSparse(
93+
AutoEnzyme(
94+
mode = Enzyme.Reverse, function_annotation = Enzyme.Duplicated
95+
)
96+
),
97+
nonbc_diffmode = AutoEnzyme(
98+
mode = Enzyme.Forward, function_annotation = Enzyme.Duplicated
7099
)
71-
),
72-
nonbc_diffmode = AutoEnzyme(mode = Enzyme.Forward, function_annotation = Enzyme.Duplicated)
73-
)
74-
jac_alg_mooncake = BVPJacobianAlgorithm(
75-
bc_diffmode = AutoSparse(AutoMooncake(; config = nothing)),
76-
nonbc_diffmode = AutoEnzyme(mode = Enzyme.Forward, function_annotation = Enzyme.Duplicated)
77-
)
78-
for jac_alg in [jac_alg_forwarddiff, jac_alg_enzyme, jac_alg_mooncake]
100+
)
101+
sol = solve(prob, MIRK4(; jac_alg = jac_alg), dt = 0.05)
102+
@test SciMLBase.successful_retcode(sol)
103+
end
104+
if run_mirk_ad_backend("MOONCAKE")
105+
using Enzyme, Mooncake
106+
jac_alg = BVPJacobianAlgorithm(
107+
bc_diffmode = AutoSparse(AutoMooncake(; config = nothing)),
108+
nonbc_diffmode = AutoEnzyme(
109+
mode = Enzyme.Forward, function_annotation = Enzyme.Duplicated
110+
)
111+
)
79112
sol = solve(prob, MIRK4(; jac_alg = jac_alg), dt = 0.05)
80113
@test SciMLBase.successful_retcode(sol)
81114
end
@@ -103,22 +136,34 @@ run_mirk_ad_group(group) = mirk_ad_group() in ("ALL", group)
103136
odef!, (boundary_two_point_a!, boundary_two_point_b!),
104137
u0, tspan; bcresid_prototype, nlls = Val(false)
105138
)
106-
jac_alg_forwarddiff = BVPJacobianAlgorithm(AutoSparse(AutoForwardDiff()))
107-
jac_alg_enzyme = BVPJacobianAlgorithm(
108-
AutoSparse(
109-
AutoEnzyme(
110-
mode = Enzyme.Forward, function_annotation = Enzyme.Duplicated
139+
140+
if run_mirk_ad_backend("FORWARDDIFF")
141+
using ForwardDiff
142+
jac_alg = BVPJacobianAlgorithm(AutoSparse(AutoForwardDiff()))
143+
sol = solve(prob, MIRK4(; jac_alg = jac_alg), dt = 0.01)
144+
@test SciMLBase.successful_retcode(sol)
145+
end
146+
if run_mirk_ad_backend("ENZYME")
147+
using Enzyme
148+
jac_alg = BVPJacobianAlgorithm(
149+
AutoSparse(
150+
AutoEnzyme(
151+
mode = Enzyme.Forward, function_annotation = Enzyme.Duplicated
152+
)
111153
)
112154
)
113-
)
114-
jac_alg_mooncake = BVPJacobianAlgorithm(
115-
AutoSparse(
116-
AutoMooncake(;
117-
config = nothing
155+
sol = solve(prob, MIRK4(; jac_alg = jac_alg), dt = 0.01)
156+
@test SciMLBase.successful_retcode(sol)
157+
end
158+
if run_mirk_ad_backend("MOONCAKE")
159+
using Mooncake
160+
jac_alg = BVPJacobianAlgorithm(
161+
AutoSparse(
162+
AutoMooncake(;
163+
config = nothing
164+
)
118165
)
119166
)
120-
)
121-
for jac_alg in [jac_alg_forwarddiff, jac_alg_enzyme, jac_alg_mooncake]
122167
sol = solve(prob, MIRK4(; jac_alg = jac_alg), dt = 0.01)
123168
@test SciMLBase.successful_retcode(sol)
124169
end

lib/BoundaryValueDiffEqMIRK/test/runtests.jl

Lines changed: 108 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -9,16 +9,88 @@ mirk_ad_multipoint_grid() = @time @safetestset "MIRK AD Multipoint Grid Tests" b
99
include("AD/ad_tests.jl")
1010
end
1111
end
12+
mirk_ad_multipoint_grid_forwarddiff() = @time @safetestset "MIRK AD Multipoint Grid ForwardDiff Tests" begin
13+
withenv(
14+
"BOUNDARYVALUEDIFFEQ_MIRK_AD_GROUP" => "MULTIPOINT_GRID",
15+
"BOUNDARYVALUEDIFFEQ_MIRK_AD_BACKEND" => "FORWARDDIFF"
16+
) do
17+
include("AD/ad_tests.jl")
18+
end
19+
end
20+
mirk_ad_multipoint_grid_enzyme() = @time @safetestset "MIRK AD Multipoint Grid Enzyme Tests" begin
21+
withenv(
22+
"BOUNDARYVALUEDIFFEQ_MIRK_AD_GROUP" => "MULTIPOINT_GRID",
23+
"BOUNDARYVALUEDIFFEQ_MIRK_AD_BACKEND" => "ENZYME"
24+
) do
25+
include("AD/ad_tests.jl")
26+
end
27+
end
28+
mirk_ad_multipoint_grid_mooncake() = @time @safetestset "MIRK AD Multipoint Grid Mooncake Tests" begin
29+
withenv(
30+
"BOUNDARYVALUEDIFFEQ_MIRK_AD_GROUP" => "MULTIPOINT_GRID",
31+
"BOUNDARYVALUEDIFFEQ_MIRK_AD_BACKEND" => "MOONCAKE"
32+
) do
33+
include("AD/ad_tests.jl")
34+
end
35+
end
1236
mirk_ad_multipoint_interpolation() = @time @safetestset "MIRK AD Multipoint Interpolation Tests" begin
1337
withenv("BOUNDARYVALUEDIFFEQ_MIRK_AD_GROUP" => "MULTIPOINT_INTERPOLATION") do
1438
include("AD/ad_tests.jl")
1539
end
1640
end
41+
mirk_ad_multipoint_interpolation_forwarddiff() = @time @safetestset "MIRK AD Multipoint Interpolation ForwardDiff Tests" begin
42+
withenv(
43+
"BOUNDARYVALUEDIFFEQ_MIRK_AD_GROUP" => "MULTIPOINT_INTERPOLATION",
44+
"BOUNDARYVALUEDIFFEQ_MIRK_AD_BACKEND" => "FORWARDDIFF"
45+
) do
46+
include("AD/ad_tests.jl")
47+
end
48+
end
49+
mirk_ad_multipoint_interpolation_enzyme() = @time @safetestset "MIRK AD Multipoint Interpolation Enzyme Tests" begin
50+
withenv(
51+
"BOUNDARYVALUEDIFFEQ_MIRK_AD_GROUP" => "MULTIPOINT_INTERPOLATION",
52+
"BOUNDARYVALUEDIFFEQ_MIRK_AD_BACKEND" => "ENZYME"
53+
) do
54+
include("AD/ad_tests.jl")
55+
end
56+
end
57+
mirk_ad_multipoint_interpolation_mooncake() = @time @safetestset "MIRK AD Multipoint Interpolation Mooncake Tests" begin
58+
withenv(
59+
"BOUNDARYVALUEDIFFEQ_MIRK_AD_GROUP" => "MULTIPOINT_INTERPOLATION",
60+
"BOUNDARYVALUEDIFFEQ_MIRK_AD_BACKEND" => "MOONCAKE"
61+
) do
62+
include("AD/ad_tests.jl")
63+
end
64+
end
1765
mirk_ad_twopoint() = @time @safetestset "MIRK AD TwoPoint Tests" begin
1866
withenv("BOUNDARYVALUEDIFFEQ_MIRK_AD_GROUP" => "TWOPOINT") do
1967
include("AD/ad_tests.jl")
2068
end
2169
end
70+
mirk_ad_twopoint_forwarddiff() = @time @safetestset "MIRK AD TwoPoint ForwardDiff Tests" begin
71+
withenv(
72+
"BOUNDARYVALUEDIFFEQ_MIRK_AD_GROUP" => "TWOPOINT",
73+
"BOUNDARYVALUEDIFFEQ_MIRK_AD_BACKEND" => "FORWARDDIFF"
74+
) do
75+
include("AD/ad_tests.jl")
76+
end
77+
end
78+
mirk_ad_twopoint_enzyme() = @time @safetestset "MIRK AD TwoPoint Enzyme Tests" begin
79+
withenv(
80+
"BOUNDARYVALUEDIFFEQ_MIRK_AD_GROUP" => "TWOPOINT",
81+
"BOUNDARYVALUEDIFFEQ_MIRK_AD_BACKEND" => "ENZYME"
82+
) do
83+
include("AD/ad_tests.jl")
84+
end
85+
end
86+
mirk_ad_twopoint_mooncake() = @time @safetestset "MIRK AD TwoPoint Mooncake Tests" begin
87+
withenv(
88+
"BOUNDARYVALUEDIFFEQ_MIRK_AD_GROUP" => "TWOPOINT",
89+
"BOUNDARYVALUEDIFFEQ_MIRK_AD_BACKEND" => "MOONCAKE"
90+
) do
91+
include("AD/ad_tests.jl")
92+
end
93+
end
2294

2395
run_tests(;
2496
env = "BOUNDARYVALUEDIFFEQ_TEST_GROUP",
@@ -44,14 +116,50 @@ run_tests(;
44116
env = MIRK_AD_ENV,
45117
body = mirk_ad_multipoint_grid,
46118
),
119+
"AD_MULTIPOINT_GRID_FORWARDDIFF" => (;
120+
env = MIRK_AD_ENV,
121+
body = mirk_ad_multipoint_grid_forwarddiff,
122+
),
123+
"AD_MULTIPOINT_GRID_ENZYME" => (;
124+
env = MIRK_AD_ENV,
125+
body = mirk_ad_multipoint_grid_enzyme,
126+
),
127+
"AD_MULTIPOINT_GRID_MOONCAKE" => (;
128+
env = MIRK_AD_ENV,
129+
body = mirk_ad_multipoint_grid_mooncake,
130+
),
47131
"AD_MULTIPOINT_INTERPOLATION" => (;
48132
env = MIRK_AD_ENV,
49133
body = mirk_ad_multipoint_interpolation,
50134
),
135+
"AD_MULTIPOINT_INTERPOLATION_FORWARDDIFF" => (;
136+
env = MIRK_AD_ENV,
137+
body = mirk_ad_multipoint_interpolation_forwarddiff,
138+
),
139+
"AD_MULTIPOINT_INTERPOLATION_ENZYME" => (;
140+
env = MIRK_AD_ENV,
141+
body = mirk_ad_multipoint_interpolation_enzyme,
142+
),
143+
"AD_MULTIPOINT_INTERPOLATION_MOONCAKE" => (;
144+
env = MIRK_AD_ENV,
145+
body = mirk_ad_multipoint_interpolation_mooncake,
146+
),
51147
"AD_TWOPOINT" => (;
52148
env = MIRK_AD_ENV,
53149
body = mirk_ad_twopoint,
54150
),
151+
"AD_TWOPOINT_FORWARDDIFF" => (;
152+
env = MIRK_AD_ENV,
153+
body = mirk_ad_twopoint_forwarddiff,
154+
),
155+
"AD_TWOPOINT_ENZYME" => (;
156+
env = MIRK_AD_ENV,
157+
body = mirk_ad_twopoint_enzyme,
158+
),
159+
"AD_TWOPOINT_MOONCAKE" => (;
160+
env = MIRK_AD_ENV,
161+
body = mirk_ad_twopoint_mooncake,
162+
),
55163
),
56164
qa = (;
57165
env = joinpath(@__DIR__, "qa"),

lib/BoundaryValueDiffEqMIRK/test/test_groups.toml

Lines changed: 32 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -12,14 +12,43 @@ versions = ["lts", "1", "pre"]
1212
# AD: different-AD-backend (Enzyme/Mooncake) compatibility tests. The heavy
1313
# optional backends live in test/AD/Project.toml (kept out of the main test env so
1414
# the Downgrade lane does not force a joint at-floor resolve of them).
15+
# The lts/pre aggregate shards pass in CI; Julia 1 is split further by backend to
16+
# avoid runner termination during the combined AD compile/solve path.
1517
[AD_MULTIPOINT_GRID]
16-
versions = ["lts", "1", "pre"]
18+
versions = ["lts", "pre"]
19+
20+
[AD_MULTIPOINT_GRID_FORWARDDIFF]
21+
versions = ["1"]
22+
23+
[AD_MULTIPOINT_GRID_ENZYME]
24+
versions = ["1"]
25+
26+
[AD_MULTIPOINT_GRID_MOONCAKE]
27+
versions = ["1"]
1728

1829
[AD_MULTIPOINT_INTERPOLATION]
19-
versions = ["lts", "1", "pre"]
30+
versions = ["lts", "pre"]
31+
32+
[AD_MULTIPOINT_INTERPOLATION_FORWARDDIFF]
33+
versions = ["1"]
34+
35+
[AD_MULTIPOINT_INTERPOLATION_ENZYME]
36+
versions = ["1"]
37+
38+
[AD_MULTIPOINT_INTERPOLATION_MOONCAKE]
39+
versions = ["1"]
2040

2141
[AD_TWOPOINT]
22-
versions = ["lts", "1", "pre"]
42+
versions = ["lts", "pre"]
43+
44+
[AD_TWOPOINT_FORWARDDIFF]
45+
versions = ["1"]
46+
47+
[AD_TWOPOINT_ENZYME]
48+
versions = ["1"]
49+
50+
[AD_TWOPOINT_MOONCAKE]
51+
versions = ["1"]
2352

2453
[QA]
2554
versions = ["lts", "1"]

0 commit comments

Comments
 (0)