// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package grad import ( "testing" "sourcedock.dev/petrbalvin/tensor/internal/core" ) // The benchmarks below pin the costs the graph machinery adds around // the kernels: one deep chain of small tensors (per-node tape cost), one // wide element-wise graph (per-edge accumulation cost), the mid-size // sweeps whose per-element work decides their parallel policy, and the // L2-norm axis backward. Inputs are fixed literals, so a run is // deterministic. // benchLit builds a leaf of n elements from fixed literals, keeping the // arithmetic well inside the domain of every op used here. func benchLit(b *testing.B, seed, n int) *Tensor { b.Helper() v := make([]float64, n) for i := range v { v[i] = 0.25 + float64((i*7+seed)%13)*0.125 } a, err := core.FromFloats(v, n) if err != nil { b.Fatal(err) } return FromArray(a, true) } // mustReshape reshapes an array for a benchmark fixture. func mustReshape(b *testing.B, a *core.Array, shape ...int) *core.Array { b.Helper() out, err := core.Reshape(a, shape...) if err != nil { b.Fatal(err) } return out } // deepChain builds a chain of n element-wise nodes over x and w and // reduces it to a scalar, the shape a training loop's tape has. func deepChain(x, w *Tensor, n int) (*Tensor, error) { h := x for i := range n { var err error switch i % 4 { case 0: h, err = h.Add(w) case 1: h, err = h.Mul(w) case 2: h, err = h.Tanh() default: h, err = h.Scale(0.25) } if err != nil { return nil, err } } return h.Sum() } // wideFan multiplies x by n independent leaves and sums the products, // so the backward folds n contributions into x's gradient. func wideFan(x *Tensor, leaves []*Tensor) (*Tensor, error) { acc, err := x.Mul(leaves[0]) if err != nil { return nil, err } for _, l := range leaves[1:] { p, err := x.Mul(l) if err != nil { return nil, err } if acc, err = acc.Add(p); err != nil { return nil, err } } return acc.Sum() } // BenchmarkTapeDeepForward measures the forward pass alone: one node // per element-wise op over an 8-element tensor. func BenchmarkTapeDeepForward(b *testing.B) { x, w := benchLit(b, 1, 8), benchLit(b, 2, 8) b.ReportAllocs() for b.Loop() { s, err := deepChain(x, w, 128) if err != nil { b.Fatal(err) } if s.Data().Len() != 1 { b.Fatal("unexpected shape") } } } // BenchmarkTapeDeepBackward measures the same chain with the reverse // sweep, where every node reads its gradient and folds into the two // shared leaves. func BenchmarkTapeDeepBackward(b *testing.B) { x, w := benchLit(b, 1, 8), benchLit(b, 2, 8) b.ReportAllocs() for b.Loop() { s, err := deepChain(x, w, 128) if err != nil { b.Fatal(err) } x.ZeroGrad() w.ZeroGrad() if err := s.Backward(); err != nil { b.Fatal(err) } } } // BenchmarkTapeWideForward builds a 64-way fan over x, one node per // leaf, and reduces it. func BenchmarkTapeWideForward(b *testing.B) { x := benchLit(b, 3, 16) leaves := make([]*Tensor, 64) for i := range leaves { leaves[i] = benchLit(b, 10+i, 16) } b.ReportAllocs() for b.Loop() { s, err := wideFan(x, leaves) if err != nil { b.Fatal(err) } if s.Data().Len() != 1 { b.Fatal("unexpected shape") } } } // BenchmarkTapeWideBackward runs the same 64-way fan with the reverse // sweep: 64 edges fold into x's gradient through the reduction tree. func BenchmarkTapeWideBackward(b *testing.B) { x := benchLit(b, 3, 16) leaves := make([]*Tensor, 64) for i := range leaves { leaves[i] = benchLit(b, 10+i, 16) } b.ReportAllocs() for b.Loop() { s, err := wideFan(x, leaves) if err != nil { b.Fatal(err) } x.ZeroGrad() for _, l := range leaves { l.ZeroGrad() } if err := s.Backward(); err != nil { b.Fatal(err) } } } // midElems is the sweep size the transcendental benchmarks use: below // the element-wise floor of 1024 per worker on a 32-worker machine, so // a sweep of this size runs on the calling goroutine under that policy // and splits under a floor scaled to its per-element cost. const midElems = 20000 // BenchmarkPowForwardMid measures the integer-exponent power over a // mid-size sweep, one math.Pow per element. func BenchmarkPowForwardMid(b *testing.B) { x := benchLit(b, 5, midElems) b.ReportAllocs() for b.Loop() { y, err := x.Pow(3) if err != nil { b.Fatal(err) } if y.Data().Len() != midElems { b.Fatal("unexpected shape") } } } // BenchmarkPowBackwardMid measures the power backward over the same // size: one Pow and one multiply per element, plus the reduction's // fill. func BenchmarkPowBackwardMid(b *testing.B) { x := benchLit(b, 5, midElems) y, err := x.Pow(3) if err != nil { b.Fatal(err) } s, err := y.Sum() if err != nil { b.Fatal(err) } b.ReportAllocs() for b.Loop() { x.ZeroGrad() if err := s.Backward(); err != nil { b.Fatal(err) } } } // BenchmarkSqrtBackwardMid measures the square-root backward over a // mid-size sweep, one divide per element. func BenchmarkSqrtBackwardMid(b *testing.B) { x := benchLit(b, 6, midElems) y, err := x.Sqrt() if err != nil { b.Fatal(err) } s, err := y.Sum() if err != nil { b.Fatal(err) } b.ReportAllocs() for b.Loop() { x.ZeroGrad() if err := s.Backward(); err != nil { b.Fatal(err) } } } // BenchmarkMeanAxisBackwardMid measures the axis-mean backward, whose // first stage divides every element of the incoming gradient. func BenchmarkMeanAxisBackwardMid(b *testing.B) { x := benchLit(b, 7, midElems) xt := FromArray(mustReshape(b, x.Data(), 200, 100), true) y, err := xt.MeanAxis(0) if err != nil { b.Fatal(err) } s, err := y.Sum() if err != nil { b.Fatal(err) } b.ReportAllocs() for b.Loop() { xt.ZeroGrad() if err := s.Backward(); err != nil { b.Fatal(err) } } } // BenchmarkSumBackwardWide measures a reduction over a large operand: // the backward fills the operand's shape with one value, and the pass // commits a hundred-thousand-element gradient into the leaf. func BenchmarkSumBackwardWide(b *testing.B) { x := benchLit(b, 8, 100000) s, err := x.Sum() if err != nil { b.Fatal(err) } b.ReportAllocs() for b.Loop() { x.ZeroGrad() if err := s.Backward(); err != nil { b.Fatal(err) } } } // BenchmarkL2NormAxisBackward measures the norm backward over a // (2000×100) tensor reduced along the leading axis: 100 lines of 2000 // elements each, long enough for the line sweep to dominate the // allocation of the output. func BenchmarkL2NormAxisBackward(b *testing.B) { x := benchLit(b, 9, 200000) xt := FromArray(mustReshape(b, x.Data(), 2000, 100), true) y, err := xt.L2NormAxis(0) if err != nil { b.Fatal(err) } s, err := y.Sum() if err != nil { b.Fatal(err) } b.ReportAllocs() for b.Loop() { xt.ZeroGrad() if err := s.Backward(); err != nil { b.Fatal(err) } } }