168 lines
4.9 KiB
Go
168 lines
4.9 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 TestTensorConcatForward(t *testing.T) {
|
|
a, _ := core.FromFloats([]float64{1, 2, 3, 4}, 2, 2)
|
|
b, _ := core.FromFloats([]float64{5, 6, 7, 8, 9, 10}, 2, 3)
|
|
|
|
out, err := FromArray(a, false).Concat(FromArray(b, false), 1)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if got := out.Data().Shape(); got[0] != 2 || got[1] != 5 {
|
|
t.Fatalf("concat shape: %v", got)
|
|
}
|
|
want := []float64{1, 2, 5, 6, 7, 3, 4, 8, 9, 10}
|
|
for i := range want {
|
|
if g := out.Data().FloatAt(i); g != want[i] {
|
|
t.Fatalf("concat[%d] = %v, want %v", i, g, want[i])
|
|
}
|
|
}
|
|
|
|
// Concatenation along the leading axis stacks the blocks.
|
|
c, _ := core.FromFloats([]float64{1, 2}, 1, 2)
|
|
d, _ := core.FromFloats([]float64{3, 4}, 1, 2)
|
|
vert, err := FromArray(c, false).Concat(FromArray(d, false), 0)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if vert.Data().Shape()[0] != 2 {
|
|
t.Fatalf("vertical shape: %v", vert.Data().Shape())
|
|
}
|
|
|
|
// Mismatched ranks and out-of-range axes error.
|
|
misrank, _ := core.Reshape(c, 2)
|
|
if _, err := FromArray(a, false).Concat(FromArray(misrank, false), 0); err == nil {
|
|
t.Fatal("rank mismatch accepted")
|
|
}
|
|
if _, err := FromArray(a, false).Concat(FromArray(b, false), 2); err == nil {
|
|
t.Fatal("out-of-range dimension accepted")
|
|
}
|
|
}
|
|
|
|
// TestTensorConcatGradients checks both backward spans against central
|
|
// differences with a weighted loss so every slot gets a distinct weight.
|
|
func TestTensorConcatGradients(t *testing.T) {
|
|
cases := []struct {
|
|
dim int
|
|
aVal, bVal []float64
|
|
aShape, bShape []int
|
|
waVal, wbVal []float64
|
|
}{
|
|
{
|
|
dim: 1,
|
|
aVal: []float64{0.5, -1, 2, 0.25}, aShape: []int{2, 2},
|
|
bVal: []float64{1.5, -0.5, 1, 2, -2, 0.75}, bShape: []int{2, 3},
|
|
waVal: []float64{0.1, -0.4, 0.9, 0.6},
|
|
wbVal: []float64{0.2, 0.3, -0.7, 0.8, 0.05, -0.6},
|
|
},
|
|
{
|
|
dim: 0,
|
|
aVal: []float64{0.3, 1, -0.25, 2}, aShape: []int{2, 2},
|
|
bVal: []float64{-1.5, 0.4, 0.9, 1, -2, 0.7}, bShape: []int{3, 2},
|
|
waVal: []float64{0.55, -0.35, 0.85, 0.15},
|
|
wbVal: []float64{0.45, -0.65, 0.95, 0.05, -0.5, 0.75},
|
|
},
|
|
}
|
|
for _, tc := range cases {
|
|
a, _ := core.FromFloats(tc.aVal, tc.aShape...)
|
|
b, _ := core.FromFloats(tc.bVal, tc.bShape...)
|
|
wa, _ := core.FromFloats(tc.waVal, tc.aShape...)
|
|
wb, _ := core.FromFloats(tc.wbVal, tc.bShape...)
|
|
|
|
at := FromArray(a, true)
|
|
bt := FromArray(b, true)
|
|
joint, err := at.Concat(bt, tc.dim)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
wc, _ := core.Concat(wa, wb, tc.dim)
|
|
scaled, err := joint.Mul(FromArray(wc, 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)
|
|
}
|
|
|
|
fa := func(v *core.Array) float64 { return weightedConcatSum(v, b, wa, wb, tc.dim) }
|
|
fb := func(v *core.Array) float64 { return weightedConcatSum(a, v, wa, wb, tc.dim) }
|
|
checkSpan(t, at.Grad(), numericGrad(fa, a))
|
|
checkSpan(t, bt.Grad(), numericGrad(fb, b))
|
|
}
|
|
}
|
|
|
|
// TestTensorConcatGradientDtype keeps each side's gradient in its own
|
|
// element type: float32 inputs never come back as float64 leaves.
|
|
func TestTensorConcatGradientDtype(t *testing.T) {
|
|
gen := core.NewGenerator(5)
|
|
af, _ := core.Float32s(gen, 6)
|
|
afArr, _ := core.Reshape(af, 2, 3)
|
|
bf, _ := core.Float32s(gen, 6)
|
|
bfArr, _ := core.Reshape(bf, 2, 3)
|
|
|
|
at := FromArray(afArr, true)
|
|
bt := FromArray(bfArr, true)
|
|
joint, err := at.Concat(bt, 1)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
loss, err := joint.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("gradient dtypes: %v and %v", at.Grad().Dtype(), bt.Grad().Dtype())
|
|
}
|
|
for i := range at.Grad().Len() {
|
|
if at.Grad().FloatAt(i) != 1 {
|
|
t.Errorf("float32 gradient slot %d: %v, want 1", i, at.Grad().FloatAt(i))
|
|
}
|
|
}
|
|
}
|
|
|
|
// weightedConcatSum evaluates Σ w∘Concat(x, y, dim) with fixed weights,
|
|
// the scalar objective whose gradients the backward is checked against.
|
|
func weightedConcatSum(x, y *core.Array, wx, wy *core.Array, dim int) float64 {
|
|
joint, err := FromArray(x, false).Concat(FromArray(y, false), dim)
|
|
if err != nil {
|
|
return math.NaN()
|
|
}
|
|
wc, _ := core.Concat(wx, wy, dim)
|
|
total := 0.0
|
|
for i := range joint.Data().Len() {
|
|
total += wc.FloatAt(i) * joint.Data().FloatAt(i)
|
|
}
|
|
return total
|
|
}
|
|
|
|
// checkSpan reports every slot where the analytic gradient drifts from
|
|
// the central-difference reference.
|
|
func checkSpan(t *testing.T, got *core.Array, ref []float64) {
|
|
t.Helper()
|
|
if got.Len() != len(ref) {
|
|
t.Fatalf("gradient length %d, reference %d", got.Len(), len(ref))
|
|
}
|
|
for i := range ref {
|
|
if math.Abs(got.FloatAt(i)-ref[i]) > 1e-5 {
|
|
t.Errorf("gradient[%d] = %v, want ≈%v", i, got.FloatAt(i), ref[i])
|
|
}
|
|
}
|
|
}
|