Files
tensor/grad/bench_perf_test.go
T

293 lines
6.8 KiB
Go
Raw 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 (
"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)
}
}
}