Files

159 lines
4.6 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 (
"math"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
// Regression pins for the MatMul adjoints: the 1-D × 2-D branch against
// central differences, real and complex, and the Newton-CG loop's last
// iteration.
// TestMatMulVectorByMatrixGradient pins the 1-D × 2-D branch of
// the MatMul adjoint (da = g·Bᵀ, db = outer(a, g)), real and complex,
// against central differences. The branch had no gradient test at all.
func TestMatMulVectorByMatrixGradient(t *testing.T) {
avec := []float64{1.5, -0.5, 2, 0.25}
bvec := []float64{
0.5, -1, 2,
1.5, 0.25, -0.75,
-2, 1, 0.5,
1, -0.5, 1.25,
}
// A weighted linear loss, so every output slot carries its own
// coefficient and a wrong routing cannot cancel against another.
w := []float64{0.7, -1.3, 2.1}
lossOf := func(a, b *core.Array) float64 {
out, err := core.MatMul2D(a, b)
if err != nil {
t.Fatalf("MatMul2D: %v", err)
}
s := 0.0
for i := range out.Len() {
s += w[i] * out.FloatAt(i)
}
return s
}
cloneWith := func(a *core.Array, i int, v float64) *core.Array {
vals := make([]float64, a.Len())
for k := range a.Len() {
vals[k] = a.FloatAt(k)
}
vals[i] = v
out, err := core.FromFloats(vals, a.Shape()...)
if err != nil {
t.Fatalf("FromFloats: %v", err)
}
return out
}
a, _ := FromFloat64s(avec, true, 4)
b, _ := FromFloat64s(bvec, true, 4, 3)
if err := backwardWeightedMatMul(t, a, b, w); err != nil {
t.Fatalf("Backward: %v", err)
}
baseA, _ := core.FromFloats(avec, 4)
baseB, _ := core.FromFloats(bvec, 4, 3)
for i := range 4 {
eps := 1e-6
want := (lossOf(cloneWith(baseA, i, avec[i]+eps), baseB) -
lossOf(cloneWith(baseA, i, avec[i]-eps), baseB)) / (2 * eps)
if math.Abs(a.Grad().FloatAt(i)-want) > 1e-5*(1+math.Abs(want)) {
t.Fatalf("da[%d] = %g, want %g", i, a.Grad().FloatAt(i), want)
}
}
for i := range 12 {
eps := 1e-6
want := (lossOf(baseA, cloneWith(baseB, i, bvec[i]+eps)) -
lossOf(baseA, cloneWith(baseB, i, bvec[i]-eps))) / (2 * eps)
if math.Abs(b.Grad().FloatAt(i)-want) > 1e-5*(1+math.Abs(want)) {
t.Fatalf("db[%d] = %g, want %g", i, b.Grad().FloatAt(i), want)
}
}
// The complex 1-D × 2-D branch against numericComplexGrad, under
// the same weighted fold backwardComplex builds: L = Σ Re(w̄·y) +
// Σ|y|²/n with the helper's own deterministic weights.
cb := FromArray(mustComplexes([]complex128{
0.5 + 0.5i, -1,
1.5, 0.25 - 0.75i,
-2 + 1i, 1,
}, 3, 2), true)
op := func(x *Tensor) (*Tensor, error) { return x.MatMul(cb) }
vals := []complex128{1 + 0.5i, -0.25 - 1i, 0.75 + 0.25i}
xt := backwardComplex(t, op, vals, 3)
closs := complexLossOf(t, op, spectralWeights(2, 11))
checkAgainstNumeric(t, xt.Grad(), numericComplexGrad(closs, xt.Data()), 1e-6)
}
// backwardWeightedMatMul builds L = w·(a·B) over fresh leaves and runs
// one backward pass.
func backwardWeightedMatMul(t *testing.T, a, b *Tensor, w []float64) error {
t.Helper()
out, err := a.MatMul(b)
if err != nil {
return err
}
wv, err := FromFloat64s(w, false, len(w))
if err != nil {
return err
}
loss, err := out.Mul(wv)
if err != nil {
return err
}
sum, err := loss.Sum()
if err != nil {
return err
}
sum.Backward()
return nil
}
// TestNewtonCGConvergesOnTheLastIteration pins that a tolerance
// met exactly on the final permitted iteration is a success, not a
// budget error whose message prints a gradient already under the
// tolerance.
func TestNewtonCGConvergesOnTheLastIteration(t *testing.T) {
f := func(x *Tensor) (*Tensor, error) {
d, err := x.Sub(mustTensorF64(1.5))
if err != nil {
return nil, err
}
return d.Mul(d)
}
x0, _ := core.FromFloats([]float64{0}, 1)
// One truncated-CG step lands within the Hessian-product rounding
// of the minimiser (about 1e-11 here); a tolerance of 1e-9 is met
// by exactly that step, on the final permitted iteration.
out, _, err := MinimiseNewtonCG(f, x0, NewtonCGOptions{MaxIterations: 1, Tolerance: 1e-9})
if err != nil {
t.Fatalf("MinimiseNewtonCG on a quadratic with one exact step: %v", err)
}
if math.Abs(out.FloatAt(0)-1.5) > 1e-9 {
t.Fatalf("minimum = %g, want 1.5", out.FloatAt(0))
}
}
// mustTensorF64 wraps one float as a no-grad tensor.
func mustTensorF64(v float64) *Tensor {
t, err := FromFloat64s([]float64{v}, false, 1)
if err != nil {
panic(err)
}
return t
}
// mustComplexes builds a complex array or fails the test.
func mustComplexes(vals []complex128, shape ...int) *core.Array {
a, err := core.FromComplexes(vals, shape...)
if err != nil {
panic(err)
}
return a
}