Files

41 lines
1.2 KiB
Go
Raw Permalink Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"strings"
"testing"
)
// Regression tests: an objective whose graph
// never reaches x left xt.Grad() nil, and the second-order helpers
// dereferenced it.
// TestHessianDisconnectedObjective pins the error: an objective that
// ignores its argument has no gradient to differentiate.
func TestHessianDisconnectedObjective(t *testing.T) {
x, err := FromFloat64s([]float64{1, 2}, true, 2)
if err != nil {
t.Fatal(err)
}
c, err := FromFloat64s([]float64{3}, true, 1)
if err != nil {
t.Fatal(err)
}
objective := func(*Tensor) (*Tensor, error) { return c.Mul(c) }
if _, err := Hessian(objective, x, HessianOptions{}); err == nil {
t.Fatal("expected an error when the objective does not depend on x")
} else if !strings.Contains(err.Error(), "does not depend") {
t.Fatalf("error = %v, want a disconnected-graph refusal", err)
}
v, err := FromFloat64s([]float64{1, 0}, true, 2)
if err != nil {
t.Fatal(err)
}
if _, err := HessianVectorProduct(objective, x, v, HessianOptions{}); err == nil {
t.Fatal("expected an error when the objective does not depend on x")
}
}