Skip to content
Merged
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
63 changes: 62 additions & 1 deletion lib/ModelingToolkitTearing/src/reassemble.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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

Expand Down
126 changes: 126 additions & 0 deletions lib/ModelingToolkitTearing/test/runtests.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down
Loading