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