Files
tensor/grad/concat_test.go
T

168 lines
4.9 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// 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])
}
}
}