From 35361ef20047616d391903cb9ff4383ff747c52d Mon Sep 17 00:00:00 2001 From: ChrisRackauckas-Claude Date: Wed, 9 Sep 2026 18:24:34 -0400 Subject: [PATCH 1/2] MTKTearing: track array differential equations as row groups through structural simplification Array equations `D(x[slice]) ~ rhs` are still scalarized into rows of the bipartite graph, so matching, Pantelides, dummy derivatives, tearing and alias elimination see exact per-element incidence. The rows now remember the array equation they came from (`ArrayEquationGroup`, `row_group`, `row_elem`), and every pass that rewrites a row in a way the array equation cannot represent (differentiation, removal, dummy derivative substitution, solving for another variable, inline linear SCCs, clock partition splits) marks the group dirty. Rows of the integer-linear subsystem that belong to a group are kept out of Gaussian elimination so they are neither reduced nor used as pivots. With `preserve_array_equations = true` on `DefaultReassembleAlgorithm` (or as a `mtkcompile` keyword) intact groups are emitted as a single array equation over scalar unknowns; dirty groups are emitted scalarized as before. The default output is unchanged. Co-authored-by: Cursor --- lib/ModelingToolkitTearing/Project.toml | 2 +- .../src/clock_inference/interface.jl | 15 + lib/ModelingToolkitTearing/src/reassemble.jl | 194 ++++++++++- .../src/stateselection_interface.jl | 59 ++++ .../src/tearingstate.jl | 317 +++++++++++++++++- lib/ModelingToolkitTearing/test/runtests.jl | 118 +++++++ 6 files changed, 685 insertions(+), 20 deletions(-) diff --git a/lib/ModelingToolkitTearing/Project.toml b/lib/ModelingToolkitTearing/Project.toml index cb11bc9..bb59ed2 100644 --- a/lib/ModelingToolkitTearing/Project.toml +++ b/lib/ModelingToolkitTearing/Project.toml @@ -1,6 +1,6 @@ name = "ModelingToolkitTearing" uuid = "6bb917b9-1269-42b9-9f7c-b0dca72083ab" -version = "1.20.6" +version = "1.21.0" authors = ["Aayush Sabharwal "] [deps] diff --git a/lib/ModelingToolkitTearing/src/clock_inference/interface.jl b/lib/ModelingToolkitTearing/src/clock_inference/interface.jl index 9550206..e7732c7 100644 --- a/lib/ModelingToolkitTearing/src/clock_inference/interface.jl +++ b/lib/ModelingToolkitTearing/src/clock_inference/interface.jl @@ -55,6 +55,7 @@ function substitute_sample_time(ci::ClockInference{TearingState}, ts::TearingSta else subrules[st] = dt neweq = substitute(eq, subrules) + dirty_array_group!(ts, i) end eqs[i] = neweq end @@ -73,6 +74,20 @@ function system_subset(ts::TearingState, ieqs::Vector{Int}, iieqs::Vector{Int}, if !isempty(ts.eqs_source) @set! ts.eqs_source = ts.eqs_source[ieqs] end + # Array blocks are only tracked once the equations are scalarized (after clock + # inference), so the subset only carries scalar rows. Any groups that exist (eagerly + # scalarized state) that are split across partitions are broken. + array_groups = map(g -> ArrayEquationGroup(g.eq, g.lhs_vars, g.dirty), ts.array_groups) + row_group = ts.row_group[ieqs] + row_elem = ts.row_elem[ieqs] + old_counts = count_group_rows(length(array_groups), ts.row_group) + new_counts = count_group_rows(length(array_groups), row_group) + for g in eachindex(array_groups) + old_counts[g] == new_counts[g] || (array_groups[g].dirty = true) + end + @set! ts.array_groups = array_groups + @set! ts.row_group = row_group + @set! ts.row_elem = row_elem if all(eq -> eq.rhs isa StateMachineOperator, MTKBase.get_eqs(ts.sys)) names = Symbol[] for eq in MTKBase.get_eqs(ts.sys) diff --git a/lib/ModelingToolkitTearing/src/reassemble.jl b/lib/ModelingToolkitTearing/src/reassemble.jl index 9dba0cf..c976772 100644 --- a/lib/ModelingToolkitTearing/src/reassemble.jl +++ b/lib/ModelingToolkitTearing/src/reassemble.jl @@ -43,6 +43,9 @@ function substitute_derivatives_algevars!( for eq in 𝑑neighbors(graph, dv) dummy_sub[dd] = v_t neweqs[eq] = substitute(neweqs[eq], dd => v_t) + # The array equation would still contain `D(x)`, which the scalarized row + # no longer does. + dirty_array_group!(ts, eq) end fullvars[dv] = v_t # If we have: @@ -56,6 +59,7 @@ function substitute_derivatives_algevars!( dx_t = D(x_t) for eq in 𝑑neighbors(graph, ddx) neweqs[eq] = substitute(neweqs[eq], fullvars[ddx] => dx_t) + dirty_array_group!(ts, eq) end fullvars[ddx] = dx_t dx = ddx @@ -498,6 +502,10 @@ function generate_system_equations!(state::TearingState, neweqs::Vector{Equation @assert length(vars_mask) == length(vscc) _escc = escc[eqs_mask] _vscc = vscc[vars_mask] + # Rows solved as part of a linear block are no longer their own array element. + for ieq in _escc + dirty_array_group!(state, ieq) + end # `linsol` is the `A \ b` term (runtime path); component `j` is solved # for the variable assigned below. The analytical path returns a # `Const`-wrapped vector instead, which is not reported. @@ -585,8 +593,22 @@ function generate_system_equations!(state::TearingState, neweqs::Vector{Equation # after the SCC-ordered block. Extra variables are likewise suffixed by the `setdiff` # append below. Algebraic placeholders left unfilled (e.g. the redundant equations of an # overdetermined system) stay `0` and are dropped by the `filter!`. + # Array equation membership of each generated equation, so that the rows of an intact + # array equation end up contiguous and in element order. + eq_group = map(Base.Fix1(row_group, state), eq_ordering) + eq_elem = map(eq_ordering) do ieq + ieq <= length(state.row_elem) ? state.row_elem[ieq] : 0 + end + for (k, g) in enumerate(eq_group) + iszero(g) && continue + if state.array_groups[g].dirty + eq_group[k] = 0 + eq_elem[k] = 0 + end + end blt_reorder_generated_equations!( - neweqs′, eq_ordering, var_ordering, eq_scc, var_sccs, findnextfn, ndsts(graph)) + neweqs′, eq_ordering, var_ordering, eq_scc, var_sccs, findnextfn, ndsts(graph); + eq_group, eq_elem) filter!(!iszero, var_ordering) var_ordering = [var_ordering; setdiff(1:ndsts(graph), var_ordering, solved_vars_set)] neweqs = neweqs′ @@ -617,12 +639,21 @@ caller drops the resulting `0`s). After the sort each SCC's variables are contig `findnextfn(v)` identifies the algebraic unknowns (the variables the placeholder fill assigns); `nvars` is `ndsts(graph)`. + +`eq_group[k]`/`eq_elem[k]` identify the array equation (and element within it) that +`neweqs′[k]` is a row of, or `0`. All rows of an array equation are tagged with the earliest +SCC among them and sorted by element, so that they are contiguous and in element order. +Since these rows are all differential equations of selected states, moving them earlier +does not affect the validity of the BLT order for the algebraic equations. """ function blt_reorder_generated_equations!( neweqs′::Vector{Equation}, eq_ordering::Vector{Int}, var_ordering::Vector{Int}, - eq_scc::Vector{Int}, var_sccs::Vector{Vector{Int}}, findnextfn, nvars::Int) + eq_scc::Vector{Int}, var_sccs::Vector{Vector{Int}}, findnextfn, nvars::Int; + eq_group::Vector{Int} = zeros(Int, length(neweqs′)), + eq_elem::Vector{Int} = zeros(Int, length(neweqs′))) n = length(neweqs′) @assert length(eq_ordering) == n && length(var_ordering) == n && length(eq_scc) == n + @assert length(eq_group) == n && length(eq_elem) == n # Position of each variable's SCC in the (topologically sorted) `var_sccs`. scc_pos = zeros(Int, nvars) @@ -660,7 +691,24 @@ function blt_reorder_generated_equations!( eq_scc[k] == typemax(Int) && continue eq_scc[k] = scc_pos[var_ordering[k]] end - perm = sortperm(eq_scc) + if any(!iszero, eq_group) + group_scc = Dict{Int, Int}() + for k in 1:n + g = eq_group[k] + iszero(g) && continue + group_scc[g] = min(get(group_scc, g, typemax(Int)), eq_scc[k]) + end + for k in 1:n + g = eq_group[k] + iszero(g) && continue + eq_scc[k] = group_scc[g] + end + perm = sortperm(1:n; by = k -> (eq_scc[k], eq_group[k], eq_elem[k])) + permute!(eq_group, perm) + permute!(eq_elem, perm) + else + perm = sortperm(eq_scc) + end permute!(neweqs′, perm) permute!(eq_ordering, perm) permute!(var_ordering, perm) @@ -1190,12 +1238,25 @@ function codegen_equation!(eg::EquationGenerator, # the docstring for `add_additional_history!`, this is an exception and needs to be # treated like a solved equation rather than a differential equation. is_highest_diff = iv isa Int && isdervar && var_to_diff[iv] === nothing + group = row_group(state, ieq) if issolvable && isdervar && (!isdisc || !is_highest_diff) var = fullvars[iv] isnothing(D) && throw(UnexpectedDifferentialError(equations(sys)[ieq])) order, lv = var_order(iv, diff_to_var) dx = D(MTKBase.simplify_shifts(fullvars[lv])) - neweq = make_differential_equation(var, dx, eq, total_sub) + if iszero(group) + neweq, _ = make_differential_equation(var, dx, eq, total_sub) + else + # The row stays an element of its array equation only if it is solved for its + # own first-order derivative and previously solved derivatives do not enter its + # RHS (they are not substituted in the array equation). + neweq, subbed = make_differential_equation(var, dx, eq, total_sub; check_sub = true) + arr_group = state.array_groups[group] + if subbed || !isequal(var, dx) || + !isequal(var, arr_group.lhs_vars[state.row_elem[ieq]]) + arr_group.dirty = true + end + end # We will add `neweq.lhs` to `total_sub`, so any equation involving it won't be # incident on it. Remove the edges incident on `iv` from the graph, and add # the replacement vertices from `ieq` so that the incidence is still correct. @@ -1220,6 +1281,7 @@ function codegen_equation!(eg::EquationGenerator, push!(var_ordering, diff_to_var[iv]) push!(eq_scc, scc_idx) elseif issolvable + iszero(group) || (state.array_groups[group].dirty = true) var = fullvars[iv] neweq = make_solved_equation(var, eq, total_sub; simplify) if neweq !== nothing @@ -1233,6 +1295,7 @@ function codegen_equation!(eg::EquationGenerator, push!(solved_vars, iv) end else + iszero(group) || (state.array_groups[group].dirty = true) neweq = make_algebraic_equation(eq, total_sub) # For the same reason as solved equations (they are effectively the same) if isdisc @@ -1261,12 +1324,16 @@ end Generate a first-order differential equation whose LHS is `dx`. `var` and `dx` represent the same variable, but `var` may be a higher-order differential and `dx` is always first-order. For example, if `var` is D(D(x)), then `dx` would be `D(x_t)`. Solve `eq` for `var`, substitute previously solved variables, and return the differential equation. + +Also return whether substituting previously solved variables changed the solved expression. +This is only computed (and otherwise `false`) if `check_sub` is `true`. """ -function make_differential_equation(var, dx, eq, total_sub) +function make_differential_equation(var, dx, eq, total_sub; check_sub::Bool = false) v1 = Symbolics.symbolic_linear_solve(eq, var)::SymbolicT v2 = Symbolics.fixpoint_sub(v1, total_sub, MTKBase.Shift) v3 = MTKBase.simplify_shifts(v2) - dx ~ v3 + subbed = check_sub && !isequal(v1, v2) + (dx ~ v3), subbed end """ @@ -1357,9 +1424,97 @@ function reorder_vars!(state::TearingState, var_eq_matching, var_sccs, eq_orderi state.structure.var_to_diff = new_var_to_diff state.structure.eq_to_diff = new_eq_to_diff state.fullvars = new_fullvars + # Rows of the new graph are the generated equations, in order. + state.row_group = map(Base.Fix1(row_group, state), eq_ordering) + state.row_elem = map(eq_ordering) do ieq + ieq <= length(state.row_elem) ? state.row_elem[ieq] : 0 + end state end +""" + $(TYPEDSIGNATURES) + +Replace runs of generated scalar equations that make up an intact array equation (see +[`ArrayEquationGroup`](@ref)) by the array equation itself. `state.row_group` and +`state.row_elem` must describe the rows of `neweqs` (see [`reorder_vars!`](@ref)); on +return they are updated so that a row is tagged with a group iff it was collapsed into that +group's array equation, which makes [`row_to_equation_indices`](@ref) well-defined. + +A group is collapsed iff it is not dirty, and all of its rows occur consecutively in element +order with each row being the differential equation of its own element. +""" +function collapse_array_equations!(state::TearingState, neweqs::Vector{Equation}) + (; array_groups, row_group, row_elem) = state + n = length(neweqs) + @assert length(row_group) == n && length(row_elem) == n + any(!iszero, row_group) || return neweqs + + out = Equation[] + sizehint!(out, n) + new_row_group = zeros(Int, n) + new_row_elem = zeros(Int, n) + i = 1 + while i <= n + g = row_group[i] + if iszero(g) || array_groups[g].dirty + push!(out, neweqs[i]) + i += 1 + continue + end + grp = array_groups[g] + m = length(grp.lhs_vars) + intact = i + m - 1 <= n + if intact + for k in 1:m + r = i + k - 1 + if row_group[r] != g || row_elem[r] != k || + !isequal(neweqs[r].lhs, grp.lhs_vars[k]) + intact = false + break + end + end + end + if !intact + grp.dirty = true + push!(out, neweqs[i]) + i += 1 + continue + end + push!(out, grp.eq) + new_row_group[i:(i + m - 1)] .= g + new_row_elem[i:(i + m - 1)] .= 1:m + i += m + end + state.row_group = new_row_group + state.row_elem = new_row_elem + return out +end + +""" + $(TYPEDSIGNATURES) + +For each equation row of `state.structure.graph` of a simplified `state`, the index of the +equation in `equations(state.sys)` it belongs to. This is the identity unless array +equations were preserved, in which case all rows of an array equation map to it. +""" +function row_to_equation_indices(state::TearingState) + (; row_group) = state + n = length(row_group) + idxs = Vector{Int}(undef, n) + ieq = 0 + prev = 0 + for i in 1:n + g = row_group[i] + if iszero(g) || g != prev + ieq += 1 + end + idxs[i] = ieq + prev = g + end + return idxs +end + """ Update the system equations, unknowns, and observables after simplification. """ @@ -1367,7 +1522,8 @@ function update_simplified_system!( state::TearingState, neweqs::Vector{Equation}, solved_eqs::Vector{Equation}, dummy_sub::Dict{SymbolicT, SymbolicT}, var_sccs::Vector{Vector{Int}}, extra_unknowns::Vector{SymbolicT}, iv::Union{SymbolicT, Nothing}, - D::Union{Differential, Shift, Nothing}; array_hack = true) + D::Union{Differential, Shift, Nothing}; array_hack = true, + preserve_array_equations = false) (; fullvars, structure, sys) = state (; solvable_graph, var_to_diff, eq_to_diff, graph) = structure @@ -1445,6 +1601,14 @@ function update_simplified_system!( sys = MTKBase.with_reversible_transformation(sys, tf) end + if preserve_array_equations && !StateSelection.is_only_discrete(structure) + neweqs = collapse_array_equations!(state, neweqs) + else + # No row of the simplified system stands for an array equation. + fill!(state.row_group, 0) + fill!(state.row_elem, 0) + end + @set! sys.eqs = neweqs @set! sys.observed = obs @@ -1523,6 +1687,16 @@ $TYPEDFIELDS equations which is solved symbolically rather than using `LinearSolve.jl`. """ analytical_linear_scc_limit::Int = 2 + """ + Whether array differential equations `D(x[slice]) ~ rhs` that survive structural + simplification as a block (see [`ArrayEquationGroup`](@ref)) are emitted as single array + equations instead of being scalarized. The unknowns of the simplified system remain + scalar; the elements of `x[slice]` occupy the unknown slots of the array equation's row + range, in element order. Requires code generation to support array equations. Can also + be passed as the `preserve_array_equations` keyword argument to the algorithm + (and hence to `mtkcompile`). + """ + preserve_array_equations::Bool = false end function (alg::DefaultReassembleAlgorithm)(state::TearingState, @@ -1530,7 +1704,9 @@ function (alg::DefaultReassembleAlgorithm)(state::TearingState, mm::Union{CLIL.SparseMatrixCLIL, Nothing}; fully_determined::Bool = true, allow_symbolic::Bool = false, - allow_parameter::Bool = true, kw...) + allow_parameter::Bool = true, + preserve_array_equations::Bool = alg.preserve_array_equations, + kw...) (; simplify, array_hack, inline_linear_sccs, analytical_linear_scc_limit) = alg (; var_eq_matching, full_var_eq_matching, var_sccs) = tearing_result @@ -1586,7 +1762,7 @@ function (alg::DefaultReassembleAlgorithm)(state::TearingState, # var_eq_matching and full_var_eq_matching are now invalidated sys = update_simplified_system!(state, neweqs, solved_eqs, dummy_sub, var_sccs, - extra_unknowns, iv, D; array_hack) + extra_unknowns, iv, D; array_hack, preserve_array_equations) else D = D::Nothing neweqs, solved_eqs, diff --git a/lib/ModelingToolkitTearing/src/stateselection_interface.jl b/lib/ModelingToolkitTearing/src/stateselection_interface.jl index 4b1e8ef..9a84a8c 100644 --- a/lib/ModelingToolkitTearing/src/stateselection_interface.jl +++ b/lib/ModelingToolkitTearing/src/stateselection_interface.jl @@ -68,6 +68,9 @@ StateSelection.get_mm(ts::TearingState) = ts.mm function StateSelection.eq_derivative!(ts::TearingState, ieq::Int; kwargs...) s = ts.structure + # Index reduction differentiating an element of an array equation means the block is + # part of a higher-index structure; emit it scalarized. + dirty_array_group!(ts, ieq) eq_diff = StateSelection.eq_derivative_graph!(s, ieq) mm = ts.mm @@ -185,6 +188,56 @@ function StateSelection.linear_subsys_adjmat!(state::TearingState; kwargs...) return mm end +""" + $TYPEDSIGNATURES + +Split `mm` into the rows that are elements of array equations (see +[`ArrayEquationGroup`](@ref)) and the rest. Returns `(rest, group_rows)` where `rest` is a +`SparseMatrixCLIL` over the same parent rows and columns. Used with +[`merge_array_group_rows`](@ref) to keep array equation rows out of Gaussian elimination +of the integer-linear subsystem: an element of an array equation must neither be replaced +by a linear combination of rows nor be used as a pivot to rewrite other rows (which would +leave the element's derivative matched to a different equation), or the array equation +can no longer be emitted as a unit. Excluding the rows is conservative: they are simply +not part of the linear subsystem the elimination works on. +""" +function split_array_group_rows(state::TearingState, mm::CLIL.SparseMatrixCLIL{Int, Int}) + empty = CLIL.SparseMatrixCLIL(mm.nparentrows, mm.ncols, Int[], Vector{Int}[], Vector{Int}[]) + isempty(state.array_groups) && return mm, empty + keep = Int[] + split = Int[] + for (i, e) in enumerate(mm.nzrows) + push!(iszero(row_group(state, e)) ? keep : split, i) + end + isempty(split) && return mm, empty + rest = CLIL.SparseMatrixCLIL( + mm.nparentrows, mm.ncols, mm.nzrows[keep], mm.row_cols[keep], mm.row_vals[keep] + ) + group_rows = CLIL.SparseMatrixCLIL( + mm.nparentrows, mm.ncols, mm.nzrows[split], mm.row_cols[split], mm.row_vals[split] + ) + return rest, group_rows +end + +""" + $TYPEDSIGNATURES + +Inverse of [`split_array_group_rows`](@ref): add the rows of `group_rows` back to `mm`, +keeping `mm.nzrows` sorted. +""" +function merge_array_group_rows( + mm::CLIL.SparseMatrixCLIL{T, Int}, group_rows::CLIL.SparseMatrixCLIL{Int, Int} + ) where {T} + isempty(group_rows.nzrows) && return mm + nzrows = vcat(mm.nzrows, group_rows.nzrows) + row_cols = vcat(mm.row_cols, group_rows.row_cols) + row_vals = Vector{T}[mm.row_vals; map(Base.Fix1(convert, Vector{T}), group_rows.row_vals)] + perm = sortperm(nzrows) + return CLIL.SparseMatrixCLIL( + mm.nparentrows, mm.ncols, nzrows[perm], row_cols[perm], row_vals[perm] + ) +end + function maybe_zeros_descend(ex::SymbolicT) @match ex begin BSImpl.AddMul(; variant) => return variant === SU.AddMulVariant.MUL @@ -464,6 +517,12 @@ function StateSelection.rm_eqs_vars!( if !isempty(state.eqs_source) deleteat!(state.eqs_source, eqs_to_rm) end + # `eqs_to_rm` was sorted and uniqued in place by `default_rm_eqs_vars!`. + for ieq in eqs_to_rm + dirty_array_group!(state, ieq) + end + deleteat!(state.row_group, eqs_to_rm) + deleteat!(state.row_elem, eqs_to_rm) @set! sys.eqs = eqs state.sys = sys diff --git a/lib/ModelingToolkitTearing/src/tearingstate.jl b/lib/ModelingToolkitTearing/src/tearingstate.jl index 7fcf295..7e02860 100644 --- a/lib/ModelingToolkitTearing/src/tearingstate.jl +++ b/lib/ModelingToolkitTearing/src/tearingstate.jl @@ -5,6 +5,11 @@ # - Is it updated in `eq_derivative!`? (if necessary) # - Is it updated in `rm_eqs_vars!`? (if necessary) # - Is it updated in `scalarize_tearing_state_eqs!`? (if necessary) +# - Is it updated in `system_subset`? (if necessary) +# +# NOTE: Checklist for passes that rewrite equations in a `TearingState` +# - If the rewrite is anything other than substituting a variable that is retained as an +# observed equation (alias/zero elimination), call `dirty_array_group!` for the row. """ $TYPEDEF @@ -49,6 +54,39 @@ end StateSelection.is_only_discrete(s::SystemStructure) = s.only_discrete +""" + $TYPEDEF + +An array-valued differential equation `D(x[slice]) ~ rhs` that is tracked as a unit through +structural simplification. Its scalarized elements are ordinary rows of the bipartite graph +(so matching, Pantelides, dummy derivatives, tearing and alias elimination see exact +per-element incidence), but the rows remember which array equation they came from. If no +pass needs to break the block — no row is differentiated or removed, no row is rewritten by +anything other than an alias/zero substitution whose eliminated variable is retained as an +observed equation, and every row ends up as the differential equation of its own element — +the array equation is emitted intact by [`update_simplified_system!`](@ref) instead of as +`length(slice)` scalar equations. + +# Fields + +$TYPEDFIELDS +""" +mutable struct ArrayEquationGroup + """The array equation in canonical form: `lhs` is `D(x[slice])`.""" + eq::Equation + """The scalar derivatives `D(x[k])`, in the order in which the rows were scalarized.""" + lhs_vars::Vector{SymbolicT} + """ + Whether a transformation broke the block, in which case the rows are emitted as scalar + equations. + """ + dirty::Bool +end + +function ArrayEquationGroup(eq::Equation, lhs_vars::Vector{SymbolicT}) + return ArrayEquationGroup(eq, lhs_vars, false) +end + """ $TYPEDEF @@ -112,12 +150,92 @@ mutable struct TearingState <: StateSelection.TransformationState{System} and put into `additional_observed`. """ analytical_derivatives::Dict{SymbolicT, SymbolicT} + """ + Array equations tracked as blocks of scalar rows. See [`ArrayEquationGroup`](@ref). + """ + array_groups::Vector{ArrayEquationGroup} + """ + For each equation row of `structure.graph`, the index into `array_groups` of the array + equation the row was scalarized from, or `0` for scalar equations. Rows appended after + construction (e.g. by `eq_derivative!`) are always scalar. + """ + row_group::Vector{Int} + """ + For each equation row, the (linear) element index of the row within its array equation. + `0` for scalar equations. + """ + row_elem::Vector{Int} +end + +function TearingState( + sys::System, fullvars::Vector{SymbolicT}, structure::SystemStructure, + extra_eqs::Vector{Equation}, param_derivative_map::Dict{SymbolicT, SymbolicT}, + no_deriv_params::Set{SymbolicT}, original_eqs::Vector{Equation}, + additional_observed::Vector{Equation}, always_present::BitVector, + statemachines::Vector{System}, eqs_source::Vector{Vector{Symbol}}, + mm::Union{Nothing, CLIL.SparseMatrixCLIL{Int, Int}}, + analytical_derivatives::Dict{SymbolicT, SymbolicT} + ) + neqs = nsrcs(structure.graph) + return TearingState( + sys, fullvars, structure, extra_eqs, param_derivative_map, no_deriv_params, + original_eqs, additional_observed, always_present, statemachines, eqs_source, mm, + analytical_derivatives, ArrayEquationGroup[], zeros(Int, neqs), zeros(Int, neqs) + ) end function Base.show(io::IO, state::TearingState) print(io, "TearingState of ", typeof(state.sys)) end +""" + $TYPEDSIGNATURES + +Index into `ts.array_groups` of the array equation that equation row `ieq` belongs to, or +`0` if it is a scalar equation (including rows appended after construction). +""" +function row_group(ts::TearingState, ieq::Int) + rg = ts.row_group + return ieq <= length(rg) ? rg[ieq] : 0 +end + +""" + $TYPEDSIGNATURES + +Mark the array equation that row `ieq` belongs to (if any) as broken, so that its rows are +emitted as scalar equations. +""" +function dirty_array_group!(ts::TearingState, ieq::Int) + g = row_group(ts, ieq) + iszero(g) && return false + ts.array_groups[g].dirty = true + return true +end + +""" + $TYPEDSIGNATURES + +Number of rows tagged with each of the `ngroups` array groups in `row_group`. +""" +function count_group_rows(ngroups::Int, row_group::Vector{Int}) + counts = zeros(Int, ngroups) + for g in row_group + iszero(g) || (counts[g] += 1) + end + return counts +end + +""" + $TYPEDSIGNATURES + +Append `n` scalar rows to the row-to-group bookkeeping of `ts`. +""" +function push_scalar_rows!(ts::TearingState, n::Int = 1) + append!(ts.row_group, Iterators.repeated(0, n)) + append!(ts.row_elem, Iterators.repeated(0, n)) + return ts +end + StateSelection.has_equations(::TearingState) = true StateSelection.equations(ts::TearingState) = equations(ts) @@ -144,6 +262,8 @@ function Base.setindex!(ev::EquationsView, v::Equation, i::Integer) end function Base.push!(ev::EquationsView, eq) push!(ev.ts.extra_eqs, eq) + push_scalar_rows!(ev.ts) + return ev end function TearingState(sys::System, source_info::Union{Nothing, MTKBase.EquationSourceInformation} = nothing; check::Bool = true, sort_eqs::Bool = true, defer_scalarization::Bool = false) @@ -160,6 +280,9 @@ function TearingState(sys::System, source_info::Union{Nothing, MTKBase.EquationS MTKBase.check_no_parameter_equations(sys) iv = MTKBase.get_iv(sys) sources = Vector{Vector{Symbol}}() + array_groups = ArrayEquationGroup[] + row_group = Int[] + row_elem = Int[] # flatten array equations if defer_scalarization # Don't scalarize eagerly — defer scalarization to after clock inference (via @@ -173,21 +296,28 @@ function TearingState(sys::System, source_info::Union{Nothing, MTKBase.EquationS sources = Vector{Vector{Symbol}}(source_info.eqs_source) end eqs = Vector{Equation}(equations(sys)) + resize!(row_group, length(eqs)) + fill!(row_group, 0) + resize!(row_elem, length(eqs)) + fill!(row_elem, 0) else if source_info !== nothing @assert length(equations(sys)) == length(source_info.eqs_source) """ Mismatch between source information provided to `TearingState` and the structure \ of the system. """ - # Eager scalarization: expand each equation and replicate its source entry. - for (eq, src) in zip(equations(sys), source_info.eqs_source) - scal_eq = MTKBase.flatten_equation(eq) - for _ in scal_eq - push!(sources, src) + end + # Eager scalarization: expand each equation (tracking array equations that can stay + # atomic) and replicate its source entry. + eqs = Equation[] + for (i, eq) in enumerate(equations(sys)) + nrows = append_equation_rows!(eqs, row_group, row_elem, array_groups, eq, iv) + if source_info !== nothing + for _ in 1:nrows + push!(sources, source_info.eqs_source[i]) end end end - eqs = MTKBase.flatten_equations(equations(sys)) init_eqs = MTKBase.flatten_equations(initialization_equations(sys)) @set! sys.initialization_eqs = init_eqs end @@ -425,8 +555,18 @@ function TearingState(sys::System, source_info::Union{Nothing, MTKBase.EquationS end filter!(Base.Fix2(!==, MTKBase.COMMON_NOTHING) ∘ last, param_derivative_map) + if !all(eqs_to_retain) + # Removing a row of an array equation (a parameter derivative equation can't be one, + # but be safe) breaks the block. + for i in findall(!, eqs_to_retain) + g = row_group[i] + iszero(g) || (array_groups[g].dirty = true) + end + end eqs = eqs[eqs_to_retain] original_eqs = original_eqs[eqs_to_retain] + row_group = row_group[eqs_to_retain] + row_elem = row_elem[eqs_to_retain] neqs = length(eqs) symbolic_incidence = symbolic_incidence[eqs_to_retain] if !isempty(sources) @@ -462,6 +602,8 @@ function TearingState(sys::System, source_info::Union{Nothing, MTKBase.EquationS eqs = eqs[sortidxs] original_eqs = original_eqs[sortidxs] symbolic_incidence = symbolic_incidence[sortidxs] + row_group = row_group[sortidxs] + row_elem = row_elem[sortidxs] if !isempty(sources) sources = sources[sortidxs] end @@ -510,7 +652,147 @@ function TearingState(sys::System, source_info::Union{Nothing, MTKBase.EquationS canonical_ranks, false) return TearingState(sys, fullvars, structure, Equation[], param_derivative_map, no_deriv_params, original_eqs, Equation[], falses(length(fullvars)), - typeof(sys)[], sources, nothing, Dict{SymbolicT, SymbolicT}()) + typeof(sys)[], sources, nothing, Dict{SymbolicT, SymbolicT}(), + array_groups, row_group, row_elem) +end + +""" + $TYPEDSIGNATURES + +If `eq` is an array equation with a first-order derivative (with respect to `iv`) of an +array unknown or a constant-index slice of one on exactly one side, return it oriented as +`D(x[slice]) ~ rhs`. Return `nothing` for scalar equations and for array equations that +cannot be tracked as a block (e.g. `D(x) .+ D(y) ~ 0`, `0 ~ f(x)`). +""" +function canonicalize_array_equation(eq::Equation, iv::SymbolicT) + (; lhs, rhs) = eq + SU.is_array_shape(SU.shape(lhs)) || return nothing + if is_array_unknown_derivative(lhs, iv) + is_array_unknown_derivative(rhs, iv) && return nothing + return eq + elseif is_array_unknown_derivative(rhs, iv) + return rhs ~ lhs + end + # Residual forms `resid ~ 0` / `0 ~ resid`, as emitted by finite-difference + # discretizations: `resid` is a broadcasted sum/difference with the derivative as one + # of its two operands. + if is_zero_array(rhs) + resid = lhs + elseif is_zero_array(lhs) + resid = rhs + else + return nothing + end + iscall(resid) || return nothing + operation(resid) === broadcast || return nothing + args = arguments(resid) + length(args) == 3 || return nothing + op = args[1] + SU.isconst(op) || return nothing + op = unwrap_const(op) + A, B = args[2], args[3] + isA = is_array_unknown_derivative(A, iv) + isB = is_array_unknown_derivative(B, iv) + isA == isB && return nothing + if op === (-) + # `D(x) - f = 0` or `f - D(x) = 0` + return isA ? (A ~ B) : (B ~ A) + elseif op === (+) + # `D(x) + f = 0` + f = isA ? B : A + return (isA ? A : B) ~ unwrap(-Symbolics.wrap(f)) + end + return nothing +end + +""" + $TYPEDSIGNATURES + +Whether `x` is a constant array (or scalar) of zeros. +""" +function is_zero_array(x::SymbolicT) + @match x begin + BSImpl.Const(; val) => val isa AbstractArray ? all(iszero, val) : SU._iszero(x) + _ => false + end +end + +""" + $TYPEDSIGNATURES + +Whether `x` is `D(arr)` where `D` is the first-order derivative with respect to `iv` and +`arr` is an array variable or a constant-index slice of one. +""" +function is_array_unknown_derivative(x::SymbolicT, iv::SymbolicT) + SU.is_array_shape(SU.shape(x)) || return false + iscall(x) || return false + f = operation(x) + f isa Differential || return false + isequal(f.x, iv) && f.order isa Int && isone(f.order) || return false + args = arguments(x) + length(args) == 1 || return false + arg = args[1] + SU.shape(arg) isa SU.ShapeVecT || return false + arr, isidx = MTKBase.split_indexed_var(arg) + if isidx + # `x[slice]`: every index must be constant so the elements are known statically + iscall(arg) || return false + operation(arg) === getindex || return false + all(SU.isconst, Iterators.drop(arguments(arg), 1)) || return false + return true + end + return MTKBase.isvariable(arg) +end + +""" + $TYPEDSIGNATURES + +Scalarize `eq` and append the resulting rows to `rows`, extending `row_group` and +`row_elem` in parallel. If `eq` can be tracked as an [`ArrayEquationGroup`](@ref) (see +[`canonicalize_array_equation`](@ref)), register it in `groups` and tag its rows. Return +the number of rows appended. +""" +function append_equation_rows!( + rows::Vector{Equation}, row_group::Vector{Int}, row_elem::Vector{Int}, + groups::Vector{ArrayEquationGroup}, eq::Equation, @nospecialize(iv::Union{SymbolicT, Nothing}) + ) + scalar_eqs = MTKBase.flatten_equation(eq) + n = length(scalar_eqs) + append!(rows, scalar_eqs) + canon = if iv isa SymbolicT && SU.is_array_shape(SU.shape(eq.lhs)) + canonicalize_array_equation(eq, iv) + else + nothing + end + if canon !== nothing + lhs_vars = vec(collect(collect(canon.lhs)::AbstractArray{SymbolicT}))::Vector{SymbolicT} + # Row `k` is the equation of element `k` (`flatten_equation` and `collect` use the + # same element order), and the elements must be distinct. + ok = length(lhs_vars) == n && allunique(lhs_vars) + # For explicit forms, additionally check that the scalarized rows are indeed + # `D(x[k]) ~ rhs[k]` or `rhs[k] ~ D(x[k])`. Residual forms are verified when the + # rows are matched: a row not solved for its own element breaks the group. + explicit = canon.lhs === eq.lhs || canon.lhs === eq.rhs + if ok && explicit + for k in 1:n + seq = scalar_eqs[k] + if !(isequal(seq.lhs, lhs_vars[k]) || isequal(seq.rhs, lhs_vars[k])) + ok = false + break + end + end + end + if ok + push!(groups, ArrayEquationGroup(canon, lhs_vars)) + g = length(groups) + append!(row_group, Iterators.repeated(g, n)) + append!(row_elem, 1:n) + return n + end + end + append!(row_group, Iterators.repeated(0, n)) + append!(row_elem, Iterators.repeated(0, n)) + return n end """ @@ -736,9 +1018,12 @@ function scalarize_tearing_state_eqs!(ts::TearingState) # Early exit has_arr_eqs = any(eq -> SU.is_array_shape(SU.shape(eq.lhs)), arr_eqs) has_arr_eqs || iszero(Graphs.ne(ts.structure.graph)) || return ts + # Already scalarized (eagerly, with array equations tracked as groups). + has_arr_eqs || isempty(ts.array_groups) || return ts arr_orig_eqs = ts.original_eqs eqs_source = ts.eqs_source + iv = MTKBase.get_iv(ts.sys) new_eqs = Equation[] sizehint!(new_eqs, length(arr_eqs)) @@ -746,14 +1031,19 @@ function scalarize_tearing_state_eqs!(ts::TearingState) sizehint!(new_orig, length(arr_orig_eqs)) new_sources = Vector{Vector{Symbol}}() sizehint!(new_sources, length(eqs_source)) + # Rows of the deferred graph are unit equations and are never grouped, so the + # bookkeeping is rebuilt from scratch here. + array_groups = ArrayEquationGroup[] + row_group = Int[] + row_elem = Int[] for i in eachindex(arr_eqs) - scalar_eqs = MTKBase.flatten_equation(arr_eqs[i]) + nrows = append_equation_rows!(new_eqs, row_group, row_elem, array_groups, arr_eqs[i], iv) scalar_orig_eqs = MTKBase.flatten_equation(arr_orig_eqs[i]) - append!(new_eqs, scalar_eqs) + @assert length(scalar_orig_eqs) == nrows append!(new_orig, scalar_orig_eqs) if !isempty(eqs_source) - for _ in scalar_orig_eqs + for _ in 1:nrows push!(new_sources, eqs_source[i]) end end @@ -802,6 +1092,9 @@ function scalarize_tearing_state_eqs!(ts::TearingState) ts.sys = sys ts.original_eqs = new_orig ts.eqs_source = new_sources + ts.array_groups = array_groups + ts.row_group = row_group + ts.row_elem = row_elem structure = ts.structure @set! structure.graph = complete(graph) @@ -1096,5 +1389,9 @@ function shift_discrete_system(ts::TearingState) @set! ts.sys.eqs = eqs @set! ts.fullvars = fullvars + # Array blocks are only tracked for continuous differential equations. + for g in ts.array_groups + g.dirty = true + end return ts end diff --git a/lib/ModelingToolkitTearing/test/runtests.jl b/lib/ModelingToolkitTearing/test/runtests.jl index 92c46cb..1e1186b 100644 --- a/lib/ModelingToolkitTearing/test/runtests.jl +++ b/lib/ModelingToolkitTearing/test/runtests.jl @@ -588,3 +588,121 @@ end MTKTearing.scalarize_tearing_state_eqs!(tss[cid]) @test !iszero(Graphs.ne(tss[cid].structure.graph)) end + +@testset "Array equation groups" begin + @testset "`TearingState` tracks array equations as groups of rows" begin + @variables x(t)[1:3] y(t) + @named sys = System([D(x) ~ -x .+ y, y ~ sum(x)], t) + ts = TearingState(sys) + # Scalar rows are still what the structure sees. + @test length(equations(ts)) == 4 + @test length(ts.array_groups) == 1 + grp = only(ts.array_groups) + @test !grp.dirty + @test isequal(grp.eq.lhs, D(x)) + @test isequal(grp.lhs_vars, [D(x[1]), D(x[2]), D(x[3])]) + @test length(ts.row_group) == length(ts.row_elem) == 4 + rows = findall(==(1), ts.row_group) + @test length(rows) == 3 + @test ts.row_elem[rows] == 1:3 + @test iszero(ts.row_group[only(setdiff(1:4, rows))]) + for (k, r) in enumerate(rows) + @test isequal(equations(ts)[r].lhs, D(x[k])) + end + end + + @testset "residual forms are canonicalized" begin + @variables u(t)[1:5] + lap = u[1:3] .- 2 .* u[2:4] .+ u[3:5] + @named sys = System([broadcast(-, D(u[2:4]), lap) ~ zeros(3), u[1] ~ 0, u[5] ~ 0], t, collect(u), []) + ts = TearingState(sys) + grp = only(ts.array_groups) + @test isequal(grp.eq.lhs, D(u[2:4])) + @test isequal(grp.eq.rhs, unwrap(lap)) + @named sys = System([broadcast(+, D(u[2:4]), lap) ~ zeros(3), u[1] ~ 0, u[5] ~ 0], t, collect(u), []) + ts = TearingState(sys) + grp = only(ts.array_groups) + @test isequal(grp.eq.lhs, D(u[2:4])) + @test isequal(grp.eq.rhs, unwrap(-lap)) + end + + @testset "ineligible array equations are not grouped" begin + @variables x(t)[1:3] y(t)[1:3] + # Algebraic array equation. + @named sys = System([D(x) ~ y, zeros(3) ~ x .+ y], t) + ts = TearingState(sys) + @test length(ts.array_groups) == 1 + @test isequal(only(ts.array_groups).eq.lhs, D(x)) + # Derivatives on both sides. + @named sys = System([D(x) ~ D(y), D(y) ~ -y], t) + ts = TearingState(sys) + @test length(ts.array_groups) == 1 + @test isequal(only(ts.array_groups).eq.lhs, D(y)) + end + + @testset "differentiating a row dirties its group" begin + @variables x(t)[1:3] y(t) + @named sys = System([D(x) ~ -x .+ y, y ~ sum(x)], t) + ts = TearingState(sys) + StateSelection.complete!(ts.structure) + r = findfirst(==(1), ts.row_group) + for v in BipartiteGraphs.𝑠neighbors(ts.structure.graph, r) + ts.structure.var_to_diff[v] === nothing || continue + StateSelection.var_derivative!(ts, v) + end + StateSelection.eq_derivative!(ts, r) + @test only(ts.array_groups).dirty + # The appended derivative row is scalar. + @test length(ts.row_group) == length(ts.row_elem) == length(equations(ts)) + @test iszero(ts.row_group[end]) + end + + @testset "alias elimination does not pivot on array equation rows" begin + @variables x(t)[1:3] y(t) z(t) + # `y ~ sum(x)` is linear in `y` and the elements of `x`; Gaussian elimination of + # the linear subsystem must not use `D(x[1]) ~ -x[1] + y` to eliminate `y` from it. + @named sys = System([D(x) ~ -x .+ y, y ~ sum(x), 0 ~ z^3 + z - y], t) + ts = TearingState(sys) + ModelingToolkit.alias_elimination!(ts) + @test !only(ts.array_groups).dirty + for (r, g) in enumerate(ts.row_group) + iszero(g) && continue + @test isequal(equations(ts)[r].lhs, D(x[ts.row_elem[r]])) + end + end + + @testset "`preserve_array_equations` emits intact groups" begin + @variables x(t)[1:3] y(t) z(t) + @named sys = System([D(x) ~ -x .+ y, y ~ sum(x), 0 ~ z^3 + z - y], t) + ssys = mtkcompile(sys; preserve_array_equations = true) + eqs = equations(ssys) + @test length(eqs) == 2 + arr = findfirst(eq -> SU.is_array_shape(SU.shape(eq.lhs)), eqs) + @test arr !== nothing + @test isequal(eqs[arr].lhs, D(x)) + @test issetequal(unknowns(ssys), [x[1], x[2], x[3], z]) + ts = ModelingToolkit.get_tearing_state(ssys) + @test MTKTearing.row_to_equation_indices(ts) == (arr == 1 ? [1, 1, 1, 2] : [1, 2, 2, 2]) + + # Through the algorithm object. + alg = MTKTearing.DefaultReassembleAlgorithm(; preserve_array_equations = true) + ssys = mtkcompile(sys; reassemble_alg = alg) + @test count(eq -> SU.is_array_shape(SU.shape(eq.lhs)), equations(ssys)) == 1 + + # Default: scalarized, as before. + ssys = mtkcompile(sys) + @test length(equations(ssys)) == 4 + @test !any(eq -> SU.is_array_shape(SU.shape(eq.lhs)), equations(ssys)) + end + + @testset "dirty groups are emitted scalarized" begin + @variables x(t)[1:3] y(t) T(t) + # `x[2] ~ y` constrains two differential variables: index reduction differentiates + # the constraint and `D(x[2])` is no longer solved from its array equation. + @named sys = System([D(x) ~ -x .+ [0, T, 0], 0 ~ x[2] - y, D(y) ~ -y], t) + ssys = mtkcompile(sys; preserve_array_equations = true) + @test !any(eq -> SU.is_array_shape(SU.shape(eq.lhs)), equations(ssys)) + ts = ModelingToolkit.get_tearing_state(ssys) + @test all(g -> g.dirty, ts.array_groups) + end +end From 1c2810702822aca896c91aa189a252f708aabc68 Mon Sep 17 00:00:00 2001 From: ChrisRackauckas-Claude Date: Wed, 9 Sep 2026 18:48:05 -0400 Subject: [PATCH 2/2] MTKTearing: exclude intact array-group rows from linear subsystem at collection Fold the split/merge helpers into linear_subsys_adjmat! via is_intact_array_group_row, and cover the mixed array-DE + linear algebraic case. Co-authored-by: Cursor --- .../src/stateselection_interface.jl | 60 +++---------------- .../src/tearingstate.jl | 10 ++++ lib/ModelingToolkitTearing/test/runtests.jl | 17 ++++++ 3 files changed, 36 insertions(+), 51 deletions(-) diff --git a/lib/ModelingToolkitTearing/src/stateselection_interface.jl b/lib/ModelingToolkitTearing/src/stateselection_interface.jl index 9a84a8c..d3686c0 100644 --- a/lib/ModelingToolkitTearing/src/stateselection_interface.jl +++ b/lib/ModelingToolkitTearing/src/stateselection_interface.jl @@ -175,7 +175,15 @@ function StateSelection.linear_subsys_adjmat!(state::TearingState; kwargs...) # ``∑ c_i * v_i = 0``, # # where ``c_i`` ∈ ℤ and ``v_i`` denotes unknowns. - if all_int_vars && Symbolics._iszero(rhs) + # + # Rows that are elements of an intact array equation (see `ArrayEquationGroup`) + # are not part of the integer-linear subsystem, even when they are integer-linear. + # Gaussian elimination of that subsystem (alias elimination, singularity removal, + # exact SCC matching) may replace a row by a linear combination of rows or use it as + # a pivot to rewrite other rows; either leaves the element's derivative matched to a + # different equation than its own, and the array equation can no longer be emitted + # as a unit. The rows keep their solvability information in `solvable_graph`. + if all_int_vars && Symbolics._iszero(rhs) && !is_intact_array_group_row(state, i) push!(linear_equations, i) push!(eadj, copy(𝑠neighbors(graph, i))) push!(cadj, copy(coeffs)) @@ -188,56 +196,6 @@ function StateSelection.linear_subsys_adjmat!(state::TearingState; kwargs...) return mm end -""" - $TYPEDSIGNATURES - -Split `mm` into the rows that are elements of array equations (see -[`ArrayEquationGroup`](@ref)) and the rest. Returns `(rest, group_rows)` where `rest` is a -`SparseMatrixCLIL` over the same parent rows and columns. Used with -[`merge_array_group_rows`](@ref) to keep array equation rows out of Gaussian elimination -of the integer-linear subsystem: an element of an array equation must neither be replaced -by a linear combination of rows nor be used as a pivot to rewrite other rows (which would -leave the element's derivative matched to a different equation), or the array equation -can no longer be emitted as a unit. Excluding the rows is conservative: they are simply -not part of the linear subsystem the elimination works on. -""" -function split_array_group_rows(state::TearingState, mm::CLIL.SparseMatrixCLIL{Int, Int}) - empty = CLIL.SparseMatrixCLIL(mm.nparentrows, mm.ncols, Int[], Vector{Int}[], Vector{Int}[]) - isempty(state.array_groups) && return mm, empty - keep = Int[] - split = Int[] - for (i, e) in enumerate(mm.nzrows) - push!(iszero(row_group(state, e)) ? keep : split, i) - end - isempty(split) && return mm, empty - rest = CLIL.SparseMatrixCLIL( - mm.nparentrows, mm.ncols, mm.nzrows[keep], mm.row_cols[keep], mm.row_vals[keep] - ) - group_rows = CLIL.SparseMatrixCLIL( - mm.nparentrows, mm.ncols, mm.nzrows[split], mm.row_cols[split], mm.row_vals[split] - ) - return rest, group_rows -end - -""" - $TYPEDSIGNATURES - -Inverse of [`split_array_group_rows`](@ref): add the rows of `group_rows` back to `mm`, -keeping `mm.nzrows` sorted. -""" -function merge_array_group_rows( - mm::CLIL.SparseMatrixCLIL{T, Int}, group_rows::CLIL.SparseMatrixCLIL{Int, Int} - ) where {T} - isempty(group_rows.nzrows) && return mm - nzrows = vcat(mm.nzrows, group_rows.nzrows) - row_cols = vcat(mm.row_cols, group_rows.row_cols) - row_vals = Vector{T}[mm.row_vals; map(Base.Fix1(convert, Vector{T}), group_rows.row_vals)] - perm = sortperm(nzrows) - return CLIL.SparseMatrixCLIL( - mm.nparentrows, mm.ncols, nzrows[perm], row_cols[perm], row_vals[perm] - ) -end - function maybe_zeros_descend(ex::SymbolicT) @match ex begin BSImpl.AddMul(; variant) => return variant === SU.AddMulVariant.MUL diff --git a/lib/ModelingToolkitTearing/src/tearingstate.jl b/lib/ModelingToolkitTearing/src/tearingstate.jl index 7e02860..70ff8e8 100644 --- a/lib/ModelingToolkitTearing/src/tearingstate.jl +++ b/lib/ModelingToolkitTearing/src/tearingstate.jl @@ -199,6 +199,16 @@ function row_group(ts::TearingState, ieq::Int) return ieq <= length(rg) ? rg[ieq] : 0 end +""" + $TYPEDSIGNATURES + +Whether row `ieq` is an element of an array equation that has not been marked dirty. +""" +function is_intact_array_group_row(ts::TearingState, ieq::Int) + g = row_group(ts, ieq) + return !iszero(g) && !ts.array_groups[g].dirty +end + """ $TYPEDSIGNATURES diff --git a/lib/ModelingToolkitTearing/test/runtests.jl b/lib/ModelingToolkitTearing/test/runtests.jl index 1e1186b..1ab60b5 100644 --- a/lib/ModelingToolkitTearing/test/runtests.jl +++ b/lib/ModelingToolkitTearing/test/runtests.jl @@ -662,6 +662,23 @@ end # `y ~ sum(x)` is linear in `y` and the elements of `x`; Gaussian elimination of # the linear subsystem must not use `D(x[1]) ~ -x[1] + y` to eliminate `y` from it. @named sys = System([D(x) ~ -x .+ y, y ~ sum(x), 0 ~ z^3 + z - y], t) + ts = TearingState(sys) + mm = StateSelection.linear_subsys_adjmat!(ts) + group_rows = findall(!iszero, ts.row_group) + @test length(mm.nzrows) == 1 + @test isempty(intersect(mm.nzrows, group_rows)) + # The rows are still solvable for their derivatives. + sgraph = ts.structure.solvable_graph + for r in group_rows + dv = findfirst(isequal(D(x[ts.row_elem[r]])), ts.fullvars) + @test dv in BipartiteGraphs.𝑠neighbors(sgraph, r) + end + # Dirty groups are ordinary scalar rows and do take part. + ts = TearingState(sys) + MTKTearing.dirty_array_group!(ts, first(findall(!iszero, ts.row_group))) + mm = StateSelection.linear_subsys_adjmat!(ts) + @test length(mm.nzrows) == 4 + ts = TearingState(sys) ModelingToolkit.alias_elimination!(ts) @test !only(ts.array_groups).dirty