From e795e6c75fa9d23ffabfa91765fe6985de2dabcc Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Beno=C3=AEt=20Legat?= Date: Thu, 17 Sep 2026 11:50:28 +0200 Subject: [PATCH] Fix hessian with min and max operators --- src/forward_over_reverse.jl | 5 +---- src/operators.jl | 8 +------ test/test_ReverseAD.jl | 45 +++++++++++++++++++++++++++++++++++++ 3 files changed, 47 insertions(+), 11 deletions(-) diff --git a/src/forward_over_reverse.jl b/src/forward_over_reverse.jl index b01bceb..68b1be5 100644 --- a/src/forward_over_reverse.jl +++ b/src/forward_over_reverse.jl @@ -324,15 +324,12 @@ function _forward_eval_ϵ( d.user_output_buffer, n_children, ) - has_hessian = eval_multivariate_hessian( + eval_multivariate_hessian( d.data.operators, d.data.operators.multivariate_operators[node.index], H, f_input, ) - # This might be `false` if we extend this code to all - # multivariate functions. - @assert has_hessian for col in 1:n_children dual = zero(P) for row in 1:n_children diff --git a/src/operators.jl b/src/operators.jl index b962e5d..4826fff 100644 --- a/src/operators.jl +++ b/src/operators.jl @@ -276,7 +276,7 @@ function eval_multivariate_hessian( H, x::AbstractVector{T}, ) where {T} - if op in (:+, :-, :ifelse) + if op in (:+, :-, :ifelse, :min, :max) return false end if op == :* @@ -340,12 +340,6 @@ function eval_multivariate_hessian( H[1, 1] = -2 * x[2] * x[1] / base H[2, 1] = (x[1]^2 - x[2]^2) / base H[2, 2] = 2 * x[2] * x[1] / base - elseif op == :min - _, i = findmin(x) - H[i, i] = one(T) - elseif op == :max - _, i = findmax(x) - H[i, i] = one(T) else id = registry.multivariate_operator_to_id[op] offset = id - registry.multivariate_user_operator_start diff --git a/test/test_ReverseAD.jl b/test/test_ReverseAD.jl index eed7c10..4709d87 100644 --- a/test/test_ReverseAD.jl +++ b/test/test_ReverseAD.jl @@ -1370,6 +1370,51 @@ function test_hessian_reinterpret_unsafe() return end +function test_hessian_min() + x, y = MOI.VariableIndex.(1:2) + model = ArrayDiff.Model() + ArrayDiff.set_objective(model, :(min($x^2, $y^2))) + evaluator = ArrayDiff.Evaluator(model, ArrayDiff.Mode(), [x, y]) + MOI.initialize(evaluator, [:Grad, :Hess]) + @test MOI.hessian_lagrangian_structure(evaluator) == [(1, 1), (2, 2)] + H = zeros(2) + MOI.eval_hessian_lagrangian(evaluator, H, [1.1, 2.3], 1.5, Float64[]) + @test isapprox(H, [3.0, 0.0]) + MOI.eval_hessian_lagrangian(evaluator, H, [2.3, 1.5], 1.2, Float64[]) + @test isapprox(H, [0.0, 2.4]) + return +end + +function test_hessian_max() + x, y = MOI.VariableIndex.(1:2) + model = ArrayDiff.Model() + ArrayDiff.set_objective(model, :(max($x^2, $y^2))) + evaluator = ArrayDiff.Evaluator(model, ArrayDiff.Mode(), [x, y]) + MOI.initialize(evaluator, [:Grad, :Hess]) + @test MOI.hessian_lagrangian_structure(evaluator) == [(1, 1), (2, 2)] + H = zeros(2) + MOI.eval_hessian_lagrangian(evaluator, H, [1.1, 2.3], 1.5, Float64[]) + @test isapprox(H, [0.0, 3.0]) + MOI.eval_hessian_lagrangian(evaluator, H, [2.3, 1.5], 1.2, Float64[]) + @test isapprox(H, [2.4, 0.0]) + return +end + +function test_hessian_ifelse() + x, y = MOI.VariableIndex.(1:2) + model = ArrayDiff.Model() + ArrayDiff.set_objective(model, :(ifelse($x < $y, $x^2, $y^2))) + evaluator = ArrayDiff.Evaluator(model, ArrayDiff.Mode(), [x, y]) + MOI.initialize(evaluator, [:Grad, :Hess]) + @test MOI.hessian_lagrangian_structure(evaluator) == [(1, 1), (2, 2)] + H = zeros(2) + MOI.eval_hessian_lagrangian(evaluator, H, [1.1, 2.3], 1.5, Float64[]) + @test isapprox(H, [3.0, 0.0]) + MOI.eval_hessian_lagrangian(evaluator, H, [2.3, 1.5], 1.2, Float64[]) + @test isapprox(H, [0.0, 2.4]) + return +end + end # module TestReverseAD.runtests()