feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,215 @@
|
||||
// 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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user