Skip to content

Commit 1159155

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 53ebab8 commit 1159155

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

@@ -457,6 +455,7 @@ function OptimizationBase.instantiate_function(
457455
end
458456

459457
function fill_lag_hessian!(θ, σ, μ, p)
458+
lag = x -> lagrangian(x, f.f, f.cons, p, μ, σ)
460459
for i in eachindex(θ)
461460
Enzyme.make_zero!(lag_bθ)
462461
Enzyme.make_zero!(lag_vdbθ[i])
@@ -465,13 +464,8 @@ function OptimizationBase.instantiate_function(
465464
lag_grad,
466465
Const(rmode),
467466
Enzyme.Duplicated(θ, lag_vdθ[i]),
468-
Enzyme.DuplicatedNoNeed(lag_bθ, lag_vdbθ[i]),
469-
Const(lagrangian),
470-
Const(f.f),
471-
Const(f.cons),
472-
Const(p),
473-
Const(σ),
474-
Const(μ)
467+
Enzyme.Duplicated(lag_bθ, lag_vdbθ[i]),
468+
Const(lag)
475469
)
476470
end
477471
return nothing
@@ -494,7 +488,9 @@ function OptimizationBase.instantiate_function(
494488
fill_lag_hessian!(θ, σ, μ, p)
495489

496490
for i in eachindex(θ)
497-
H[i, :] .= lag_vdbθ[i]
491+
vec_lagv = lag_vdbθ[i]
492+
H[i, 1:i] .= @view(vec_lagv[1:i])
493+
H[1:i, i] .= @view(vec_lagv[1:i])
498494
end
499495
return
500496
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)