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)
|
|||
|
|
}
|
|||
|
|
}
|
|||
|
|
}
|