diff --git a/lib/ModelingToolkitTearing/src/reassemble.jl b/lib/ModelingToolkitTearing/src/reassemble.jl index 9dba0cf..2ca279e 100644 --- a/lib/ModelingToolkitTearing/src/reassemble.jl +++ b/lib/ModelingToolkitTearing/src/reassemble.jl @@ -712,9 +712,65 @@ function safe_ldiv(A, b) shape = SU.promote_shape(safe_ldiv, SU.shape(A), SU.shape(b)) ) end - return CommonSolve.solve(LinearProblem(A, b)).u + return numeric_ldiv!(A, b) end +""" + $TYPEDSIGNATURES + +Solve the numeric linear system emitted for an inlined linear SCC. + +`A` and `b` are the scratch buffers built by the `ArrayMaker` in +[`get_linear_scc_linsol`](@ref), rewritten entry by entry on every call. The `LinearCache` +is kept in task local storage and reused, so a steady state call neither builds a +`LinearProblem` nor allocates a factorization. + +The cache is task local because the emitted expression is shared by every problem built +from it, while `A` and `b` are not. + +One cache is kept per array type and size, so a solve site can share with another of the +same shape. The solution is therefore copied back into `b`, which belongs to this solve +site alone, rather than returning the cache's own buffer. +""" +function numeric_ldiv!(A::AbstractMatrix, b::AbstractVector) + # the lookup is necessarily type unstable, so the solve goes behind a function barrier + return solve_into!(get_inline_linsolve_cache(A, b), A, b) +end + +function solve_into!(cache, A::AbstractMatrix, b::AbstractVector) + # Fill the cache's own buffers and assign them back, rather than writing through + # `cache.A` in place. The assignment is what runs LinearSolve's invalidation. Under + # `ForwardDiff` that matters twice over: a `DualLinearCache` keeps the primal matrix in + # a separate array and the partials behind their own validity flag, and an in place + # write reaches neither. The next solve then silently uses the previous call's matrix + # and partials, so both the values and the derivatives come back wrong. + Awork = cache.A + copyto!(Awork, A) + cache.A = Awork + bwork = cache.b + copyto!(bwork, b) + cache.b = bwork + copyto!(b, CommonSolve.solve!(cache).u) + return b +end + +function get_inline_linsolve_cache(A::AbstractMatrix, b::AbstractVector) + tls = task_local_storage() + key = (INLINE_LINSOLVE_CACHE, typeof(A), size(A)) + cache = get(tls, key, nothing) + if cache === nothing + # `A` and `b` are views into the diffcache buffers of whichever problem happens + # to call first, and the cache keeps what it is handed, so it needs copies it owns. + # `copy` rather than `Matrix`/`Vector` so the cache keeps the array type it was + # given, which matters for anything living off the CPU. + cache = CommonSolve.init(LinearProblem(copy(A), copy(b))) + tls[key] = cache + end + return cache +end + +numeric_ldiv!(A, b) = CommonSolve.solve(LinearProblem(A, b)).u + function SU.promote_symtype(::typeof(safe_ldiv), TA::SU.TypeT, TB::SU.TypeT) return Vector{Real} end @@ -730,6 +786,11 @@ end const INLINE_LINEAR_SCC_OP = safe_ldiv +""" +Tag distinguishing this package's entries in task local storage. See [`numeric_ldiv!`](@ref). +""" +const INLINE_LINSOLVE_CACHE = :ModelingToolkitTearing_inline_linsolve_cache + """ $TYPEDSIGNATURES diff --git a/lib/ModelingToolkitTearing/test/runtests.jl b/lib/ModelingToolkitTearing/test/runtests.jl index 92c46cb..3de3e54 100644 --- a/lib/ModelingToolkitTearing/test/runtests.jl +++ b/lib/ModelingToolkitTearing/test/runtests.jl @@ -11,6 +11,7 @@ import SymbolicUtils as SU using SymbolicUtils: unwrap using Setfield using ForwardDiff +using LinearAlgebra import ModelingToolkitBase as MTKBase @testset "`InferredDiscrete` validation" begin @@ -142,6 +143,131 @@ end end end +@testset "Inline linear SCC solve does not allocate per call" begin + @variables x(t) = 1.0 y(t) = 1.0 z(t) = 1.0 w(t) = 1.0 q(t) = 1.0 + reassemble_alg = MTKTearing.DefaultReassembleAlgorithm(; inline_linear_sccs = true) + # `A` depends on `x`, so the SCC cannot be solved symbolically and is emitted as a + # runtime solve. `D(q)` reads the SCC, so the solve is live in the RHS. + eqs = [ + D(x) ~ 2t + 1, + (2 + x) * y + x * z + w ~ 4, + (4 + x) * y + 3z + 2w ~ 7, + 2x * y + (3 + x) * z + w ~ 10, + D(q) ~ 2w + 3z + y, + ] + @mtkcompile sys = System(eqs, t) reassemble_alg = reassemble_alg + + blk = only(MTKTearing.inline_linear_systems(sys)) + @test SU.operation(unwrap(blk.expression)) === MTKTearing.INLINE_LINEAR_SCC_OP + + prob = ODEProblem(sys, [], (0.0, 1.0)) + du = similar(prob.u0) + f! = prob.f.f + f!(du, prob.u0, prob.p, 0.0) + # the solve used to build a fresh `LinearProblem` and `LinearCache` per call, which + # cost ~1.4 kB and 23 allocations regardless of the size of the system + @test @allocated(f!(du, prob.u0, prob.p, 0.0)) < 512 + + # the same system without inlining is the reference for the values + @mtkcompile refsys = System(eqs, t) + refprob = ODEProblem(refsys, [], (0.0, 1.0)) + refdu = similar(refprob.u0) + refprob.f.f(refdu, refprob.u0, refprob.p, 0.0) + @test du[1] ≈ refdu[1] +end + +@testset "`numeric_ldiv!` solves in place and stays rank tolerant" begin + A = [2.0 1.0; 1.0 3.0] + b = [1.0, 2.0] + expected = A \ b + u = MTKTearing.numeric_ldiv!(copy(A), b) + @test u ≈ expected + # the solution is written into `b`, which is scratch in the emitted code + @test b ≈ expected + @test u === b + + # reusing the cache must not carry anything over between solves + A2 = [4.0 1.0; 2.0 5.0] + b2 = [3.0, 1.0] + @test MTKTearing.numeric_ldiv!(copy(A2), b2) ≈ A2 \ [3.0, 1.0] + + # a singular system still goes through LinearSolve's default algorithm, which returns + # the minimum-norm least squares solution + S = [1.0 2.0; 2.0 4.0] + @test MTKTearing.numeric_ldiv!(copy(S), [1.0, 2.0]) ≈ pinv(S) * [1.0, 2.0] + + # and the cache recovers for the next well conditioned solve of the same shape + @test MTKTearing.numeric_ldiv!(copy(A), [1.0, 2.0]) ≈ expected +end + +@testset "Inline linear SCC derivatives stay correct across repeated solves" begin + @variables x(t) = 1.0 y(t) = 1.0 z(t) = 1.0 w(t) = 1.0 q(t) = 1.0 + reassemble_alg = MTKTearing.DefaultReassembleAlgorithm(; inline_linear_sccs = true) + @mtkcompile sys = System( + [ + D(x) ~ 2t + 1, + (2 + x) * y + x * z + w ~ 4, + (4 + x) * y + 3z + 2w ~ 7, + 2x * y + (3 + x) * z + w ~ 10, + D(q) ~ 2w + 3z + y, + ], + t + ) reassemble_alg = reassemble_alg + + jac(prob, tv) = ForwardDiff.jacobian( + (du, u) -> (prob.f.f(du, u, prob.p, tv); nothing), + similar(prob.u0), copy(prob.u0) + ) + + prob1 = ODEProblem(sys, [x => 1.0, q => 1.0], (0.0, 1.0)) + prob2 = ODEProblem(sys, [x => 2.5, q => 3.0], (0.0, 1.0)) + + # the solve cache is task local, so a fresh task builds it from scratch and its first + # solve is the reference value + J1 = fetch(Threads.@spawn jac(prob1, 0.3)) + J2 = fetch(Threads.@spawn jac(prob2, 0.3)) + # the two problems must actually differ or this cannot detect a stale cache + @test !isapprox(J1, J2) + + # evaluating `prob1` on a cache already primed by `prob2` must still give `prob1`'s + # jacobian. Writing `A` into the cache without letting LinearSolve invalidate it leaves + # both the primal matrix and the partials of the previous solve in place, so the answer + # is the previous system's, and nothing warns. + @test fetch(Threads.@spawn (jac(prob2, 0.3); jac(prob1, 0.3))) ≈ J1 + @test fetch(Threads.@spawn all(i -> jac(prob1, 0.3) ≈ J1, 1:5)) + @test fetch(Threads.@spawn all(i -> jac(prob1, 0.3) ≈ J1 && jac(prob2, 0.3) ≈ J2, 1:5)) +end + +@testset "Two inline linear SCCs of the same shape don't share a result buffer" begin + @variables x(t) = 1.0 q(t) = 1.0 + @variables y(t) = 1.0 z(t) = 1.0 w(t) = 1.0 + @variables a(t) = 1.0 c(t) = 1.0 d(t) = 1.0 + reassemble_alg = MTKTearing.DefaultReassembleAlgorithm(; inline_linear_sccs = true) + # two independent SCCs that both reduce to 2x2, so they share a cache + eqs = [ + D(x) ~ 2t + 1, + (2 + x) * y + x * z + w ~ 4, + (4 + x) * y + 3z + 2w ~ 7, + 2x * y + (3 + x) * z + w ~ 10, + (2 + x) * a + x * c + d ~ 5, + (4 + x) * a + 3c + 2d ~ 8, + 2x * a + (3 + x) * c + d ~ 11, + D(q) ~ 2w + 3z + y + 2d + 3c + a, + ] + @mtkcompile sys = System(eqs, t) reassemble_alg = reassemble_alg + @test length(MTKTearing.inline_linear_systems(sys)) == 2 + + prob = ODEProblem(sys, [], (0.0, 1.0)) + du = similar(prob.u0) + prob.f.f(du, prob.u0, prob.p, 0.0) + + @mtkcompile refsys = System(eqs, t) + refprob = ODEProblem(refsys, [], (0.0, 1.0)) + refdu = similar(refprob.u0) + refprob.f.f(refdu, refprob.u0, refprob.p, 0.0) + @test du[1] ≈ refdu[1] +end + @testset "`inline_linear_systems` diagnostic" begin @variables x(t) y(t) z(t) w(t) q(t) ralg = MTKTearing.DefaultReassembleAlgorithm(; inline_linear_sccs = true)