Files

39 lines
1.1 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 "testing"
// TestHessianVectorProductZeroDirectionShape pins the repair: a zero
// direction returned a flat length-n vector while
// every other answer carries the point's own shape.
func TestHessianVectorProductZeroDirectionShape(t *testing.T) {
x0 := []float64{0.5, -1.25, 2, -2}
xt, err := FromFloat64s(x0, false, 2, 2)
if err != nil {
t.Fatalf("FromFloat64s: %v", err)
}
f := func(z *Tensor) (*Tensor, error) { return z.Sum() }
v, err := FromFloat64s(make([]float64, 4), false, 2, 2)
if err != nil {
t.Fatalf("FromFloat64s: %v", err)
}
hv, err := HessianVectorProduct(f, xt, v, HessianOptions{})
if err != nil {
t.Fatalf("HessianVectorProduct: %v", err)
}
want := []int{2, 2}
got := hv.Shape()
for i := range want {
if got[i] != want[i] {
t.Fatalf("H·0 shape = %v, want %v", got, want)
}
}
for i := range 4 {
if hv.FloatAt(i) != 0 {
t.Fatalf("H·0 = %v, want the zero matrix", hv.FloatAt(i))
}
}
}