// Copyright (c) 2026 Petr Balvín (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) } } }