Files

287 lines
8.0 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"
)
// The cross-rank gradient sweep (the machine that catches regressions
// hiding in untested shapes, like the LayerNorm affine reduction):
// every listed differentiable op runs through a finite-difference
// check on several input ranks and both float element types.
type sweepCase struct {
name string
ranks [][]int // the shapes the op must answer for
run func(x *Tensor) (*Tensor, error)
}
// sweepValue is deterministic, sign-varying and comfortably away from
// kinks and saturation boundaries.
func sweepValue(i int) float64 {
return math.Sin(float64(i%17)*0.7)*2 + 0.25
}
func sweepPattern(n int) []float64 {
out := make([]float64, n)
for i := range out {
out[i] = 0.5*float64(i%4) - 0.75
}
return out
}
func sweepMatrixPattern(width int) *core.Array {
vals := sweepPattern(width * width)
arr, _ := core.FromFloats(vals, width, width)
return arr
}
// sweepPositive keeps logs and divisions inside their real domains no
// matter how the signed sweep values land, shaped like the input.
func sweepPositive(shape []int) *core.Array {
n := numEl(shape)
out := make([]float64, n)
for i := range out {
out[i] = 3 + float64(i%3)
}
arr, _ := core.FromFloats(out, shape...)
return arr
}
// sweepSecond derives a second operand from an independent pattern:
// paired cases need two leaves but stay deterministic.
func sweepSecond(shape []int) (*Tensor, error) {
total := 1
for _, d := range shape {
total *= d
}
vals := make([]float64, total)
for i := range vals {
vals[i] = math.Cos(float64(i%11))*1.5 - 0.5
}
a, err := core.FromFloats(vals, shape...)
if err != nil {
return nil, err
}
return FromArray(a, true), nil
}
func numEl(shape []int) int {
n := 1
for _, d := range shape {
n *= d
}
return n
}
// checkOp runs one case per shape and dtype: forward on a gradient
// leaf, weighted-sum loss with a fixed mask so every slot gets its own
// coefficient, then analytic-vs-central-difference compare.
func checkOp(t *testing.T, tc sweepCase, dt core.Dtype) {
t.Helper()
for _, dims := range tc.ranks {
n := numEl(dims)
build := func(reqGrad bool) *Tensor {
vals := make([]float64, n)
for i := range vals {
vals[i] = sweepValue(i + len(dims))
}
a, err := core.FromFloats(vals, dims...)
if err != nil {
t.Fatal(err)
}
if dt == core.Float32 {
a32, cerr := core.Astype(a, core.Float32)
if cerr != nil {
t.Fatal(cerr)
}
a = a32
}
return FromArray(a, reqGrad)
}
x := build(true)
out, err := tc.run(x)
if err != nil {
t.Fatalf("%s %v %v: forward: %v", tc.name, dims, dt, err)
}
// The mask matches the OUTPUT shape: reducing ops return fewer
// slots than their input carries.
outN := out.Data().Len()
maskVals := sweepPattern(outN)
mArr, _ := core.FromFloats(maskVals, out.Data().Shape()...)
scaled, err := out.Mul(FromArray(mArr, false))
if err != nil {
t.Fatalf("%s %v %v: loss mul: %v", tc.name, dims, dt, err)
}
loss, err := scaled.Sum()
if err != nil {
t.Fatalf("%s %v %v: loss sum: %v", tc.name, dims, dt, err)
}
if err := loss.Backward(); err != nil {
t.Fatalf("%s %v %v: backward: %v", tc.name, dims, dt, err)
}
ref := numericGrad(func(v *core.Array) float64 {
o, rerr := tc.run(FromArray(v, false))
if rerr != nil {
return math.NaN()
}
total := 0.0
for i := range maskVals {
total += mArr.FloatAt(i) * o.Data().FloatAt(i)
}
return total
}, x.Data())
got := x.Grad()
if got.Len() != len(ref) {
t.Fatalf("%s %v %v: gradient length %d, reference %d",
tc.name, dims, dt, got.Len(), len(ref))
}
scale := 1.0
for _, r := range ref {
if s := math.Abs(r); s > scale {
scale = s
}
}
tol := 1e-4
if dt == core.Float32 {
tol = 8e-2
}
for i := range ref {
if math.Abs(got.FloatAt(i)-ref[i]) > tol*scale {
t.Errorf("%s %v %v: grad[%d] = %v, want ≈%v",
tc.name, dims, dt, i, got.FloatAt(i), ref[i])
}
}
}
}
func TestGradientSweepAcrossRanksAndDtypes(t *testing.T) {
shapes234 := [][]int{{4}, {2, 3}, {2, 2, 2}}
simple := []sweepCase{
{name: "Neg", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Neg() }},
{name: "Exp", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Exp() }},
{name: "Sigmoid", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Sigmoid() }},
{name: "Tanh", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Tanh() }},
{name: "Abs", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Abs() }},
{name: "Pow3", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Pow(3) }},
{name: "Scale", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Scale(-1.75) }},
{name: "ClipInterior", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) { return x.Clip(-2.75, 2.75) }},
{name: "LogShifted", ranks: shapes234, run: func(x *Tensor) (*Tensor, error) {
up, err := x.Add(FromArray(sweepPositive(x.Data().Shape()), false))
if err != nil {
return nil, err
}
return up.Log()
}},
{name: "TransposeAxesReverse", ranks: [][]int{{2, 3}, {2, 3, 2}}, run: func(x *Tensor) (*Tensor, error) {
d := x.Data().NDim()
perm := make([]int, d)
for i := range perm {
perm[i] = d - 1 - i
}
return x.TransposeAxes(perm...)
}},
{name: "ReshapeFlatten", ranks: [][]int{{2, 3}, {2, 2, 2}}, run: func(x *Tensor) (*Tensor, error) {
return x.Reshape(x.Data().Len())
}},
{name: "SumAxisZero", ranks: [][]int{{2, 3}, {2, 3, 2}}, run: func(x *Tensor) (*Tensor, error) {
return x.SumAxis(0)
}},
{name: "MeanAxisLast", ranks: [][]int{{2, 3}, {2, 3, 2}}, run: func(x *Tensor) (*Tensor, error) {
return x.MeanAxis(x.Data().NDim() - 1)
}},
// The L2 norm's backward runs a dedicated float32 sweep beside
// the float64 one; the two dtype legs below drive both.
{name: "L2NormAxisLast", ranks: [][]int{{4}, {2, 3}, {2, 3, 2}}, run: func(x *Tensor) (*Tensor, error) {
return x.L2NormAxis(x.Data().NDim() - 1)
}},
}
elementPairs := []struct {
name string
op func(a, b *Tensor) (*Tensor, error)
}{
{"Add", func(a, b *Tensor) (*Tensor, error) { return a.Add(b) }},
{"Sub", func(a, b *Tensor) (*Tensor, error) { return a.Sub(b) }},
{"Mul", func(a, b *Tensor) (*Tensor, error) { return a.Mul(b) }},
}
for _, ep := range elementPairs {
simple = append(simple, sweepCase{
name: ep.name,
ranks: shapes234,
run: func(x *Tensor) (*Tensor, error) {
other, err := sweepSecond(x.Data().Shape())
if err != nil {
return nil, err
}
return ep.op(x, other)
},
})
}
// Division keeps both operands positive via the shared shift.
simple = append(simple, sweepCase{
name: "DivShifted",
ranks: shapes234,
run: func(x *Tensor) (*Tensor, error) {
other, err := sweepSecond(x.Data().Shape())
if err != nil {
return nil, err
}
lift, err := other.Add(FromArray(sweepPositive(x.Data().Shape()), false))
if err != nil {
return nil, err
}
return x.Div(lift)
},
})
// Column concatenation against half of a second leaf.
simple = append(simple, sweepCase{
name: "ConcatColumns",
ranks: [][]int{{2, 4}},
run: func(x *Tensor) (*Tensor, error) {
extra, err := sweepSecond(x.Data().Shape())
if err != nil {
return nil, err
}
halves, err := extra.Slice(1, 0, x.Data().Shape()[1]/2)
if err != nil {
return nil, err
}
return x.Concat(halves, 1)
},
})
// A matmul product collapsed by an axis sum, the inference
// backbone's gradient path.
simple = append(simple, sweepCase{
name: "MatMulSumRows",
ranks: [][]int{{3, 4}},
run: func(x *Tensor) (*Tensor, error) {
cols := x.Data().Shape()[1]
product, err := x.MatMul(FromArray(sweepMatrixPattern(cols), false))
if err != nil {
return nil, err
}
return product.SumAxis(0)
},
})
for _, tc := range simple {
for _, dt := range []core.Dtype{core.Float, core.Float32} {
checkOp(t, tc, dt)
}
}
}