diff --git a/docs/Project.toml b/docs/Project.toml index 39d5ddb615..df14347ec6 100644 --- a/docs/Project.toml +++ b/docs/Project.toml @@ -45,7 +45,7 @@ CairoMakie = "0.15" CommonSolve = "0.2" ControlSystemsBase = "1.20" ControlSystemsMTK = "2.6" -DataInterpolations = "8" +DataInterpolations = "8, 9" Distributions = "0.25" Documenter = "1" DynamicQuantities = "1" diff --git a/lib/ModelingToolkitBase/Project.toml b/lib/ModelingToolkitBase/Project.toml index f28ef17bfe..915c9e207b 100644 --- a/lib/ModelingToolkitBase/Project.toml +++ b/lib/ModelingToolkitBase/Project.toml @@ -40,6 +40,7 @@ OffsetArrays = "6fe1bfb0-de20-5000-8ca7-80f57d26f881" OrderedCollections = "bac558e1-5e72-5ebc-8fee-abe8a469f55d" PreallocationTools = "d236fae5-4411-538c-8e31-a6e3d9e00b46" PrecompileTools = "aea7be01-6a6a-4083-8856-8a6e6704d82a" +Printf = "de0858da-6303-5e67-8744-51eddeeeb8d7" REPL = "3fa0cd96-eef1-5676-8a61-b3b8758bbffb" Random = "9a3f8284-a2c9-5f02-9a11-845980a1fd5c" ReadOnlyDicts = "795d4caa-f5a7-4580-b5d8-c01d53451803" @@ -162,6 +163,7 @@ OrdinaryDiffEqTsit5 = "1, 2" Pkg = "1" PreallocationTools = "0.4.27, 1" PrecompileTools = "1.2.1" +Printf = "1" Pyomo = "0.1.0" REPL = "1" Random = "1" @@ -172,7 +174,7 @@ ReferenceTests = "0.10" RuntimeGeneratedFunctions = "0.5.12" SCCNonlinearSolve = "1.13" SafeTestsets = "0.1" -SciMLBase = "3.19" +SciMLBase = "3.38" SciMLPublic = "1.0.0" SciMLStructures = "1.7" Serialization = "1" diff --git a/lib/ModelingToolkitBase/src/ModelingToolkitBase.jl b/lib/ModelingToolkitBase/src/ModelingToolkitBase.jl index 0597d84105..899aab1aa3 100644 --- a/lib/ModelingToolkitBase/src/ModelingToolkitBase.jl +++ b/lib/ModelingToolkitBase/src/ModelingToolkitBase.jl @@ -21,6 +21,10 @@ using PrecompileTools, Reexport import BandedMatrices: BandedMatrices, BandedMatrix, bandwidths end +import SciMLBase +import SciMLBase: diagnose_symbolic_instability +using Printf: @sprintf + import SymbolicUtils import SymbolicUtils as SU import SymbolicUtils: iscall, arguments, operation, maketerm, promote_symtype, diff --git a/lib/ModelingToolkitBase/src/debugging.jl b/lib/ModelingToolkitBase/src/debugging.jl index 82e4c80627..bbc410f7a4 100644 --- a/lib/ModelingToolkitBase/src/debugging.jl +++ b/lib/ModelingToolkitBase/src/debugging.jl @@ -99,3 +99,122 @@ function get_assertions_expr(sys::AbstractSystem) end return term end + +function SciMLBase.diagnose_symbolic_instability(sys::AbstractSystem, u, uprev) + diagnosis = String[] + + #check for assertion failures + unks = unknowns(sys) + curr_substitution_map = Dict{SymbolicT, SymbolicT}(zip(unks, u)) + prev_substitution_map = Dict{SymbolicT, SymbolicT}(zip(unknowns(sys), uprev)) + + for (cond, msg) in assertions(sys) + subclauses = String[] + find_failing_subterms(cond, prev_substitution_map, curr_substitution_map, subclauses) + if !isempty(subclauses) + push!(diagnosis, "\n\nAssertion violated: $cond - \"$msg\"") + append!(diagnosis, subclauses) + end + end + + #find singularity causes in equations + singularities = String[] + visited = IdDict{SymbolicT, Nothing}() + subber = SymbolicUtils.IRSubstituter{true}(get_irstructure(sys), prev_substitution_map) + for eq in full_equations(sys) + find_singular_subterms(eq, eq.rhs, subber, singularities, visited) + end + if !isempty(singularities) + push!(diagnosis, "\nSymbolic Analysis of MTK System:") + append!(diagnosis, singularities) + end + + return isempty(diagnosis) ? "" : join(diagnosis, "\n") +end + +function find_singular_subterms(eq, expr, sub_map, diagnosis, visited) + expr = unwrap(expr) + !SymbolicUtils.iscall(expr) && return diagnosis + op = SymbolicUtils.operation(expr) + args = SymbolicUtils.arguments(expr) + haskey(visited, expr) && return diagnosis + visited[expr] = nothing + + if op === (/) #division, singular if we divide by small thing + d = Symbolics.value(sub_map(args[2])) + if d isa Number && abs(d) < 1e-10 + push!(diagnosis, "in equation $eq: division by very small value $(args[2]) ≈ $(@sprintf("%.4g", d)) leads to singularity.") + end + elseif op === log #singular if we log small thing + x = Symbolics.value(sub_map(args[1])) + if x isa Number && x <= 1e-10 + push!(diagnosis, "in equation $eq: log of $(args[1]) = $(@sprintf("%.4g", x)) near/at singularity (derivative blows up).") + end + elseif op === sqrt + x = Symbolics.value(sub_map(args[1])) + if x isa Number && x < 1e-10 + push!(diagnosis, "in equation $eq: sqrt of $(args[1]) = $(@sprintf("%.4g", x)) near/at singularity (derivative blows up).") + end + elseif op === (^) + e = Symbolics.value(sub_map(args[2])) + b = Symbolics.value(sub_map(args[1])) + if e isa Number && b isa Number #two cases + if e < 0 && abs(b) < 1e-10 + push!(diagnosis, "in equation $eq: ($(args[1])) raised to power $e with base ≈ $(@sprintf("%.4g", b)) going to 0; result diverges.") + elseif e > 0 && abs(b) > 1 + push!(diagnosis, "in equation $eq: ($(args[1]) ≈ $(@sprintf("%.4g", b))) raised to power $e - base magnitude is large and being amplified.") + end + end + end + + for arg in args + find_singular_subterms(eq, arg, sub_map, diagnosis, visited) + end + return diagnosis +end + +function find_failing_subterms(cond, prev_map, curr_map, diagnosis) + c = Symbolics.unwrap(cond) + !SymbolicUtils.iscall(c) && return diagnosis + op = SymbolicUtils.operation(c) + args = SymbolicUtils.arguments(c) + + if (op === (<) || op === (>) || op === (<=) || op === (>=)) && length(args) == 2 + #compare using previous non-nan values to find violating subclauses, then output current values + lhs = Symbolics.value(Symbolics.substitute(args[1], prev_map)) + rhs = Symbolics.value(Symbolics.substitute(args[2], prev_map)) + if lhs isa Number && rhs isa Number + # small margin -> violated + margin = (op === (<) || op === (<=)) ? rhs - lhs : lhs - rhs + if margin <= 1e-6 + push!(diagnosis, " subclause `$c` violated: $(clause_values(c, curr_map))") + end + end + elseif op === (!=) && length(args) == 2 + lhs = Symbolics.value(Symbolics.substitute(args[1], prev_map)) + rhs = Symbolics.value(Symbolics.substitute(args[2], prev_map)) + if lhs isa Number && rhs isa Number && abs(lhs - rhs) <= 1e-6 + push!(diagnosis, " subclause `$c` violated: $(clause_values(c, curr_map))") + end + elseif op === (==) && length(args) == 2 + lhs = Symbolics.value(Symbolics.substitute(args[1], prev_map)) + rhs = Symbolics.value(Symbolics.substitute(args[2], prev_map)) + if lhs isa Number && rhs isa Number && abs(lhs - rhs) > 1e-6 + push!(diagnosis, " subclause `$c` violated: $(clause_values(c, curr_map))") + end + else #recurse + for arg in args + find_failing_subterms(arg, prev_map, curr_map, diagnosis) + end + end + return diagnosis +end + +function clause_values(c, curr_map) + parts = String[] + for v in Symbolics.get_variables(c) + val = Symbolics.value(Symbolics.substitute(v, curr_map)) + push!(parts, val isa Number ? "$v = $(@sprintf("%.4g", val))" : "$v = $val") + end + return join(parts, ", ") +end diff --git a/lib/ModelingToolkitBase/test/optimization/Project.toml b/lib/ModelingToolkitBase/test/optimization/Project.toml index 7c89fb85f8..fa13aadfdc 100644 --- a/lib/ModelingToolkitBase/test/optimization/Project.toml +++ b/lib/ModelingToolkitBase/test/optimization/Project.toml @@ -32,7 +32,7 @@ ModelingToolkitBase = {path = "../.."} [compat] CasADi = "1.0.7" -DataInterpolations = "8.8" +DataInterpolations = "8.8, 9" OrdinaryDiffEqExplicitTableaus = "2" OrdinaryDiffEqImplicitTableaus = "2" SafeTestsets = "0.1, 1"