// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package grad import ( "testing" core "sourcedock.dev/petrbalvin/tensor/internal/core" ) // Backward benchmarks guard the tape overhead around the kernels: the // two-layer graph is the smallest shape where per-node costs and the // matmul backward both show. func benchVals(b *testing.B, seed, n int) []float64 { b.Helper() v := make([]float64, n) for i := range v { v[i] = float64(i%13)*float64(seed%3)*0.25 + float64(i%5) - 2 } return v } func benchTensor(b *testing.B, seed int, shape ...int) *Tensor { b.Helper() n := 1 for _, d := range shape { n *= d } a, err := core.FromFloats(benchVals(b, seed, n), shape...) if err != nil { b.Fatal(err) } return FromArray(a, true) } // BenchmarkBackwardTwoLayer runs forward and backward over // (32×64)·(64×32), then tanh, then ·(32×10), then sum. func BenchmarkBackwardTwoLayer(b *testing.B) { x := benchTensor(b, 1, 32, 64) w1 := benchTensor(b, 2, 64, 32) w2 := benchTensor(b, 3, 32, 10) b.ReportAllocs() for b.Loop() { h, err := x.MatMul(w1) if err != nil { b.Fatal(err) } t, err := h.Tanh() if err != nil { b.Fatal(err) } y, err := t.MatMul(w2) if err != nil { b.Fatal(err) } s, err := y.Sum() if err != nil { b.Fatal(err) } if err := s.Backward(); err != nil { b.Fatal(err) } } } // BenchmarkForwardOnly isolates the graph construction from the // backward sweep. func BenchmarkForwardOnly(b *testing.B) { x := benchTensor(b, 1, 32, 64) w1 := benchTensor(b, 2, 64, 32) w2 := benchTensor(b, 3, 32, 10) b.ReportAllocs() for b.Loop() { h, err := x.MatMul(w1) if err != nil { b.Fatal(err) } t, err := h.Tanh() if err != nil { b.Fatal(err) } if _, err := t.MatMul(w2); err != nil { b.Fatal(err) } } } // BenchmarkGradMatMulBackward isolates one MatMul node's backward // sweep on (128×128) operands. func BenchmarkGradMatMulBackward(b *testing.B) { a := benchTensor(b, 4, 128, 128) c := benchTensor(b, 5, 128, 128) y, err := a.MatMul(c) if err != nil { b.Fatal(err) } s, err := y.Sum() if err != nil { b.Fatal(err) } b.ResetTimer() b.ReportAllocs() for b.Loop() { a.ZeroGrad() c.ZeroGrad() if err := s.Backward(); err != nil { b.Fatal(err) } } }