// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: MIT package core import ( "math" "strings" "testing" ) func TestFlatten(t *testing.T) { a := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3) got, err := Flatten(a, 0, -1) if err != nil { t.Fatal(err) } want, _ := FromFloats([]float64{1, 2, 3, 4, 5, 6}, 6) if got.String() != want.String() { t.Errorf("flatten all: got %s, want %s", got, want) } // Flatten only middle dimension. got, err = Flatten(a, 1, 1) if err != nil { t.Fatal(err) } if got.Shape()[0] != 2 || got.Shape()[1] != 3 { t.Errorf("flatten 1: shape = %v", got.Shape()) } // Out-of-range. if _, err := Flatten(a, 0, 5); err == nil { t.Error("flatten: expected error for out-of-range endDim") } } func TestSqueeze(t *testing.T) { a, _ := FromFloats([]float64{1, 2, 3, 4}, 1, 2, 1, 2) got, err := Squeeze(a, 0) if err != nil { t.Fatal(err) } if len(got.Shape()) != 3 || got.Shape()[0] != 2 { t.Errorf("squeeze dim 0: shape = %v", got.Shape()) } // Squeeze all (-1). got, err = Squeeze(a, -1) if err != nil { t.Fatal(err) } if len(got.Shape()) != 2 || got.Shape()[0] != 2 { t.Errorf("squeeze all: shape = %v", got.Shape()) } // Cannot squeeze non-1 dimension. if _, err := Squeeze(a, 1); err == nil { t.Error("squeeze: expected error for non-1 dim") } } func TestUnsqueeze(t *testing.T) { a := mustFromFloats(t, []float64{1, 2, 3}, 3) got, err := Unsqueeze(a, 0) if err != nil { t.Fatal(err) } if len(got.Shape()) != 2 || got.Shape()[0] != 1 || got.Shape()[1] != 3 { t.Errorf("unsqueeze 0: shape = %v", got.Shape()) } // Negative dim. got, err = Unsqueeze(a, -1) if err != nil { t.Fatal(err) } if got.Shape()[0] != 3 || got.Shape()[1] != 1 { t.Errorf("unsqueeze -1: shape = %v", got.Shape()) } } func TestTransposeAxes(t *testing.T) { a, _ := FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3) got, err := TransposeAxes(a, 1, 0) if err != nil { t.Fatal(err) } if got.Shape()[0] != 3 || got.Shape()[1] != 2 { t.Errorf("transpose: shape = %v", got.Shape()) } // Check element: a[1,0] = 4 should land at out[0,1]. v, _ := FloatAt(got, 0, 1) if v != 4 { t.Errorf("transpose element: got %v, want 4", v) } // Invalid permutation. if _, err := TransposeAxes(a, 0, 0); err == nil { t.Error("transpose: expected error for duplicate dims") } } func TestCopy(t *testing.T) { a := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2) got := Copy(a) if got.Shape()[0] != 2 || got.Shape()[1] != 2 { t.Errorf("copy: shape = %v", got.Shape()) } } func TestPadConstant(t *testing.T) { a := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2) // 2-D: pad = (left, right, top, bottom). got, err := Pad(a, []int{1, 1, 1, 1}, "constant", 0) if err != nil { t.Fatal(err) } if got.Shape()[0] != 4 || got.Shape()[1] != 4 { t.Errorf("pad constant: shape = %v", got.Shape()) } // Corners should be zero. c, _ := FloatAt(got, 0, 0) if c != 0 { t.Errorf("pad corner: got %v, want 0", c) } // Centre should be original a[0,0] = 1. c, _ = FloatAt(got, 1, 1) if c != 1 { t.Errorf("pad centre: got %v, want 1", c) } // Wrong pad length. if _, err := Pad(a, []int{1, 1, 1}, "constant", 0); err == nil { t.Error("pad: expected error for odd-length pad") } // Unknown mode. if _, err := Pad(a, []int{1, 1, 1, 1}, "weird", 0); err == nil { t.Error("pad: expected error for unknown mode") } } func TestPadReflect(t *testing.T) { a := mustFromFloats(t, []float64{1, 2, 3}, 1, 3) got, err := Pad(a, []int{2, 0, 0, 0}, "reflect", 0) if err != nil { t.Fatal(err) } // Reflect without repeating edge: at index 0 we mirror index 2 -> 3, // at index 1 we mirror index 1 -> 2. v, _ := FloatAt(got, 0, 0) if v != 3 { t.Errorf("pad reflect [0]: got %v, want 3", v) } v, _ = FloatAt(got, 0, 1) if v != 2 { t.Errorf("pad reflect [1]: got %v, want 2", v) } // Original index 0 = 1 should land at the position equal to the pre-pad. v, _ = FloatAt(got, 0, 2) if v != 1 { t.Errorf("pad reflect [2]: got %v, want 1", v) } } func TestPadReplicate(t *testing.T) { a := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2) got, err := Pad(a, []int{2, 1, 0, 0}, "replicate", 0) if err != nil { t.Fatal(err) } // Shape (2, 5). At (0, 0) replicate original (0, 0) = 1. v, _ := FloatAt(got, 0, 0) if v != 1 { t.Errorf("pad replicate [0,0]: got %v, want 1", v) } v, _ = FloatAt(got, 0, 1) if v != 1 { t.Errorf("pad replicate [0,1]: got %v, want 1", v) } // Post-pad replicates last column (index 1) for the last 1 column. v, _ = FloatAt(got, 1, 4) if v != 4 { t.Errorf("pad replicate [1,4]: got %v, want 4", v) } } func TestPadCircular(t *testing.T) { a := mustFromFloats(t, []float64{1, 2, 3, 4}, 2, 2) got, err := Pad(a, []int{2, 1, 0, 0}, "circular", 0) if err != nil { t.Fatal(err) } // Shape (2, 5). Pad left=2, right=1 on dim 1. // Row 0 = [1, 2]. Wrapped with pre=2 -> [..., 1, 2, 1, 2, 1] then trim to 5: // index 0 = wrap(0-2=-2) = 0 -> 1 // index 1 = wrap(0-1=-1) = 1 -> 2 // index 2 = 1 // index 3 = 2 // index 4 = wrap(0+2=2) mod 2 = 0 -> 1 v, _ := FloatAt(got, 0, 0) if v != 1 { t.Errorf("pad circular [0,0]: got %v, want 1", v) } v, _ = FloatAt(got, 0, 1) if v != 2 { t.Errorf("pad circular [0,1]: got %v, want 2", v) } v, _ = FloatAt(got, 0, 4) if v != 1 { t.Errorf("pad circular [0,4]: got %v, want 1", v) } } func TestGather(t *testing.T) { a, _ := FromFloats([]float64{10, 20, 30, 40, 50, 60}, 2, 3) idx := mustFromInts(t, []int64{2, 0}, 2, 1) got, err := Gather(a, 1, idx) if err != nil { t.Fatal(err) } // Expect [a[0,2], a[1,0]] = [30, 40]. v0, _ := FloatAt(got, 0, 0) v1, _ := FloatAt(got, 1, 0) if v0 != 30 || v1 != 40 { t.Errorf("gather: got [%v, %v], want [30, 40]", v0, v1) } // Out-of-range index. badIdx := mustFromInts(t, []int64{99, 0}, 2, 1) if _, err := Gather(a, 1, badIdx); err == nil { t.Error("gather: expected error for out-of-range index") } // Wrong dtype. badDtype := mustFromFloats(t, []float64{1, 1}, 2, 1) if _, err := Gather(a, 1, badDtype); err == nil { t.Error("gather: expected error for non-int index") } } func TestScatter(t *testing.T) { a, _ := FromFloats([]float64{10, 20, 30, 40, 50, 60}, 2, 3) idx := mustFromInts(t, []int64{2, 0}, 2, 1) src := mustFromFloats(t, []float64{300, 400}, 2, 1) got, err := Scatter(a, 1, idx, src) if err != nil { t.Fatal(err) } // Expect [10, 20, 300, 400, 50, 60]. for i, want := range []float64{10, 20, 300, 400, 50, 60} { v, _ := FloatAt(got, i/3, i%3) if v != want { t.Errorf("scatter [%d]: got %v, want %v", i, v, want) } } // Index/src shape mismatch. badIdx := mustFromInts(t, []int64{0, 0, 0}, 3, 1) if _, err := Scatter(a, 0, badIdx, idx); err == nil { t.Error("scatter: expected error for shape mismatch") } } func TestNonzero(t *testing.T) { a, _ := FromFloats([]float64{0, 1, 0, 2, 3, 0}, 2, 3) got, err := Nonzero(a) if err != nil { t.Fatal(err) } if len(got) != 2 { t.Fatalf("nonzero dims: got %d, want 2", len(got)) } // Flat layout [0, 1, 0, 2, 3, 0] -> row 0 [0,1,0], row 1 [2,3,0]. // Non-zero positions: flat 1 = (0,1), flat 3 = (1,0), flat 4 = (1,1). want := [][2]int{{0, 1}, {1, 0}, {1, 1}} for i, w := range want { if got[0][i] != w[0] || got[1][i] != w[1] { t.Errorf("nonzero [%d]: got (%d, %d), want %v", i, got[0][i], got[1][i], w) } } // Complex rejected. c, _ := FromComplexes([]complex128{1, 0}, 2) if _, err := Nonzero(c); err == nil { t.Error("nonzero: expected error for complex input") } } func TestTake(t *testing.T) { a := mustFromFloats(t, []float64{10, 20, 30, 40, 50}, 5) idx := mustFromInts(t, []int64{4, 0, 2}, 3) got, err := Take(a, idx) if err != nil { t.Fatal(err) } for i, want := range []float64{50, 10, 30} { v, _ := FloatAt(got, i) if v != want { t.Errorf("take [%d]: got %v, want %v", i, v, want) } } // Out-of-range. badIdx := mustFromInts(t, []int64{99}, 1) if _, err := Take(a, badIdx); err == nil { t.Error("take: expected error for out-of-range index") } // Non-1-D indices. badShape := mustFromInts(t, []int64{0, 0}, 2, 1) if _, err := Take(a, badShape); err == nil { t.Error("take: expected error for non-1-D indices") } } func TestCumSum(t *testing.T) { a := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3) got, err := CumSum(a, 1) if err != nil { t.Fatal(err) } want, _ := FromFloats([]float64{1, 3, 6, 4, 9, 15}, 2, 3) for i := range 6 { v, _ := FloatAt(got, i/3, i%3) w, _ := FloatAt(want, i/3, i%3) if v != w { t.Errorf("cumsum [%d]: got %v, want %v", i, v, w) } } // Out-of-range dim. if _, err := CumSum(a, 5); err == nil || !strings.Contains(err.Error(), "out of range") { t.Errorf("cumsum: expected out-of-range error, got %v", err) } } func TestCumProd(t *testing.T) { a := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3) got, err := CumProd(a, 0) if err != nil { t.Fatal(err) } want, _ := FromFloats([]float64{1, 2, 3, 4, 10, 18}, 2, 3) for i := range 6 { v, _ := FloatAt(got, i/3, i%3) w, _ := FloatAt(want, i/3, i%3) if v != w { t.Errorf("cumprod [%d]: got %v, want %v", i, v, w) } } } func TestProd(t *testing.T) { a := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6}, 2, 3) got, err := Prod(a, 1, false) if err != nil { t.Fatal(err) } // Prod across dim 1: row 0 = [1,2,3] -> 6, row 1 = [4,5,6] -> 120. for i, w := range []float64{6, 120} { v, _ := FloatAt(got, i) if v != w { t.Errorf("prod [%d]: got %v, want %v", i, v, w) } } // With keepDim. got, err = Prod(a, 1, true) if err != nil { t.Fatal(err) } if got.Shape()[0] != 2 || got.Shape()[1] != 1 { t.Errorf("prod keepDim: shape = %v", got.Shape()) } } func TestNorm(t *testing.T) { a := mustFromFloats(t, []float64{3, 4}, 2) got, err := Norm(a, 2, 0, false) if err != nil { t.Fatal(err) } if math.Abs(got.RawFloats()[0]-5) > 1e-9 { t.Errorf("L2 norm [3,4]: got %v, want 5", got.RawFloats()[0]) } // L1. got, _ = Norm(a, 1, 0, false) if got.RawFloats()[0] != 7 { t.Errorf("L1 norm: got %v, want 7", got.RawFloats()[0]) } // L-inf. got, _ = Norm(a, math.Inf(1), 0, false) if got.RawFloats()[0] != 4 { t.Errorf("L-inf norm: got %v, want 4", got.RawFloats()[0]) } // Negative p. if _, err := Norm(a, -1, 0, false); err == nil { t.Error("norm: expected error for negative p") } } func TestTrace(t *testing.T) { a := mustFromFloats(t, []float64{1, 2, 3, 4, 5, 6, 7, 8, 9}, 3, 3) got, err := Trace(a) if err != nil { t.Fatal(err) } if got != 15 { t.Errorf("trace: got %v, want 15", got) } // Non-square. bad, _ := FromFloats([]float64{1, 2, 3, 4}, 2, 2) // 2x2 is square; try non-square. bad2, _ := FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3) if _, err := Trace(bad2); err == nil { t.Error("trace: expected error for non-square") } // Just to silence "unused" for bad. _ = bad } func TestDiagonal(t *testing.T) { a, _ := FromFloats([]float64{1, 2, 3, 4, 5, 6, 7, 8, 9}, 3, 3) main, err := Diagonal(a, 0) if err != nil { t.Fatal(err) } for i, w := range []float64{1, 5, 9} { v, _ := FloatAt(main, i) if v != w { t.Errorf("diag main [%d]: got %v, want %v", i, v, w) } } // Super-diagonal offset 1. sup, err := Diagonal(a, 1) if err != nil { t.Fatal(err) } for i, w := range []float64{2, 6} { v, _ := FloatAt(sup, i) if v != w { t.Errorf("diag +1 [%d]: got %v, want %v", i, v, w) } } // Sub-diagonal offset -1. sub, err := Diagonal(a, -1) if err != nil { t.Fatal(err) } for i, w := range []float64{4, 8} { v, _ := FloatAt(sub, i) if v != w { t.Errorf("diag -1 [%d]: got %v, want %v", i, v, w) } } } func TestKron(t *testing.T) { a, _ := FromFloats([]float64{1, 2}, 1, 2) b, _ := FromFloats([]float64{0, 5, 6, 7}, 2, 2) got, err := Kron(a, b) if err != nil { t.Fatal(err) } if got.Shape()[0] != 2 || got.Shape()[1] != 4 { t.Errorf("kron: shape = %v", got.Shape()) } // Expect: a = [[1, 2]], b = [[0,5],[6,7]] // Result = [[0,5,0,10],[6,7,12,14]]. for i, w := range []float64{0, 5, 0, 10, 6, 7, 12, 14} { v, _ := FloatAt(got, i/4, i%4) if v != w { t.Errorf("kron [%d]: got %v, want %v", i, v, w) } } // Int input: exercises setFromValue with int dtype. aInt, _ := FromInts([]int64{1, 2}, 1, 2) gotInt, err := Kron(aInt, b) if err != nil { t.Fatal(err) } v0, _ := IntAt(gotInt, 0, 0) if v0 != 0 { t.Errorf("kron int [0,0]: got %v, want 0", v0) } } // TestPadReflectRejectsOversizedPads pins the reflect guard: a pad of // the full axis length has no source to mirror, and the single fold // used to emit a negative offset into the payload (panic). func TestPadReflectRejectsOversizedPads(t *testing.T) { a := mustFromFloats(t, []float64{1, 2, 3}, 3) if _, err := Pad(a, []int{6, 0}, "reflect", 0); err == nil { t.Error("Pad reflect pad=6 on length 3: expected error") } // A pad of n-1 is still mirrorable and must keep working. got, err := Pad(a, []int{2, 0}, "reflect", 0) if err != nil { t.Fatalf("Pad reflect pad=2: %v", err) } want := mustFromFloats(t, []float64{3, 2, 1, 2, 3}, 5) if !Equal(want, got) { t.Errorf("Pad reflect pad=2: %s", got) } } // TestScanFloat32CarryNarrows pins the float32 scan's carry rule: the // running value narrows to float32 at every step, exactly as the value // the walk stored fed the next one, so a term below the current // precision is absorbed before the next term lands on it. func TestScanFloat32CarryNarrows(t *testing.T) { // 1 + 1e-9 rounds back to 1, so the -1 lands on a clean zero. sums := mustFromFloat32s(t, []float32{1, 1e-9, -1}, 3) cs, err := CumSum(sums, 0) if err != nil { t.Fatalf("CumSum float32: %v", err) } for i, want := range []float32{1, 1, 0} { if got := cs.RawFloat32s()[i]; got != want { t.Errorf("CumSum float32 [%d] = %v, want %v", i, got, want) } } // A longer chain keeps the rule: the absorbed terms never // accumulate into the running sum. chain := mustFromFloat32s(t, []float32{1, 1e-9, 1e-9, 1e-9, -1}, 5) cs2, err := CumSum(chain, 0) if err != nil { t.Fatalf("CumSum float32 chain: %v", err) } for i, want := range []float32{1, 1, 1, 1, 0} { if got := cs2.RawFloat32s()[i]; got != want { t.Errorf("CumSum float32 chain [%d] = %v, want %v", i, got, want) } } // The product twin narrows the same way, so a product wider than // float32 is rounded before the next factor multiplies it. vals := []float32{1 + 1.0/4096, 1 + 1.0/4096, 1 + 1.0/4096, 1 + 1.0/4096} prod := mustFromFloat32s(t, vals, len(vals)) cp, err := CumProd(prod, 0) if err != nil { t.Fatalf("CumProd float32: %v", err) } var want []float32 acc := 1.0 for i, v := range vals { if i == 0 { acc = float64(v) } else { acc = float64(float32(acc)) * float64(v) } want = append(want, float32(acc)) } for i := range want { if got := cp.RawFloat32s()[i]; got != want[i] { t.Errorf("CumProd float32 [%d] = %v, want %v", i, got, want[i]) } } }