Skip to content

Commit 77941f8

Browse files
Merge pull request #4768 from SciML/as/generated-fn-opts
refactor: Introduce `GeneratedFunctionOptions` for the code-generation entry points
2 parents cb2920e + cf1ad2b commit 77941f8

8 files changed

Lines changed: 104 additions & 44 deletions

File tree

Project.toml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -86,7 +86,7 @@ Libdl = "1"
8686
LinearAlgebra = "1"
8787
LinearSolve = "3.66, 4, 5"
8888
Logging = "1"
89-
ModelingToolkitBase = "1.53"
89+
ModelingToolkitBase = "1.54"
9090
ModelingToolkitStandardLibrary = "2.20"
9191
ModelingToolkitTearing = "1.19.2"
9292
Moshi = "0.3.6"

src/ModelingToolkit.jl

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -127,7 +127,8 @@ end
127127

128128
using ModelingToolkitBase: COMMON_SENTINEL, COMMON_NOTHING, COMMON_MISSING,
129129
COMMON_TRUE, COMMON_FALSE, COMMON_INF
130-
using ModelingToolkitBase: build_function_wrapper, BuildFunctionWrapperOptions
130+
using ModelingToolkitBase: build_function_wrapper, BuildFunctionWrapperOptions,
131+
GeneratedFunctionOptions
131132

132133
@recompile_invalidations begin
133134
include("linearization.jl")
@@ -136,6 +137,7 @@ using ModelingToolkitBase: build_function_wrapper, BuildFunctionWrapperOptions
136137

137138
include("problems/docs.jl")
138139
include("systems/codegen.jl")
140+
include("systems/codegen_compat.jl")
139141
include("problems/semilinearodeproblem.jl")
140142
include("problems/sccnonlinearproblem.jl")
141143

src/linearization.jl

Lines changed: 9 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -101,7 +101,10 @@ function linearization_function(
101101
u0 = state_values(prob)
102102

103103
ps = parameters(sys)
104-
h = build_explicit_observed_function(sys, outputs; eval_expression, eval_module)
104+
h = build_explicit_observed_function(
105+
sys, outputs,
106+
GeneratedFunctionOptions(; expression = Val{false}, eval_expression, eval_module)
107+
)
105108

106109
initialization_kwargs = (;
107110
abstol = initialization_abstol, reltol = initialization_reltol,
@@ -586,11 +589,14 @@ function linearize_symbolic(
586589
ps = parameters(sys; initial_parameters = true)
587590
p = Tuple(reorder_parameters(sys, ps))
588591

589-
fun_result = generate_rhs(sys; expression = Val{true})
592+
fun_result = generate_rhs(sys, GeneratedFunctionOptions(; expression = Val{true}))
590593
fun_expr = fun_result isa Tuple ? fun_result[1] : fun_result
591594
fun = eval_or_rgf(fun_expr; eval_expression, eval_module)
592595

593-
h = build_explicit_observed_function(sys, outputs; eval_expression, eval_module)
596+
h = build_explicit_observed_function(
597+
sys, outputs,
598+
GeneratedFunctionOptions(; expression = Val{false}, eval_expression, eval_module)
599+
)
594600
if split
595601
dx = fun(sts, p, t)
596602
y = h(sts, p, t)

src/problems/sccnonlinearproblem.jl

Lines changed: 25 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -12,9 +12,10 @@ const SCC_EXPLICITFUN_CACHE_OUT = unwrap(only(@parameters __outₘₜₖ::Vector
1212

1313
function CacheWriter(
1414
sys::AbstractSystem, buffer_types::Vector{TypeT},
15-
exprs::SCCCacheVarsExprsElT, solsyms;
16-
eval_expression = false, eval_module = @__MODULE__, sparse = false
15+
exprs::SCCCacheVarsExprsElT, solsyms, opts::GeneratedFunctionOptions;
16+
sparse = false
1717
)
18+
(; eval_expression, eval_module) = opts
1819
rps = reorder_parameters(sys) # 1 arg to use the cached version
1920
cache_writes = SymbolicT[]
2021
for (i, T) in enumerate(buffer_types)
@@ -46,9 +47,8 @@ function CacheWriter(
4647
BuildFunctionWrapperOptions(;
4748
p_start = length(solsyms) + 2, p_end = length(rps) + length(solsyms) + 1,
4849
compress_args = [2:(length(solsyms) + 1)],
49-
codegen_function_options = Symbolics.CodegenFunctionOptions(;
50-
expression = Val{true},
51-
iip_config = (true, false)
50+
codegen_function_options = ConstructionBase.setproperties(
51+
opts.codegen, (; iip_config = (true, false))
5252
)
5353
)
5454
)
@@ -57,6 +57,19 @@ function CacheWriter(
5757
return CacheWriter{Any}(fn)
5858
end
5959

60+
# Backward-compatibility keyword method. The positional `opts::GeneratedFunctionOptions`
61+
# method above is the primary; this wrapper preserves the historical keyword API.
62+
function CacheWriter(
63+
sys::AbstractSystem, buffer_types::Vector{TypeT},
64+
exprs::SCCCacheVarsExprsElT, solsyms;
65+
eval_expression = false, eval_module = @__MODULE__, sparse = false
66+
)
67+
return CacheWriter(
68+
sys, buffer_types, exprs, solsyms,
69+
GeneratedFunctionOptions(; eval_expression, eval_module); sparse
70+
)
71+
end
72+
6073
# This phrasing allows us to precompile the calls
6174
@noinline function __explicitfun_copy_states_helper(buffer, sols)
6275
offset = 0
@@ -465,8 +478,11 @@ function SCCNonlinearFunction{iip}(
465478
end
466479
rps = reorder_parameters(subsys)
467480
f = generate_rhs(
468-
subsys; expression = Val{false}, wrap_gfw = Val{true}, cachesyms,
469-
eval_expression, eval_module,
481+
subsys,
482+
GeneratedFunctionOptions(;
483+
expression = Val{false}, wrap_gfw = Val{true}, eval_expression, eval_module
484+
);
485+
cachesyms
470486
)
471487

472488
return NonlinearFunction{iip}(f; sys = subsys)
@@ -596,8 +612,8 @@ function SciMLBase.SCCNonlinearProblem{iip}(
596612
push!(
597613
explicitfuns,
598614
CacheWriter(
599-
sys, decomposition.cachetypes, cacheexprs, solsyms;
600-
eval_expression, eval_module
615+
sys, decomposition.cachetypes, cacheexprs, solsyms,
616+
GeneratedFunctionOptions(; eval_expression, eval_module)
601617
)
602618
)
603619
end

src/problems/semilinearodeproblem.jl

Lines changed: 8 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -20,19 +20,22 @@
2020
_M = concrete_massmatrix(M; sparse, u0)
2121
dvs = unknowns(sys)
2222

23+
codegen_opts = GeneratedFunctionOptions(;
24+
expression, wrap_gfw = Val{true}, eval_expression, eval_module,
25+
codegen_function_options = Symbolics.CodegenFunctionOptions(; kwargs...)
26+
)
27+
2328
f1,
2429
f2 = generate_semiquadratic_functions(
25-
sys, A, B, C; stiff_linear, stiff_quadratic,
26-
stiff_nonlinear, expression, wrap_gfw = Val{true},
27-
eval_expression, eval_module, kwargs...
30+
sys, A, B, C, codegen_opts;
31+
stiff_linear, stiff_quadratic, stiff_nonlinear
2832
)
2933

3034
if jac
3135
check_symbolic_ad_allowed(sys)
3236
Cjac = (C === nothing || !stiff_nonlinear) ? nothing : Symbolics.jacobian(C, dvs)
3337
_jac = generate_semiquadratic_jacobian(
34-
sys, A, B, C, Cjac; sparse, expression,
35-
wrap_gfw = Val{true}, eval_expression, eval_module, kwargs...
38+
sys, A, B, C, Cjac, codegen_opts; sparse
3639
)
3740
_W_sparsity = get_semiquadratic_W_sparsity(
3841
sys, A, B, C, Cjac; stiff_linear, stiff_quadratic, stiff_nonlinear, mm = M

src/systems/codegen.jl

Lines changed: 18 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -204,10 +204,13 @@ $SEMILINEAR_A_B_C_CONSTRAINT
204204
$(MTKBase.EXPERIMENTAL_WARNING)
205205
"""
206206
function generate_semiquadratic_functions(
207-
sys::System, A, B, C; stiff_linear = true,
208-
stiff_quadratic = false, stiff_nonlinear = false, expression = Val{true}, wrap_gfw = Val{false},
209-
eval_expression = false, eval_module = @__MODULE__, kwargs...
207+
sys::System, A, B, C, opts::GeneratedFunctionOptions;
208+
stiff_linear::Bool = true, stiff_quadratic::Bool = false,
209+
stiff_nonlinear::Bool = false
210210
)
211+
(; eval_expression, eval_module) = opts
212+
expression = expression_val(opts)
213+
wrap_gfw = wrap_gfw_val(opts)
211214
if A === nothing && B === nothing
212215
throw(ArgumentError("Cannot generate split form for the system - it has no linear or quadratic part."))
213216
end
@@ -337,36 +340,28 @@ function generate_semiquadratic_functions(
337340
sys, nothing, [Any[Symbolics.DEFAULT_OUTSYM, dvs]; ps; Any[iv]],
338341
BuildFunctionWrapperOptions(;
339342
u_arg = 2, p_start = 3, extra_assignments = f1_iip_ir,
340-
codegen_function_options = Symbolics.CodegenFunctionOptions(;
341-
expression = Val{true}, iip_config = (true, false), kwargs...
342-
)
343+
codegen_function_options = ConstructionBase.setproperties(opts.codegen, (; iip_config = (true, false)))
343344
)
344345
)[1]
345346
f2_iip = build_function_wrapper(
346347
sys, nothing, [Any[Symbolics.DEFAULT_OUTSYM, dvs]; ps; Any[iv]],
347348
BuildFunctionWrapperOptions(;
348349
u_arg = 2, p_start = 3, extra_assignments = f2_iip_ir,
349-
codegen_function_options = Symbolics.CodegenFunctionOptions(;
350-
expression = Val{true}, iip_config = (true, false), kwargs...
351-
)
350+
codegen_function_options = ConstructionBase.setproperties(opts.codegen, (; iip_config = (true, false)))
352351
)
353352
)[1]
354353
f1_oop = build_function_wrapper(
355354
sys, f1_expr, [Any[dvs]; ps; Any[iv]],
356355
BuildFunctionWrapperOptions(;
357356
u_arg = 1,
358-
codegen_function_options = Symbolics.CodegenFunctionOptions(;
359-
expression = Val{true}, iip_config = (true, false), kwargs...
360-
)
357+
codegen_function_options = ConstructionBase.setproperties(opts.codegen, (; iip_config = (true, false)))
361358
)
362359
)[1]
363360
f2_oop = build_function_wrapper(
364361
sys, f2_expr, [Any[dvs]; ps; Any[iv]],
365362
BuildFunctionWrapperOptions(;
366363
u_arg = 1,
367-
codegen_function_options = Symbolics.CodegenFunctionOptions(;
368-
expression = Val{true}, iip_config = (true, false), kwargs...
369-
)
364+
codegen_function_options = ConstructionBase.setproperties(opts.codegen, (; iip_config = (true, false)))
370365
)
371366
)[1]
372367

@@ -402,10 +397,13 @@ $SEMILINEAR_A_B_C_CONSTRAINT
402397
$EXPERIMENTAL_WARNING
403398
"""
404399
function generate_semiquadratic_jacobian(
405-
sys::System, A, B, C, Cjac; sparse = false, stiff_linear = true, stiff_quadratic = false,
406-
stiff_nonlinear = false, expression = Val{true}, wrap_gfw = Val{false},
407-
eval_expression = false, eval_module = @__MODULE__, kwargs...
400+
sys::System, A, B, C, Cjac, opts::GeneratedFunctionOptions;
401+
sparse::Bool = false, stiff_linear::Bool = true, stiff_quadratic::Bool = false,
402+
stiff_nonlinear::Bool = false
408403
)
404+
(; eval_expression, eval_module) = opts
405+
expression = expression_val(opts)
406+
wrap_gfw = wrap_gfw_val(opts)
409407
if sparse
410408
error("Sparse analytical jacobians for split ODEs is not implemented.")
411409
end
@@ -536,18 +534,14 @@ function generate_semiquadratic_jacobian(
536534
sys, nothing, [Any[Symbolics.DEFAULT_OUTSYM, dvs]; ps; Any[iv]],
537535
BuildFunctionWrapperOptions(;
538536
u_arg = 2, p_start = 3, extra_assignments = iip_ir,
539-
codegen_function_options = Symbolics.CodegenFunctionOptions(;
540-
expression = Val{true}, iip_config = (true, false), kwargs...
541-
)
537+
codegen_function_options = ConstructionBase.setproperties(opts.codegen, (; iip_config = (true, false)))
542538
)
543539
)
544540
j_oop, _ = build_function_wrapper(
545541
sys, oop_expr, [Any[dvs]; ps; Any[iv]],
546542
BuildFunctionWrapperOptions(;
547543
u_arg = 1,
548-
codegen_function_options = Symbolics.CodegenFunctionOptions(;
549-
expression = Val{true}, iip_config = (true, false), kwargs...
550-
)
544+
codegen_function_options = ConstructionBase.setproperties(opts.codegen, (; iip_config = (true, false)))
551545
)
552546
)
553547
return maybe_compile_function(

src/systems/codegen_compat.jl

Lines changed: 35 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,35 @@
1+
# Backwards-compatibility keyword-argument wrappers for the code-generation entry points
2+
# defined in `ModelingToolkit` (as opposed to `ModelingToolkitBase`). See
3+
# `ModelingToolkitBase`'s `systems/codegen_compat.jl` for the rationale: the primary methods
4+
# take a positional `GeneratedFunctionOptions` plus strictly-typed function-specific keyword
5+
# arguments; these wrappers accept the historical loose keyword arguments and forward.
6+
7+
function generate_semiquadratic_functions(
8+
sys::System, A, B, C; stiff_linear = true, stiff_quadratic = false,
9+
stiff_nonlinear = false, expression = Val{true}, wrap_gfw = Val{false},
10+
eval_expression = false, eval_module = @__MODULE__, kwargs...
11+
)
12+
return generate_semiquadratic_functions(
13+
sys, A, B, C,
14+
GeneratedFunctionOptions(;
15+
expression, wrap_gfw, eval_expression, eval_module,
16+
codegen_function_options = Symbolics.CodegenFunctionOptions(; kwargs...)
17+
);
18+
stiff_linear, stiff_quadratic, stiff_nonlinear
19+
)
20+
end
21+
22+
function generate_semiquadratic_jacobian(
23+
sys::System, A, B, C, Cjac; sparse = false, stiff_linear = true,
24+
stiff_quadratic = false, stiff_nonlinear = false, expression = Val{true},
25+
wrap_gfw = Val{false}, eval_expression = false, eval_module = @__MODULE__, kwargs...
26+
)
27+
return generate_semiquadratic_jacobian(
28+
sys, A, B, C, Cjac,
29+
GeneratedFunctionOptions(;
30+
expression, wrap_gfw, eval_expression, eval_module,
31+
codegen_function_options = Symbolics.CodegenFunctionOptions(; kwargs...)
32+
);
33+
sparse, stiff_linear, stiff_quadratic, stiff_nonlinear
34+
)
35+
end

src/systems/solver_nlprob.jl

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -90,5 +90,9 @@ function (nlp::NLStep_probmap)(nlsol)
9090
end
9191

9292
function generate_nlprobmap(sys::System, nlsys::System)
93-
return NLStep_probmap(build_explicit_observed_function(nlsys, unknowns(sys)))
93+
return NLStep_probmap(
94+
build_explicit_observed_function(
95+
nlsys, unknowns(sys), GeneratedFunctionOptions(; expression = Val{false})
96+
)
97+
)
9498
end

0 commit comments

Comments
 (0)