// 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 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]) } } }