39 lines
1.1 KiB
Go
39 lines
1.1 KiB
Go
// 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))
|
|
}
|
|
}
|
|
}
|