398 lines
11 KiB
Go
398 lines
11 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|||
|
|
// SPDX-License-Identifier: MIT
|
||
|
|
|
||
|
|
package core
|
||
|
|
|
||
|
|
import (
|
||
|
|
"math"
|
||
|
|
"strings"
|
||
|
|
"testing"
|
||
|
|
)
|
||
|
|
|
||
|
|
func TestLinspace(t *testing.T) {
|
||
|
|
v, err := Linspace(0, 1, 5)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
want := []float64{0, 0.25, 0.5, 0.75, 1}
|
||
|
|
for i := range 5 {
|
||
|
|
if g, _ := FloatAt(v, i); math.Abs(g-want[i]) > 1e-12 {
|
||
|
|
t.Errorf("Linspace[%d]: got %v, want %v", i, g, want[i])
|
||
|
|
}
|
||
|
|
}
|
||
|
|
one, _ := Linspace(5, 9, 1)
|
||
|
|
if v, _ := FloatAt(one, 0); v != 5 {
|
||
|
|
t.Errorf("Linspace n=1: got %v, want 5", v)
|
||
|
|
}
|
||
|
|
empty, _ := Linspace(0, 1, 0)
|
||
|
|
if empty.Len() != 0 {
|
||
|
|
t.Errorf("Linspace n=0: len %d", empty.Len())
|
||
|
|
}
|
||
|
|
if _, err := Linspace(0, 1, -1); err == nil {
|
||
|
|
t.Error("Linspace: expected error for negative n")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestRepeat(t *testing.T) {
|
||
|
|
a, _ := FromInts([]int64{1, 2, 3}, 3)
|
||
|
|
r, err := Repeat(a, 2, 0)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if !Equal(mustFromInts(t, []int64{1, 1, 2, 2, 3, 3}, 6), r) {
|
||
|
|
t.Errorf("Repeat: %v", r.RawInts())
|
||
|
|
}
|
||
|
|
m, _ := FromInts([]int64{1, 2, 3, 4}, 2, 2)
|
||
|
|
r2, err := Repeat(m, 2, 0)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
// Repeating along dim 0 duplicates whole rows.
|
||
|
|
if !Equal(mustFromInts(t, []int64{1, 2, 1, 2, 3, 4, 3, 4}, 4, 2), r2) {
|
||
|
|
t.Errorf("Repeat dim0: %v", r2.RawInts())
|
||
|
|
}
|
||
|
|
if _, err := Repeat(m, 2, 5); err == nil {
|
||
|
|
t.Error("Repeat: expected error for out-of-range dim")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestTile(t *testing.T) {
|
||
|
|
a, _ := FromInts([]int64{1, 2, 3}, 3)
|
||
|
|
tw, err := Tile(a, 2)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if !Equal(mustFromInts(t, []int64{1, 2, 3, 1, 2, 3}, 6), tw) {
|
||
|
|
t.Errorf("Tile: %v", tw.RawInts())
|
||
|
|
}
|
||
|
|
m, _ := FromInts([]int64{1, 2, 3, 4}, 2, 2)
|
||
|
|
t4, err := Tile(m, 2, 2)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
want := []int64{1, 2, 1, 2, 3, 4, 3, 4, 1, 2, 1, 2, 3, 4, 3, 4}
|
||
|
|
if !Equal(mustFromInts(t, want, 4, 4), t4) {
|
||
|
|
t.Errorf("Tile 2x2: %v", t4.RawInts())
|
||
|
|
}
|
||
|
|
if _, err := Tile(a, -1); err == nil {
|
||
|
|
t.Error("Tile: expected error for negative reps")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestFlip(t *testing.T) {
|
||
|
|
a, _ := FromInts([]int64{1, 2, 3, 4}, 2, 2)
|
||
|
|
f, err := Flip(a)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if !Equal(mustFromInts(t, []int64{4, 3, 2, 1}, 2, 2), f) {
|
||
|
|
t.Errorf("Flip all: %v", f.RawInts())
|
||
|
|
}
|
||
|
|
f0, err := Flip(a, 0)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if !Equal(mustFromInts(t, []int64{3, 4, 1, 2}, 2, 2), f0) {
|
||
|
|
t.Errorf("Flip dim0: %v", f0.RawInts())
|
||
|
|
}
|
||
|
|
if _, err := Flip(a, 9); err == nil {
|
||
|
|
t.Error("Flip: expected error for out-of-range dim")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestRoll(t *testing.T) {
|
||
|
|
a, _ := FromInts([]int64{1, 2, 3, 4, 5}, 5)
|
||
|
|
r, err := Roll(a, 2, 0)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if !Equal(mustFromInts(t, []int64{4, 5, 1, 2, 3}, 5), r) {
|
||
|
|
t.Errorf("Roll +2: %v", r.RawInts())
|
||
|
|
}
|
||
|
|
rNeg, err := Roll(a, -1, 0)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if !Equal(mustFromInts(t, []int64{2, 3, 4, 5, 1}, 5), rNeg) {
|
||
|
|
t.Errorf("Roll -1: %v", rNeg.RawInts())
|
||
|
|
}
|
||
|
|
if _, err := Roll(a, 1, 3); err == nil {
|
||
|
|
t.Error("Roll: expected error for out-of-range dim")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestUnique(t *testing.T) {
|
||
|
|
a, _ := FromInts([]int64{3, 1, 2, 1, 3, 5}, 6)
|
||
|
|
u, err := Unique(a)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if !Equal(mustFromInts(t, []int64{1, 2, 3, 5}, 4), u) {
|
||
|
|
t.Errorf("Unique int: %v", u.RawInts())
|
||
|
|
}
|
||
|
|
f, _ := FromFloats([]float64{2.5, math.NaN(), 1.5, 2.5, math.NaN()}, 5)
|
||
|
|
uf, err := Unique(f)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
// Sorted: 1.5, 2.5, then one NaN.
|
||
|
|
if v, _ := FloatAt(uf, 0); v != 1.5 {
|
||
|
|
t.Errorf("Unique float[0]: %v", v)
|
||
|
|
}
|
||
|
|
if v, _ := FloatAt(uf, 1); v != 2.5 {
|
||
|
|
t.Errorf("Unique float[1]: %v", v)
|
||
|
|
}
|
||
|
|
if v, _ := FloatAt(uf, 2); !math.IsNaN(v) {
|
||
|
|
t.Errorf("Unique float[2]: %v, want NaN", v)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestArgwhere(t *testing.T) {
|
||
|
|
a, _ := FromInts([]int64{0, 5, 0, 7}, 2, 2)
|
||
|
|
aw, err := Argwhere(a)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if aw.Shape()[0] != 2 || aw.Shape()[1] != 2 {
|
||
|
|
t.Fatalf("Argwhere shape: %v", aw.Shape())
|
||
|
|
}
|
||
|
|
// Nonzeros at (0,1) and (1,1).
|
||
|
|
if v, _ := IntAt(aw, 0, 0); v != 0 {
|
||
|
|
t.Errorf("Argwhere[0,0]: %v", v)
|
||
|
|
}
|
||
|
|
if v, _ := IntAt(aw, 0, 1); v != 1 {
|
||
|
|
t.Errorf("Argwhere[0,1]: %v", v)
|
||
|
|
}
|
||
|
|
if v, _ := IntAt(aw, 1, 1); v != 1 {
|
||
|
|
t.Errorf("Argwhere[1,1]: %v", v)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestAstype(t *testing.T) {
|
||
|
|
i, _ := FromInts([]int64{1, 2, 3}, 3)
|
||
|
|
f, err := Astype(i, Float)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if f.Dtype() != Float || f.RawFloats()[0] != 1 {
|
||
|
|
t.Errorf("Astype int to float: %s %v", f.Dtype(), f.RawFloats())
|
||
|
|
}
|
||
|
|
back, err := Astype(f, Int)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if !Equal(back, i) {
|
||
|
|
t.Errorf("Astype round-trip: %v", back.RawInts())
|
||
|
|
}
|
||
|
|
c, _ := FromComplexes([]complex128{1 + 2i, 3 - 4i}, 2)
|
||
|
|
cf, err := Astype(c, Float)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if cf.RawFloats()[0] != 1 || cf.RawFloats()[1] != 3 {
|
||
|
|
t.Errorf("Astype complex to float: %v", cf.RawFloats())
|
||
|
|
}
|
||
|
|
if _, err := Astype(c, Int); err == nil || !strings.Contains(err.Error(), "narrow") {
|
||
|
|
t.Errorf("Astype complex to int: %v", err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
// TestAstypeNarrowDtypes pins the conversion matrix the narrow element
|
||
|
|
// types add: exact widening out of them, range-checked narrowing into
|
||
|
|
// them with the loud first-failure error, the zero-test into bool, the
|
||
|
|
// exact 0/1 widening out of bool, and the legacy pairs keeping their
|
||
|
|
// historical cast semantics beside all of it.
|
||
|
|
func TestAstypeNarrowDtypes(t *testing.T) {
|
||
|
|
// Bool is the zero-test target: NaN reads true, and bool widens out
|
||
|
|
// as 0/1, exact into complex.
|
||
|
|
f, _ := FromFloats([]float64{0, 1.5, math.NaN()}, 3)
|
||
|
|
b, err := Astype(f, Bool)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if got := b.RawBools(); got[0] || !got[1] || !got[2] {
|
||
|
|
t.Errorf("Astype float to bool = %v", got)
|
||
|
|
}
|
||
|
|
cx, _ := FromComplexes([]complex128{0, 3i}, 2)
|
||
|
|
cxb, err := Astype(cx, Bool)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if got := cxb.RawBools(); got[0] || !got[1] {
|
||
|
|
t.Errorf("Astype complex to bool = %v", got)
|
||
|
|
}
|
||
|
|
cb, err := Astype(b, Complex)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if got := cb.RawComplexes(); got[0] != 0 || got[1] != 1 || got[2] != 1 {
|
||
|
|
t.Errorf("Astype bool to complex = %v", got)
|
||
|
|
}
|
||
|
|
bi, err := Astype(b, Int)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if got := bi.RawInts(); got[0] != 0 || got[1] != 1 || got[2] != 1 {
|
||
|
|
t.Errorf("Astype bool to int = %v", got)
|
||
|
|
}
|
||
|
|
|
||
|
|
// Exact widening out of the narrow integers.
|
||
|
|
i8, _ := FromInt8s([]int8{-128, 0, 127}, 3)
|
||
|
|
wInt, err := Astype(i8, Int)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if got := wInt.RawInts(); got[0] != -128 || got[2] != 127 {
|
||
|
|
t.Errorf("Astype int8 to int = %v", got)
|
||
|
|
}
|
||
|
|
wC, err := Astype(i8, Complex)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if got := wC.RawComplexes(); got[0] != complex(-128, 0) || got[2] != complex(127, 0) {
|
||
|
|
t.Errorf("Astype int8 to complex = %v", got)
|
||
|
|
}
|
||
|
|
u32, _ := FromUint32s([]uint32{4294967295}, 1)
|
||
|
|
wF, err := Astype(u32, Float)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if got := wF.RawFloats(); got[0] != 4294967295 {
|
||
|
|
t.Errorf("Astype uint32 to float = %v", got)
|
||
|
|
}
|
||
|
|
wI, err := Astype(u32, Int)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if got := wI.RawInts(); got[0] != 4294967295 {
|
||
|
|
t.Errorf("Astype uint32 to int = %v", got)
|
||
|
|
}
|
||
|
|
|
||
|
|
// Narrow to narrow: a contained value set widens exactly, a
|
||
|
|
// non-contained one fails at the first offending index.
|
||
|
|
i16, err := Astype(i8, Int16)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if got := i16.RawInt16s(); got[0] != -128 || got[2] != 127 {
|
||
|
|
t.Errorf("Astype int8 to int16 = %v", got)
|
||
|
|
}
|
||
|
|
if _, err := Astype(i8, Uint8); err == nil ||
|
||
|
|
!strings.Contains(err.Error(), "Astype: value -128 at index 0 does not fit uint8") {
|
||
|
|
t.Errorf("Astype int8 to uint8: %v", err)
|
||
|
|
}
|
||
|
|
if _, err := Astype(u32, Int32); err == nil ||
|
||
|
|
!strings.Contains(err.Error(), "Astype: value 4294967295 at index 0 does not fit int32") {
|
||
|
|
t.Errorf("Astype uint32 to int32: %v", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
// An int source checks exactly in int64 space, and the reported
|
||
|
|
// failure is the lowest index even when several elements overflow.
|
||
|
|
big, _ := FromInts([]int64{5, 300, -2}, 3)
|
||
|
|
if _, err := Astype(big, Int8); err == nil ||
|
||
|
|
!strings.Contains(err.Error(), "Astype: value 300 at index 1 does not fit int8") {
|
||
|
|
t.Errorf("Astype int to int8: %v", err)
|
||
|
|
}
|
||
|
|
two, _ := FromInts([]int64{400, 300}, 2)
|
||
|
|
if _, err := Astype(two, Int8); err == nil ||
|
||
|
|
!strings.Contains(err.Error(), "Astype: value 400 at index 0 does not fit int8") {
|
||
|
|
t.Errorf("Astype int to int8 first failure: %v", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
// A float source must be finite, integral and in range; the error
|
||
|
|
// names the value in the source's own float64 space.
|
||
|
|
ff, _ := FromFloats([]float64{2, -3, 127}, 3)
|
||
|
|
fi8, err := Astype(ff, Int8)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if got := fi8.RawInt8s(); got[0] != 2 || got[1] != -3 || got[2] != 127 {
|
||
|
|
t.Errorf("Astype float to int8 = %v", got)
|
||
|
|
}
|
||
|
|
fb, _ := FromFloats([]float64{1, 2.5}, 2)
|
||
|
|
if _, err := Astype(fb, Int8); err == nil ||
|
||
|
|
!strings.Contains(err.Error(), "Astype: value 2.5 at index 1 does not fit int8") {
|
||
|
|
t.Errorf("Astype float to int8: %v", err)
|
||
|
|
}
|
||
|
|
fInf, _ := FromFloats([]float64{math.Inf(1)}, 1)
|
||
|
|
if _, err := Astype(fInf, Uint16); err == nil ||
|
||
|
|
!strings.Contains(err.Error(), "Astype: value +Inf at index 0 does not fit uint16") {
|
||
|
|
t.Errorf("Astype +Inf to uint16: %v", err)
|
||
|
|
}
|
||
|
|
fNaN, _ := FromFloats([]float64{math.NaN()}, 1)
|
||
|
|
if _, err := Astype(fNaN, Int8); err == nil ||
|
||
|
|
!strings.Contains(err.Error(), "Astype: value NaN at index 0 does not fit int8") {
|
||
|
|
t.Errorf("Astype NaN to int8: %v", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
// A complex source into a narrow numeric target keeps the
|
||
|
|
// historical loud refusal; only bool reaches it by zero-test.
|
||
|
|
if _, err := Astype(cx, Uint8); err == nil ||
|
||
|
|
!strings.Contains(err.Error(), "cannot narrow complex to uint8") {
|
||
|
|
t.Errorf("Astype complex to uint8: %v", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
// Legacy pairs keep their cast semantics beside the new rules:
|
||
|
|
// float to int truncates, with no range error.
|
||
|
|
f3, _ := FromFloats([]float64{3.9, -1.2}, 2)
|
||
|
|
li, err := Astype(f3, Int)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if got := li.RawInts(); got[0] != 3 || got[1] != -1 {
|
||
|
|
t.Errorf("Astype float to int legacy truncation = %v", got)
|
||
|
|
}
|
||
|
|
|
||
|
|
// Same dtype copies through cloneArray for a narrow dtype; Equal
|
||
|
|
// still defaults to the complex payload for these dtypes, so the
|
||
|
|
// comparison reads the payload directly.
|
||
|
|
u8, _ := FromUint8s([]uint8{1, 2, 3}, 3)
|
||
|
|
cp, err := Astype(u8, Uint8)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if cp.Dtype() != Uint8 {
|
||
|
|
t.Fatalf("Astype uint8 to uint8 dtype = %s", cp.Dtype())
|
||
|
|
}
|
||
|
|
if got := cp.RawUint8s(); got[0] != 1 || got[1] != 2 || got[2] != 3 {
|
||
|
|
t.Errorf("Astype uint8 to uint8 = %v", got)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestItem(t *testing.T) {
|
||
|
|
a, _ := FromFloats([]float64{3.25}, 1)
|
||
|
|
v, err := Item(a)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if v != 3.25 {
|
||
|
|
t.Errorf("Item: %v", v)
|
||
|
|
}
|
||
|
|
b, _ := FromFloats([]float64{1, 2}, 2)
|
||
|
|
if _, err := Item(b); err == nil {
|
||
|
|
t.Error("Item: expected error for multi-element array")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestDiag(t *testing.T) {
|
||
|
|
a, _ := FromInts([]int64{1, 2, 3, 4}, 2, 2)
|
||
|
|
d, err := Diag(a)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if !Equal(mustFromInts(t, []int64{1, 4}, 2), d) {
|
||
|
|
t.Errorf("Diag 2-D: %v", d.RawInts())
|
||
|
|
}
|
||
|
|
v, _ := FromInts([]int64{5, 6}, 2)
|
||
|
|
m, err := Diag(v)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatal(err)
|
||
|
|
}
|
||
|
|
if m.Shape()[0] != 2 || m.Shape()[1] != 2 {
|
||
|
|
t.Fatalf("Diag 1-D shape: %v", m.Shape())
|
||
|
|
}
|
||
|
|
if val, _ := IntAt(m, 1, 1); val != 6 {
|
||
|
|
t.Errorf("Diag 1-D [1,1]: %v", val)
|
||
|
|
}
|
||
|
|
}
|