feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,292 @@
|
||||
// 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)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user