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