287 lines
8.0 KiB
Go
287 lines
8.0 KiB
Go
// 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)
|
|
}
|
|
}
|
|
}
|