293 lines
6.8 KiB
Go
293 lines
6.8 KiB
Go
// 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)
|
||
}
|
||
}
|
||
}
|