// Copyright (c) 2026 Petr BalvĂ­n (https://petrbalvin.org) // SPDX-License-Identifier: MIT package core import ( "math" "testing" ) func TestGrid(t *testing.T) { a, _ := FromFloats([]float64{0, 1}, 2) b, _ := FromFloats([]float64{10, 20, 30}, 3) xg, yg, err := Grid(a, b) if err != nil { t.Fatal(err) } if xG := xg.Shape(); xG[0] != 3 || xG[1] != 2 { t.Fatalf("xGrid shape: %v", xG) } // X varies along columns. if v, _ := FloatAt(xg, 0, 1); v != 1 { t.Errorf("X grid [0,1]: %v", v) } // Y repeats along rows. if v, _ := FloatAt(yg, 2, 1); v != 30 { t.Errorf("Y grid [2,1]: %v", v) } } func TestCrossProduct(t *testing.T) { u, _ := FromFloats([]float64{1, 0, 0}, 3) v, _ := FromFloats([]float64{0, 1, 0}, 3) c, err := CrossProduct(u, v) if err != nil { t.Fatal(err) } want, _ := FromFloats([]float64{0, 0, 1}, 3) if !Equal(want, c) { t.Errorf("cross: got %v", c.RawFloats()) } } func TestIntegrate(t *testing.T) { y, _ := FromFloats([]float64{0, 1, 2, 3}, 4) v, err := Integrate(y, 1) if err != nil { t.Fatal(err) } if math.Abs(v-4.5) > 1e-9 { t.Errorf("trapezoid of ramp: %v, want 4.5", v) } ci, err := CumulativeIntegrate(y, 1) if err != nil { t.Fatal(err) } if ci.Len() != 4 || ci.FloatAt(3) != 4.5 || ci.FloatAt(0) != 0 { t.Errorf("cumulative integrate: %v", ci.RawFloats()) } } func TestInterpolate(t *testing.T) { xs, _ := FromFloats([]float64{0, 10}, 2) ys, _ := FromFloats([]float64{0, 100}, 2) q, _ := FromFloats([]float64{-5, 5, 15}, 3) out, err := Interpolate(xs, ys, q) if err != nil { t.Fatal(err) } for i, want := range []float64{0, 50, 100} { if v := out.FloatAt(i); math.Abs(v-want) > 1e-9 { t.Errorf("interp[%d]: %v, want %v", i, v, want) } } } func TestMoveAxis(t *testing.T) { a, _ := FromFloats(make([]float64, 24), 2, 3, 4) for i := range a.Len() { a.SetFloatAt(i, float64(i)) } moved, err := MoveAxis(a, 2, 0) if err != nil { t.Fatal(err) } if moved.Shape()[0] != 4 { t.Fatalf("MoveAxis shape: %v", moved.Shape()) } // Element at output [0, d, h] equals input [d, h, 0]. vOut, _ := FloatAt(moved, 0, 1, 1) vIn, _ := FloatAt(a, 1, 1, 0) if vOut != vIn { t.Errorf("MoveAxis mapping broken: %v vs %v", vOut, vIn) } } func TestSearchSorted(t *testing.T) { hay, _ := FromFloats([]float64{10, 20, 30, 40}, 4) needles, _ := FromFloats([]float64{15, 5, 50, 20}, 4) out, err := SearchSorted(hay, needles) if err != nil { t.Fatal(err) } // The rightmost rule: an exact hit lands after its equals, so 20 // inserts at 2, not before its twin at 1. want := []int64{1, 0, 4, 2} for i := range want { if v := out.RawInts()[i]; v != want[i] { t.Errorf("searchsorted[%d]: %v, want %v", i, v, want[i]) } } // Duplicates in the haystack: every equal element is skipped. dupHay, _ := FromFloats([]float64{1, 2, 2, 3}, 4) dupNeedles, _ := FromFloats([]float64{2}, 1) dupOut, err := SearchSorted(dupHay, dupNeedles) if err != nil { t.Fatal(err) } if got := dupOut.RawInts()[0]; got != 3 { t.Errorf("searchsorted duplicates: %d, want 3", got) } if _, err := SearchSorted( mustFromComplexes(t, []complex128{1}, 1), mustFromFloats(t, []float64{1}, 1)); err == nil { t.Error("searchsorted complex haystack: expected error") } } func TestAssignBins(t *testing.T) { edges, err0 := FromFloats([]float64{0, 10, 20}, 3) if err0 != nil { t.Fatal(err0) } vals, _ := FromFloats([]float64{-3, 5, 12, 25}, 4) bins, err := AssignBins(vals, edges) if err != nil { t.Fatal(err) } wantBins := []int64{0, 0, 1, 1} for i := range wantBins { if v := bins.RawInts()[i]; v != wantBins[i] { t.Errorf("bin[%d]: %v, want %v", i, v, wantBins[i]) } } } func TestCovarianceCorrelation(t *testing.T) { x, _ := FromFloats([]float64{1, 2, 3, 4}, 4) y, _ := FromFloats([]float64{2, 4, 6, 8}, 4) cov, err := Covariance(x, y) if err != nil { t.Fatal(err) } if math.Abs(cov-10.0/3.0) > 1e-9 { t.Errorf("covariance: %v, want 10/3", cov) } r, err := Correlation(x, y) if err != nil { t.Fatal(err) } if math.Abs(r-1) > 1e-9 { t.Errorf("correlation of identical trend: %v, want 1", r) } yFlip, _ := FromFloats([]float64{8, 6, 4, 2}, 4) rNeg, _ := Correlation(x, yFlip) if math.Abs(rNeg+1) > 1e-9 { t.Errorf("anti-correlation: %v, want -1", rNeg) } } // TestBytesDegenerateInputs pins the guard that used to panic: Bytes // on non-int payloads. func TestBytesDegenerateInputs(t *testing.T) { f := mustFromFloats(t, []float64{1, 2}, 2) if b := f.Bytes(); b != nil { t.Errorf("Bytes on float array: got %v, want nil", b) } i := mustFromInts(t, []int64{65, 66}, 2) if b := i.Bytes(); string(b) != "AB" { t.Errorf("Bytes on int array: got %q, want %q", b, "AB") } } // TestSpMulValidatesAndPromotes pins SpMul's guards: shape mismatch, // out-of-range indices and mixed dtypes used to panic through nil // payloads instead of erroring or promoting. func TestSpMulValidatesAndPromotes(t *testing.T) { vals := mustFromFloats(t, []float64{2, 3}, 2) idx := mustFromInts(t, []int64{0, 1, 1, 0}, 2, 2) sp := &SparseCOO{Indices: idx, Values: vals, Shape: []int{2, 2}} wrongShape := mustFromFloats(t, []float64{1, 2, 3}, 3) if _, err := SpMul(sp, wrongShape); err == nil { t.Error("SpMul shape mismatch: expected error") } intDense := mustFromInts(t, []int64{10, 20, 30, 40}, 2, 2) got, err := SpMul(sp, intDense) if err != nil { t.Fatalf("SpMul mixed dtype: %v", err) } if got.Dtype() != Float { t.Fatalf("SpMul promote dtype: %s", got.Dtype()) } // Entry 0: value 2 at (0,1) -> 2*20 = 40; entry 1: value 3 at // (1,0) -> 3*30 = 90. if v, _ := FloatAt(got, 0, 1); v != 40 { t.Errorf("SpMul [0,1]: %v, want 40", v) } if v, _ := FloatAt(got, 1, 0); v != 90 { t.Errorf("SpMul [1,0]: %v, want 90", v) } }