41 lines
1.2 KiB
Go
41 lines
1.2 KiB
Go
// 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")
|
||
|
|
}
|
||
|
|
}
|