Skip to content

Commit 91ee292

Browse files
Stabilize Enzyme Lagrangian Hessians on Julia 1.12
Co-Authored-By: Chris Rackauckas <accounts@chrisrackauckas.com> Co-Authored-By: Claude <noreply@anthropic.com> Claude-Session: https://chatgpt.com/codex/tasks/01a03a17-ad6f-7131-82fc-d0fd57ea6512
1 parent 247c7c6 commit 91ee292

2 files changed

Lines changed: 12 additions & 14 deletions

File tree

lib/OptimizationBase/ext/OptimizationEnzymeExt.jl

Lines changed: 9 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -80,11 +80,9 @@ function lagrangian(x, _f::Function, cons::Function, p, λ, σ = one(eltype(x)))
8080
return σ * _f(x, p) + dot(λ, res)
8181
end
8282

83-
function lag_grad(mode, x, dx, lagrangian::Function, _f::Function, cons::Function, p, σ, λ)
84-
Enzyme.autodiff_deferred(
85-
mode, Const(lagrangian), Active, Enzyme.Duplicated(x, dx),
86-
Const(_f), Const(cons), Const(p), Const(λ), Const(σ)
87-
)
83+
function lag_grad(mode, x, dx, f)
84+
Enzyme.make_zero!(dx)
85+
Enzyme.autodiff(mode, Const(f), Active, Enzyme.Duplicated(x, dx))
8886
return nothing
8987
end
9088

@@ -499,6 +497,7 @@ function OptimizationBase.instantiate_function(
499497
end
500498

501499
function fill_lag_hessian!(θ, σ, μ, p)
500+
lag = x -> lagrangian(x, f.f, f.cons, p, μ, σ)
502501
for i in eachindex(θ)
503502
Enzyme.make_zero!(lag_bθ)
504503
Enzyme.make_zero!(lag_vdbθ[i])
@@ -507,13 +506,8 @@ function OptimizationBase.instantiate_function(
507506
lag_grad,
508507
Const(rmode),
509508
Enzyme.Duplicated(θ, lag_vdθ[i]),
510-
Enzyme.DuplicatedNoNeed(lag_bθ, lag_vdbθ[i]),
511-
Const(lagrangian),
512-
Const(f.f),
513-
Const(f.cons),
514-
Const(p),
515-
Const(σ),
516-
Const(μ)
509+
Enzyme.Duplicated(lag_bθ, lag_vdbθ[i]),
510+
Const(lag)
517511
)
518512
end
519513
return nothing
@@ -536,7 +530,9 @@ function OptimizationBase.instantiate_function(
536530
fill_lag_hessian!(θ, σ, μ, p)
537531

538532
for i in eachindex(θ)
539-
H[i, :] .= lag_vdbθ[i]
533+
vec_lagv = lag_vdbθ[i]
534+
H[i, 1:i] .= @view(vec_lagv[1:i])
535+
H[1:i, i] .= @view(vec_lagv[1:i])
540536
end
541537
return
542538
end

lib/OptimizationBase/test/AD/enzyme_lagrangian_hessian.jl

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -46,5 +46,7 @@ function check_lagrangian_hessian(N)
4646
end
4747

4848
@testset "Enzyme Lagrangian Hessian" begin
49-
foreach(check_lagrangian_hessian, (1:10..., 20, 40, 60))
49+
@testset "N = $N" for N in (1:10..., 20, 40, 60)
50+
check_lagrangian_hessian(N)
51+
end
5052
end

0 commit comments

Comments
 (0)