Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
60 changes: 36 additions & 24 deletions lib/ModelingToolkitBase/src/utils.jl
Original file line number Diff line number Diff line change
Expand Up @@ -936,6 +936,18 @@ function collect_vars!(
return nothing
end

# Break the inference cycle between the mutually recursive collectors. Julia 1.10
# can otherwise miscompile the cycle after unrelated method additions. The
# abstractly typed cell keeps downstream `collect_vars!` dispatch dynamic, while
# the explicit keyword preserves metadata recursion's depth-zero semantics.
const _COLLECT_VARS_DISPATCH = Ref{Function}(collect_vars!)

@noinline Base.@constprop :none function _call_collect_vars!(
unknowns, parameters, expr, iv
)
return _COLLECT_VARS_DISPATCH[](unknowns, parameters, expr, iv; depth = 0)
end

"""
$(TYPEDSIGNATURES)

Expand Down Expand Up @@ -966,7 +978,7 @@ function collect_var!(unknowns::OrderedSet{SymbolicT}, parameters::OrderedSet{Sy
any(!SU.isconst, Iterators.drop(arguments(var), 1))
)
for arg in Iterators.drop(arguments(var), 1)
collect_vars!(unknowns, parameters, arg, iv)
_call_collect_vars!(unknowns, parameters, arg, iv)
end
var = arr
end
Expand All @@ -975,7 +987,7 @@ function collect_var!(unknowns::OrderedSet{SymbolicT}, parameters::OrderedSet{Sy
if iscalledparameter(var)
callable = getcalledparameter(var)
push!(parameters, callable)
collect_vars!(unknowns, parameters, arguments(var), iv)
_call_collect_vars!(unknowns, parameters, arguments(var), iv)
elseif isparameter(var) || (iscall(var) && isparameter(operation(var)))
push!(parameters, var)
else
Expand All @@ -984,55 +996,55 @@ function collect_var!(unknowns::OrderedSet{SymbolicT}, parameters::OrderedSet{Sy
# Add also any parameters that appear only as defaults in the var
if hasdefault(var) && (def = getdefault(var)) !== missing
if def isa SymbolicT
collect_vars!(unknowns, parameters, def, iv)
_call_collect_vars!(unknowns, parameters, def, iv)
elseif def isa Num
collect_vars!(unknowns, parameters, def, iv)
_call_collect_vars!(unknowns, parameters, def, iv)
elseif def isa Arr{Num, 1}
collect_vars!(unknowns, parameters, def, iv)
_call_collect_vars!(unknowns, parameters, def, iv)
elseif def isa Arr{Num, 2}
collect_vars!(unknowns, parameters, def, iv)
_call_collect_vars!(unknowns, parameters, def, iv)
elseif def isa CallAndWrap{Num}
collect_vars!(unknowns, parameters, def, iv)
_call_collect_vars!(unknowns, parameters, def, iv)
elseif def isa CallAndWrap{Arr{Num, 1}}
collect_vars!(unknowns, parameters, def, iv)
_call_collect_vars!(unknowns, parameters, def, iv)
elseif def isa CallAndWrap{Arr{Num, 2}}
collect_vars!(unknowns, parameters, def, iv)
_call_collect_vars!(unknowns, parameters, def, iv)
elseif def isa Arr
collect_vars!(unknowns, parameters, def, iv)
_call_collect_vars!(unknowns, parameters, def, iv)
elseif def isa CallAndWrap
collect_vars!(unknowns, parameters, def, iv)
_call_collect_vars!(unknowns, parameters, def, iv)
else
collect_vars!(unknowns, parameters, def, iv)
_call_collect_vars!(unknowns, parameters, def, iv)
end
end
# Add also any parameters that appear only in the bounds of the var
if hasbounds(var)
(lo, hi) = getbounds(var)
if lo isa SymbolicT
collect_vars!(unknowns, parameters, lo, iv)
_call_collect_vars!(unknowns, parameters, lo, iv)
elseif lo isa Num
collect_vars!(unknowns, parameters, lo, iv)
_call_collect_vars!(unknowns, parameters, lo, iv)
elseif lo isa Arr{Num, 1}
collect_vars!(unknowns, parameters, lo, iv)
_call_collect_vars!(unknowns, parameters, lo, iv)
elseif lo isa Arr{Num, 2}
collect_vars!(unknowns, parameters, lo, iv)
_call_collect_vars!(unknowns, parameters, lo, iv)
elseif lo isa Arr
collect_vars!(unknowns, parameters, lo, iv)
_call_collect_vars!(unknowns, parameters, lo, iv)
else
collect_vars!(unknowns, parameters, lo, iv)
_call_collect_vars!(unknowns, parameters, lo, iv)
end
if hi isa SymbolicT
collect_vars!(unknowns, parameters, hi, iv)
_call_collect_vars!(unknowns, parameters, hi, iv)
elseif hi isa Num
collect_vars!(unknowns, parameters, hi, iv)
_call_collect_vars!(unknowns, parameters, hi, iv)
elseif hi isa Arr{Num, 1}
collect_vars!(unknowns, parameters, hi, iv)
_call_collect_vars!(unknowns, parameters, hi, iv)
elseif hi isa Arr{Num, 2}
collect_vars!(unknowns, parameters, hi, iv)
_call_collect_vars!(unknowns, parameters, hi, iv)
elseif hi isa Arr
collect_vars!(unknowns, parameters, hi, iv)
_call_collect_vars!(unknowns, parameters, hi, iv)
else
collect_vars!(unknowns, parameters, hi, iv)
_call_collect_vars!(unknowns, parameters, hi, iv)
end
end
return nothing
Expand Down
67 changes: 47 additions & 20 deletions lib/ModelingToolkitBase/test/dq_units.jl
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ using ModelingToolkitBase, OrdinaryDiffEq, JumpProcesses, DynamicQuantities
using Symbolics
import SymbolicUtils as SU
using Test
using DataStructures: OrderedSet
MT = ModelingToolkitBase
using ModelingToolkitBase: t, D
@parameters τ [unit = u"s"] γ
Expand Down Expand Up @@ -273,30 +274,56 @@ let
@test MT.get_unit(x_mat) == u"1"
end

# Issue #4211: collect_vars! must discover parameters in defaults with DynamicQuantities loaded
# On Julia 1.10, loading DynamicQuantities could cause collect_vars! to fail to discover
# parameters used in variable defaults due to a method invalidation bug.
@testset "Issue #4211: collect_vars! discovers parameters in defaults" begin
using DataStructures: OrderedSet
struct CollectorInvalidationString
value::String
end

struct CollectorInvalidationExpression
parameter::Symbolics.SymbolicT
end

@parameters X0_test
@variables X_test(t) = X0_test
function collect_default_parameters(var)
unknowns = OrderedSet{Symbolics.SymbolicT}()
parameters = OrderedSet{Symbolics.SymbolicT}()
MT.collect_vars!(
unknowns, parameters, Symbolics.unwrap(var), Symbolics.unwrap(t),
Symbolics.Operator; depth = 0
)
return parameters
end

us = OrderedSet{Symbolics.SymbolicT}()
ps = OrderedSet{Symbolics.SymbolicT}()
MT.collect_vars!(us, ps, Symbolics.unwrap(X_test), Symbolics.unwrap(t), Symbolics.Operator; depth = 0)
@testset "collect_vars! survives late method invalidation" begin
@parameters x0_test y0_test scale_test
@variables x_test(t) = x0_test y_test(t) = scale_test * y0_test

# X0_test should be discovered in X_test's default value
@test Symbolics.unwrap(X0_test) in ps
@test Symbolics.unwrap(x0_test) in collect_default_parameters(x_test)

# Test with expression in default
@parameters a_test b_test
@variables Y_test(t) = a_test + 2 * b_test
# This late definition exercises the Julia 1.10 invalidation that caused
# parameters in defaults to disappear after loading unrelated packages.
@eval Base.convert(::Type{Symbol}, value::CollectorInvalidationString) =
Symbol(value.value)

empty!(us)
empty!(ps)
MT.collect_vars!(us, ps, Symbolics.unwrap(Y_test), Symbolics.unwrap(t), Symbolics.Operator; depth = 0)
x_parameters = collect_default_parameters(x_test)
y_parameters = collect_default_parameters(y_test)

@test Symbolics.unwrap(x0_test) in x_parameters
@test Symbolics.unwrap(y0_test) in y_parameters
@test Symbolics.unwrap(scale_test) in y_parameters

custom_default = MT.setdefault(
x_test, CollectorInvalidationExpression(Symbolics.unwrap(scale_test))
)
@eval function MT.collect_vars!(
unknowns::OrderedSet{Symbolics.SymbolicT},
parameters::OrderedSet{Symbolics.SymbolicT},
expr::CollectorInvalidationExpression,
::Union{Symbolics.SymbolicT, Nothing};
depth = 0
)
push!(parameters, expr.parameter)
return nothing
end

@test Symbolics.unwrap(a_test) in ps
@test Symbolics.unwrap(b_test) in ps
@test Symbolics.unwrap(scale_test) in collect_default_parameters(custom_default)
@test Symbolics.unwrap(x0_test) in collect_default_parameters(x_test)
end
Loading