The Zygote rule for getindex(::ODESolution, sym) handles a scalar symbol but not an array symbol, and there is no rule at all for DAESolution. Indexing a solution by an array unknown inside a Zygote loss therefore fails, while the same loss written with the scalar elements works.
using ModelingToolkit, OrdinaryDiffEq, SciMLSensitivity, Zygote
using ModelingToolkit: t_nounits as t, D_nounits as D
using SymbolicIndexingInterface: setp_oop
@variables x(t)[1:2]
@parameters p[1:2] = [1.0, 2.0]
@mtkcompile sys = System([D(x[1]) ~ -p[1] * x[1], D(x[2]) ~ -p[2] * x[2]], t)
prob = ODEProblem(sys, [x => [1.0, 1.0]], (0.0, 1.0))
set_p = setp_oop(prob, p)
sa = InterpolatingAdjoint(autojacvec = ReverseDiffVJP(true))
sol(ps) = solve(remake(prob; p = set_p(prob, ps)), Tsit5(); saveat = 0.1, sensealg = sa)
Zygote.gradient(ps -> sum(abs2, reduce(hcat, [sol(ps)[x[1]], sol(ps)[x[2]]])), [1.0, 2.0]) # works
Zygote.gradient(ps -> sum(abs2, reduce(hcat, sol(ps)[x])), [1.0, 2.0])
# TypeError: non-boolean (Num) used in boolean context, in findall called from the getindex pullback
With a DAEProblem of the same system (0 ~ -p[i] * x[i] - D(x[i]), DFBDF()), the array symbol fails with Cannot convert an object of type Bool to an object of type Vector{Float64} and the scalar symbols with invalid index: (x(t))[2]: the solution is a DAESolution, for which no symbolic getindex rule exists, so Zygote falls back to array indexing with a symbolic index.
In ext/SciMLBaseZygoteExt.jl (and the ChainRules extension) the scalar rule looks up variable_index(VA, sym) and scatters Δ[j] into position i; for an array symbol variable_index returns a vector of indices and the returned values are vectors per time point, so the pullback needs to scatter each Δ[j][k] into idx[k]. The same rule with AbstractODESolution in the signature would cover DAESolution.
Versions: SciMLBase 3.53.1, SciMLSensitivity 7.119.3, Zygote 0.7.13, ModelingToolkit 11.42.0, Julia 1.12.5.
The Zygote rule for
getindex(::ODESolution, sym)handles a scalar symbol but not an array symbol, and there is no rule at all forDAESolution. Indexing a solution by an array unknown inside a Zygote loss therefore fails, while the same loss written with the scalar elements works.With a
DAEProblemof the same system (0 ~ -p[i] * x[i] - D(x[i]),DFBDF()), the array symbol fails withCannot convert an object of type Bool to an object of type Vector{Float64}and the scalar symbols withinvalid index: (x(t))[2]: the solution is aDAESolution, for which no symbolicgetindexrule exists, so Zygote falls back to array indexing with a symbolic index.In
ext/SciMLBaseZygoteExt.jl(and the ChainRules extension) the scalar rule looks upvariable_index(VA, sym)and scattersΔ[j]into positioni; for an array symbolvariable_indexreturns a vector of indices and the returned values are vectors per time point, so the pullback needs to scatter eachΔ[j][k]intoidx[k]. The same rule withAbstractODESolutionin the signature would coverDAESolution.Versions: SciMLBase 3.53.1, SciMLSensitivity 7.119.3, Zygote 0.7.13, ModelingToolkit 11.42.0, Julia 1.12.5.