|
| 1 | +module GlobalDiffEqSciMLSensitivityExt |
| 2 | + |
| 3 | +import GlobalDiffEq, LinearAlgebra, QuadGK, SciMLBase, SciMLSensitivity |
| 4 | + |
| 5 | +function GlobalDiffEq._default_quadrature_sensealg() |
| 6 | + return SciMLSensitivity.QuadratureAdjoint(autojacvec = true) |
| 7 | +end |
| 8 | + |
| 9 | +GlobalDiffEq._is_quadrature_adjoint(::SciMLSensitivity.QuadratureAdjoint) = true |
| 10 | + |
| 11 | +function GlobalDiffEq._adjoint_solution( |
| 12 | + sol, sensealg::SciMLSensitivity.QuadratureAdjoint, adjoint_alg, direction; |
| 13 | + abstol, reltol |
| 14 | + ) |
| 15 | + terminal_gradient! = let direction = direction |
| 16 | + function (out, u, p, t, i) |
| 17 | + copyto!(out, direction) |
| 18 | + return nothing |
| 19 | + end |
| 20 | + end |
| 21 | + terminal_time = sol.prob.tspan[2] |
| 22 | + adjoint_prob = SciMLSensitivity.ODEAdjointProblem( |
| 23 | + sol, sensealg, adjoint_alg, [terminal_time], terminal_gradient! |
| 24 | + ) |
| 25 | + adjoint_sol = SciMLBase.solve( |
| 26 | + adjoint_prob, adjoint_alg; |
| 27 | + abstol, reltol, dense = true, save_everystep = true |
| 28 | + ) |
| 29 | + SciMLBase.successful_retcode(adjoint_sol) || |
| 30 | + throw(ErrorException("the adjoint solve failed with retcode $(adjoint_sol.retcode)")) |
| 31 | + return adjoint_sol |
| 32 | +end |
| 33 | + |
| 34 | +function GlobalDiffEq._defect_projection(sol, adjoint_sol; abstol, reltol) |
| 35 | + prob = sol.prob |
| 36 | + isinplace = SciMLBase.isinplace(prob) |
| 37 | + |
| 38 | + function integrand(t) |
| 39 | + u = sol(t, continuity = :right) |
| 40 | + du = sol(t, Val{1}, continuity = :right) |
| 41 | + rhs = if isinplace |
| 42 | + value = similar(u) |
| 43 | + prob.f(value, u, prob.p, t) |
| 44 | + value |
| 45 | + else |
| 46 | + prob.f(u, prob.p, t) |
| 47 | + end |
| 48 | + lambda = adjoint_sol(t, continuity = :right) |
| 49 | + return LinearAlgebra.dot(lambda, du - rhs) |
| 50 | + end |
| 51 | + |
| 52 | + projection = zero(eltype(prob.u0)) |
| 53 | + for i in 1:(length(sol.t) - 1) |
| 54 | + left = sol.t[i] |
| 55 | + right = sol.t[i + 1] |
| 56 | + left == right && continue |
| 57 | + interval_projection, _ = QuadGK.quadgk( |
| 58 | + integrand, left, right; atol = abstol, rtol = reltol |
| 59 | + ) |
| 60 | + projection += interval_projection |
| 61 | + end |
| 62 | + return projection |
| 63 | +end |
| 64 | + |
| 65 | +end |
0 commit comments