Files
tensor/grad/batched_test.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

189 lines
5.0 KiB
Go

// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package grad
import (
"math"
"testing"
"sourcedock.dev/petrbalvin/tensor/internal/core"
)
func TestMatMulBatchedForward(t *testing.T) {
a, _ := core.FromFloats([]float64{
1, 0,
0, 1,
3, 4,
5, 6,
}, 2, 2, 2)
b, _ := core.FromFloats([]float64{
1, 1,
1, 0,
2, 0,
0, 2,
}, 2, 2, 2)
out, err := FromArray(a, false).MatMulBatched(FromArray(b, false))
if err != nil {
t.Fatal(err)
}
want := []float64{1, 1, 1, 0, 6, 8, 10, 12}
for i := range want {
if g := out.Data().FloatAt(i); g != want[i] {
t.Fatalf("slot %d = %v, want %v", i, g, want[i])
}
}
// Rank and batch mismatches error loudly.
flat, _ := core.Reshape(a, 8)
if _, err := FromArray(flat, false).MatMulBatched(FromArray(b, false)); err == nil {
t.Fatal("rank-2 operand accepted")
}
c, _ := core.FromFloats(make([]float64, 4), 1, 2, 2)
if _, err := FromArray(a, false).MatMulBatched(FromArray(c, false)); err == nil {
t.Fatal("batch-size mismatch accepted")
}
d, _ := core.FromFloats(make([]float64, 12), 2, 3, 2)
if _, err := FromArray(a, false).MatMulBatched(FromArray(d, false)); err == nil {
t.Fatal("inner-dimension mismatch accepted")
}
}
// TestMatMulBatchedGradients finite-difference checks both operands on
// a weighted sum objective so every batch slot earns its own weight.
func TestMatMulBatchedGradients(t *testing.T) {
aVal := []float64{0.5, -1, 2, 0.25, 1.5, -0.5, 0.75, 1.25, -0.25, 0.5, -1.5, 2}
bVal := []float64{1, -0.25, 0.75, 2, -1, 0.125, 0.5, -2}
weight := sweepPattern(12) // covers (2, 3, 2) outputs
aArr, _ := core.FromFloats(aVal, 2, 3, 2)
bArr, _ := core.FromFloats(bVal, 2, 2, 2)
mArr, _ := core.FromFloats(weight, 2, 3, 2)
at := FromArray(aArr, true)
bt := FromArray(bArr, true)
out, err := at.MatMulBatched(bt)
if err != nil {
t.Fatal(err)
}
scaled, err := out.Mul(FromArray(mArr, false))
if err != nil {
t.Fatal(err)
}
loss, err := scaled.Sum()
if err != nil {
t.Fatal(err)
}
if err := loss.Backward(); err != nil {
t.Fatal(err)
}
objective := func(av, bv []float64) float64 {
x, _ := core.FromFloats(av, 2, 3, 2)
y, _ := core.FromFloats(bv, 2, 2, 2)
o, rerr := FromArray(x, false).MatMulBatched(FromArray(y, false))
if rerr != nil {
return math.NaN()
}
total := 0.0
for i := range weight {
total += mArr.FloatAt(i) * o.Data().FloatAt(i)
}
return total
}
checkSpan(t, at.Grad(), numericGrad(func(v *core.Array) float64 {
return objective(flatten(v), bVal)
}, aArr))
checkSpan(t, bt.Grad(), numericGrad(func(v *core.Array) float64 {
return objective(aVal, flatten(v))
}, bArr))
}
func flatten(v *core.Array) []float64 {
out := make([]float64, v.Len())
for i := range v.Len() {
out[i] = v.FloatAt(i)
}
return out
}
// TestMatMulBatchedGradientsFloat32 runs the weighted-sum
// finite-difference check on float32 operands: the forward rounds to
// float32 while the backward widens the accessors and answers float64
// gradients, and both agree with the float64 reference.
func TestMatMulBatchedGradientsFloat32(t *testing.T) {
aVal := []float32{0.5, -1, 2, 0.25, 1.5, -0.5, 0.75, 1.25, -0.25, 0.5, -1.5, 2}
bVal := []float32{1, -0.25, 0.75, 2, -1, 0.125, 0.5, -2}
weight := sweepPattern(12)
aArr, err := core.FromFloat32s(aVal, 2, 3, 2)
if err != nil {
t.Fatal(err)
}
bArr, err := core.FromFloat32s(bVal, 2, 2, 2)
if err != nil {
t.Fatal(err)
}
mArr, err := core.FromFloats(weight, 2, 3, 2)
if err != nil {
t.Fatal(err)
}
at := FromArray(aArr, true)
bt := FromArray(bArr, true)
out, err := at.MatMulBatched(bt)
if err != nil {
t.Fatal(err)
}
scaled, err := out.Mul(FromArray(mArr, false))
if err != nil {
t.Fatal(err)
}
loss, err := scaled.Sum()
if err != nil {
t.Fatal(err)
}
if err := loss.Backward(); err != nil {
t.Fatal(err)
}
if at.Grad().Dtype() != core.Float32 || bt.Grad().Dtype() != core.Float32 {
t.Fatalf("float32 leaves carry %s and %s gradients, want float32",
at.Grad().Dtype(), bt.Grad().Dtype())
}
// The reference differentiates the same batched product evaluated
// in float64 over the identical operand values.
objective := func(av, bv []float64) float64 {
x, _ := core.FromFloats(av, 2, 3, 2)
y, _ := core.FromFloats(bv, 2, 2, 2)
o, rerr := FromArray(x, false).MatMulBatched(FromArray(y, false))
if rerr != nil {
return math.NaN()
}
total := 0.0
for i := range weight {
total += mArr.FloatAt(i) * o.Data().FloatAt(i)
}
return total
}
aRef, _ := core.FromFloats(widen32(aVal), 2, 3, 2)
bRef, _ := core.FromFloats(widen32(bVal), 2, 2, 2)
checkSpan(t, at.Grad(), numericGrad(func(v *core.Array) float64 {
return objective(flatten(v), widen32(bVal))
}, aRef))
checkSpan(t, bt.Grad(), numericGrad(func(v *core.Array) float64 {
return objective(widen32(aVal), flatten(v))
}, bRef))
}
// widen32 widens a float32 slice exactly, the view the backward's own
// accessors read.
func widen32(v []float32) []float64 {
out := make([]float64, len(v))
for i, x := range v {
out[i] = float64(x)
}
return out
}