Skip to content

Commit 6d9273a

Browse files
add rootcause analysis of symbolic system
1 parent cd3468b commit 6d9273a

3 files changed

Lines changed: 73 additions & 4 deletions

File tree

lib/OrdinaryDiffEqCore/Project.toml

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -39,12 +39,14 @@ SymbolicIndexingInterface = "2efcf032-c050-4f8e-a9bb-153293bab1f5"
3939
TruncatedStacktraces = "781d530d-4396-4725-bb49-402e4bee1e77"
4040

4141
[weakdeps]
42+
ModelingToolkit = "961ee093-0014-501f-94e3-6117800e7a78"
4243
Mooncake = "da2b9cff-9c12-43a0-ae48-6db2b0edb7d6"
4344
Polyester = "f517fe37-dbe3-4b94-8317-1923a5111588"
4445
SparseArrays = "2f01184e-e22b-5df5-ae63-d93ebab69eaf"
4546

4647
[extensions]
4748
OrdinaryDiffEqCoreMooncakeExt = "Mooncake"
49+
OrdinaryDiffEqModelingToolkitExt = "ModelingToolkit"
4850
OrdinaryDiffEqCorePolyesterExt = "Polyester"
4951
OrdinaryDiffEqCoreSparseArraysExt = "SparseArrays"
5052

Lines changed: 54 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,54 @@
1+
module OrdinaryDiffEqModelingToolkitExt
2+
3+
using OrdinaryDiffEqCore, ModelingToolkit
4+
using Printf: @sprintf
5+
6+
function OrdinaryDiffEqCore.system_singularity_rootcause(sys, u)
7+
substitution_map = Dict(zip(unknowns(sys), u))
8+
diagnosis = String[]
9+
for eq in equations(sys)
10+
find_singular_subterms(eq, eq.rhs, substitution_map, diagnosis)
11+
end
12+
return diagnosis
13+
end
14+
15+
function find_singular_subterms(eq, expr, sub_map, diagnosis)
16+
expr = Symbolics.unwrap(expr)
17+
!SymbolicUtils.iscall(expr) && return diagnosis
18+
op = SymbolicUtils.operation(expr)
19+
args = SymbolicUtils.arguments(expr)
20+
21+
if op === (/) #division, singular if we divide by small thing
22+
d = Symbolics.value(Symbolics.substitute(args[2], sub_map))
23+
if d isa Number && abs(d) < 1e-10
24+
push!(diagnosis, "in equation $eq: division by very small value $(args[2])$(@sprintf("%.4g", d)) leads to singularity.")
25+
end
26+
elseif op === log #singular if we log small thing
27+
x = Symbolics.value(Symbolics.substitute(args[1], sub_map))
28+
if x isa Number && x <= 1e-10
29+
push!(diagnosis, "in equation $eq: log of $(args[1]) = $(@sprintf("%.4g", x)) near/at singularity (derivative blows up).")
30+
end
31+
elseif op === sqrt
32+
x = Symbolics.value(Symbolics.substitute(args[1], sub_map))
33+
if x isa Number && x < 1e-10
34+
push!(diagnosis, "in equation $eq: sqrt of $(args[1]) = $(@sprintf("%.4g", x)) near/at singularity (derivative blows up).")
35+
end
36+
elseif op === (^)
37+
e = Symbolics.value(Symbolics.substitute(args[2], sub_map))
38+
b = Symbolics.value(Symbolics.substitute(args[1], sub_map))
39+
if e isa Number && b isa Number #two cases
40+
if e < 0 && abs(b) < 1e-10
41+
push!(diagnosis, "in equation $eq: ($(args[1])) raised to power $e with base ≈ $(@sprintf("%.4g", b)) going to 0; result diverges.")
42+
elseif e > 0 && abs(b) > 1
43+
push!(diagnosis, "in equation $eq: ($(args[1])$(@sprintf("%.4g", b))) raised to power $e - base magnitude is large and being amplified.")
44+
end
45+
end
46+
end
47+
48+
for arg in args
49+
find_singular_subterms(eq, arg, sub_map, diagnosis)
50+
end
51+
return diagnosis
52+
end
53+
54+
end

lib/OrdinaryDiffEqCore/src/integrators/integrator_utils.jl

Lines changed: 17 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -615,6 +615,8 @@ end
615615
# overrides this with a method that calls calc_J to get a fresh Jacobian.
616616
get_fresh_jacobian(integrator, cache) = cache.J
617617

618+
system_singularity_rootcause(sys, u) = ""
619+
618620
function SciMLBase.log_instability(integrator::ODEIntegrator)
619621
W = _get_W(integrator)
620622
u = integrator.u
@@ -684,6 +686,11 @@ function SciMLBase.log_instability(integrator::ODEIntegrator)
684686
sym_eqs = (sys !== nothing && hasfield(typeof(sys), :eqs)) ? getfield(sys, :eqs) : nothing
685687
sym_vars = (sys !== nothing && hasfield(typeof(sys), :unknowns)) ? getfield(sys, :unknowns) : nothing
686688

689+
# symbolic analysis
690+
#skip jac analysis if this isn't empty
691+
u_vals = length(nan_inf_idxs) == 0 ? u : integrator.uprev
692+
symbolic_analysis = system_singularity_rootcause(sys, u_vals)
693+
687694
# diagnostic message construction
688695
diagnostic = String[]
689696
if !isempty(nan_inf_idxs) #state vars
@@ -712,7 +719,7 @@ function SciMLBase.log_instability(integrator::ODEIntegrator)
712719
end
713720
end
714721

715-
if bad_entries !== nothing && !isempty(bad_entries) #Jacobian analysis
722+
if bad_entries !== nothing && !isempty(bad_entries) && isempty(symbolic_analysis) #Jacobian analysis (skipped if we have symbolic analysis)
716723
has_nonfinite = false
717724
has_large = false
718725
for (_, _, v) in bad_entries
@@ -730,7 +737,7 @@ function SciMLBase.log_instability(integrator::ODEIntegrator)
730737
for (i, j, v) in first(bad_entries, 5)
731738
push!(example_strs, "J[$i,$j] = $(@sprintf("%.4g", v))")
732739
end
733-
push!(diagnostic, "Jacobian row(s) $singularity_rows have $entry_desc entries (e.g. $(join(example_strs, ", "))), suggesting a singularity in those equation(s)")
740+
push!(diagnostic, "\nJacobian row(s) $singularity_rows have $entry_desc entries (e.g. $(join(example_strs, ", "))), suggesting a singularity in those equation(s)")
734741
if sym_eqs !== nothing
735742
for row in singularity_rows
736743
if row <= length(sym_eqs)
@@ -740,7 +747,7 @@ function SciMLBase.log_instability(integrator::ODEIntegrator)
740747
end
741748
# jac cols
742749
if !isempty(singularity_cols)
743-
push!(diagnostic, "Jacobian column(s) $singularity_cols have $entry_desc entries, suggesting those state component(s) are diverging")
750+
push!(diagnostic, "\nJacobian column(s) $singularity_cols have $entry_desc entries, suggesting those state component(s) are diverging")
744751
if sym_vars !== nothing
745752
for col in singularity_cols
746753
if col <= length(sym_vars)
@@ -750,7 +757,13 @@ function SciMLBase.log_instability(integrator::ODEIntegrator)
750757
end
751758
end
752759
end
753-
diagnostic = isempty(diagnostic) ? "" : "\nDiagnostics:\n" * join(diagnostic, "\n") * "."
760+
761+
diagnostic = isempty(diagnostic) ? "" : "\n\nDiagnostics:\n" * join(diagnostic, "\n") * "."
762+
763+
if !isempty(symbolic_analysis)
764+
diagnostic *= "\n\nSymbolic Analysis of MTK System:\n" * join(symbolic_analysis, "\n")
765+
end
766+
754767
return diagnostic
755768
end
756769

0 commit comments

Comments
 (0)