diff --git a/lib/OptimizationBase/ext/OptimizationEnzymeExt.jl b/lib/OptimizationBase/ext/OptimizationEnzymeExt.jl index 1bb7d68ef..a1fa857e0 100644 --- a/lib/OptimizationBase/ext/OptimizationEnzymeExt.jl +++ b/lib/OptimizationBase/ext/OptimizationEnzymeExt.jl @@ -167,12 +167,13 @@ function OptimizationBase.instantiate_function( function hess(res, θ, p = p) Enzyme.make_zero!(bθ) Enzyme.make_zero!.(vdbθ) + θ_arr = θ isa Array ? θ : Array(θ) Enzyme.autodiff( fmode, inner_grad, Const(rmode), - Enzyme.BatchDuplicated(θ, vdθ), + Enzyme.BatchDuplicated(θ_arr, vdθ), Enzyme.BatchDuplicatedNoNeed(bθ, vdbθ), Const(f.f), Const(p) @@ -193,17 +194,24 @@ function OptimizationBase.instantiate_function( function fgh!(G, H, θ, p = p) vdθ = Tuple((Array(r) for r in eachrow(I(length(θ)) * one(eltype(θ))))) vdbθ = Tuple(zeros(eltype(θ), length(θ)) for i in eachindex(θ)) + θ_arr = θ isa Array ? θ : Array(θ) + G_arr = G isa Array ? G : Array(G) + Enzyme.make_zero!(G_arr) Enzyme.autodiff( fmode, inner_grad, Const(rmode), - Enzyme.BatchDuplicated(θ, vdθ), - Enzyme.BatchDuplicatedNoNeed(G, vdbθ), + Enzyme.BatchDuplicated(θ_arr, vdθ), + Enzyme.BatchDuplicatedNoNeed(G_arr, vdbθ), Const(f.f), Const(p) ) + if !(G isa Array) + copyto!(G, G_arr) + end + for i in eachindex(θ) H[i, :] .= vdbθ[i] end diff --git a/lib/OptimizationBase/test/AD/adtests.jl b/lib/OptimizationBase/test/AD/adtests.jl index 83729e1b9..b519523be 100644 --- a/lib/OptimizationBase/test/AD/adtests.jl +++ b/lib/OptimizationBase/test/AD/adtests.jl @@ -114,6 +114,20 @@ optprob.cons_h(H3, x0) optprob.lag_h(H4, x0, σ, μ) @test H4 ≈ σ * H2 + μ[1] * H3[1] rtol = 1.0e-6 + # Test non-Vector AbstractVector (e.g. SubArray) for AutoEnzyme hess and fgh! + x_view = @view zeros(4)[1:2] + optprob_view = OptimizationBase.instantiate_function( + OptimizationFunction(rosenbrock, OptimizationBase.AutoEnzyme()), x_view, + OptimizationBase.AutoEnzyme(), nothing, 0, h = true, fgh = true + ) + H_view = Array{Float64}(undef, 2, 2) + G_view = Array{Float64}(undef, 2) + optprob_view.hess(H_view, x_view) + @test H1 == H_view + optprob_view.fgh(G_view, H_view, x_view) + @test G1 == G_view + @test H1 == H_view + G2 = Array{Float64}(undef, 2) H2 = Array{Float64}(undef, 2, 2)