216 lines
5.7 KiB
Go
216 lines
5.7 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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)
|
||
|
|
}
|
||
|
|
}
|