// Copyright (c) 2026 Petr BalvĂ­n (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 }