189 lines
5.0 KiB
Go
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
|
||
|
|
}
|