feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,38 @@
|
||||
// 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))
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user