318 lines
9.2 KiB
Go
318 lines
9.2 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|
// SPDX-License-Identifier: MIT
|
|
|
|
package core
|
|
|
|
import (
|
|
"slices"
|
|
"strings"
|
|
"testing"
|
|
)
|
|
|
|
func TestIndexing(t *testing.T) {
|
|
a := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 2, 3)
|
|
if v, err := IntAt(a, 0, 0); err != nil || v != 1 {
|
|
t.Fatalf("IntAt(0,0): %d %v", v, err)
|
|
}
|
|
if v, err := IntAt(a, 1, 2); err != nil || v != 6 {
|
|
t.Fatalf("IntAt(1,2): %d %v", v, err)
|
|
}
|
|
// Row-major storage: (0,2) is the third element.
|
|
if v, _ := IntAt(a, 0, 2); v != 3 {
|
|
t.Fatalf("IntAt(0,2): %d", v)
|
|
}
|
|
|
|
f := mustFromFloats(t, []float64{1.5}, 1)
|
|
if _, err := IntAt(f, 0); err == nil || !strings.Contains(err.Error(), "not int") {
|
|
t.Fatalf("IntAt on float: %v", err)
|
|
}
|
|
// FloatAt widens every real dtype, int included: the widening-reader
|
|
// contract the dtype surface settled.
|
|
if v, err := FloatAt(a, 0, 0); err != nil || v != 1 {
|
|
t.Fatalf("FloatAt on int: %v %v", v, err)
|
|
}
|
|
if _, err := IntAt(a, 2, 0); err == nil || !strings.Contains(err.Error(), "out of range") {
|
|
t.Fatalf("index out of range: %v", err)
|
|
}
|
|
if _, err := IntAt(a, 0); err == nil || !strings.Contains(err.Error(), "does not match the shape") {
|
|
t.Fatalf("wrong arity: %v", err)
|
|
}
|
|
}
|
|
|
|
func TestFunctionalUpdate(t *testing.T) {
|
|
a := mustFromInts(t, []int64{1, 2}, 2)
|
|
updated, err := WithInt(a, 9, 0)
|
|
if err != nil {
|
|
t.Fatalf("WithInt: %v", err)
|
|
}
|
|
if v, _ := IntAt(updated, 0); v != 9 {
|
|
t.Fatalf("updated: %d", v)
|
|
}
|
|
// The receiver stays untouched.
|
|
if v, _ := IntAt(a, 0); v != 1 {
|
|
t.Fatalf("receiver mutated: %d", v)
|
|
}
|
|
|
|
f := mustFromFloats(t, []float64{1.5, 2.5}, 2)
|
|
uf, err := WithFloat(f, 9.5, 1)
|
|
if err != nil {
|
|
t.Fatalf("WithFloat: %v", err)
|
|
}
|
|
if v, _ := FloatAt(uf, 1); v != 9.5 {
|
|
t.Fatalf("WithFloat value: %v", v)
|
|
}
|
|
if v, _ := FloatAt(f, 1); v != 2.5 {
|
|
t.Fatalf("WithFloat receiver mutated: %v", v)
|
|
}
|
|
|
|
if _, err := WithInt(f, 1, 0); err == nil || !strings.Contains(err.Error(), "not int") {
|
|
t.Fatalf("WithInt on float: %v", err)
|
|
}
|
|
if _, err := WithFloat(a, 1.0, 0); err == nil || !strings.Contains(err.Error(), "not float") {
|
|
t.Fatalf("WithFloat on int: %v", err)
|
|
}
|
|
if _, err := WithInt(a, 1, 5); err == nil || !strings.Contains(err.Error(), "out of range") {
|
|
t.Fatalf("WithInt range: %v", err)
|
|
}
|
|
}
|
|
|
|
func TestSlice(t *testing.T) {
|
|
a := mustFromInts(t, []int64{1, 2, 3, 4}, 4)
|
|
|
|
s, err := Slice(a, 0, 1, 3)
|
|
if err != nil {
|
|
t.Fatalf("Slice: %v", err)
|
|
}
|
|
if s.Len() != 2 {
|
|
t.Fatalf("Slice len: %d", s.Len())
|
|
}
|
|
if v, _ := IntAt(s, 0); v != 2 {
|
|
t.Fatalf("Slice(0): %d", v)
|
|
}
|
|
if v, _ := IntAt(s, 1); v != 3 {
|
|
t.Fatalf("Slice(1): %d", v)
|
|
}
|
|
|
|
// A functional update through the view leaves the source untouched:
|
|
// WithInt copies rather than writing the shared payload.
|
|
updated, _ := WithInt(s, 99, 0)
|
|
if v, _ := IntAt(a, 1); v != 2 {
|
|
t.Fatalf("WithInt through a slice view mutated the source: %d", v)
|
|
}
|
|
if v, _ := IntAt(updated, 0); v != 99 {
|
|
t.Fatalf("updated slice: %d", v)
|
|
}
|
|
|
|
// Slicing a dimension of a 2-D array picks whole rows.
|
|
m := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 2, 3)
|
|
rows, err := Slice(m, 0, 1, 2)
|
|
if err != nil {
|
|
t.Fatalf("Slice rows: %v", err)
|
|
}
|
|
want := mustFromInts(t, []int64{4, 5, 6}, 1, 3)
|
|
if !Equal(want, rows) {
|
|
t.Fatalf("rows: %s", rows)
|
|
}
|
|
|
|
cols, err := Slice(m, 1, 0, 2)
|
|
if err != nil {
|
|
t.Fatalf("Slice cols: %v", err)
|
|
}
|
|
wantCols := mustFromInts(t, []int64{1, 2, 4, 5}, 2, 2)
|
|
if !Equal(wantCols, cols) {
|
|
t.Fatalf("cols: %s", cols)
|
|
}
|
|
|
|
// An empty range yields an empty array of the right shape.
|
|
none, err := Slice(a, 0, 2, 2)
|
|
if err != nil || none.Len() != 0 || none.Shape()[0] != 0 {
|
|
t.Fatalf("empty slice: %s %v", none, err)
|
|
}
|
|
|
|
if _, err := Slice(a, 1, 0, 1); err == nil || !strings.Contains(err.Error(), "dimension 1 is out of range") {
|
|
t.Fatalf("Slice dim: %v", err)
|
|
}
|
|
if _, err := Slice(a, 0, 3, 2); err == nil || !strings.Contains(err.Error(), "out of range") {
|
|
t.Fatalf("Slice reversed: %v", err)
|
|
}
|
|
if _, err := Slice(a, 0, 0, 5); err == nil || !strings.Contains(err.Error(), "out of range") {
|
|
t.Fatalf("Slice beyond: %v", err)
|
|
}
|
|
}
|
|
|
|
// TestWideningReadersAndUpdates pins the widening-reader contract of
|
|
// the dtypes: IntAt reads the integer class and bool with exact widenings,
|
|
// FloatAt reads every real dtype, ComplexAt keeps its gate, WithInt
|
|
// writes every integer-class dtype with the implicit-store cast, and the
|
|
// kind gates the suite pins stay loud.
|
|
func TestWideningReadersAndUpdates(t *testing.T) {
|
|
i8, err := FromInt8s([]int8{-2, 3}, 2)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
bl, err := FromBools([]bool{true, false}, 2)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if v, err := IntAt(i8, 0); err != nil || v != -2 {
|
|
t.Fatalf("IntAt int8: %d %v", v, err)
|
|
}
|
|
if v, err := IntAt(bl, 0); err != nil || v != 1 {
|
|
t.Fatalf("IntAt bool: %d %v", v, err)
|
|
}
|
|
if v, err := FloatAt(i8, 1); err != nil || v != 3 {
|
|
t.Fatalf("FloatAt int8: %v %v", v, err)
|
|
}
|
|
if v, err := FloatAt(bl, 0); err != nil || v != 1 {
|
|
t.Fatalf("FloatAt bool: %v %v", v, err)
|
|
}
|
|
f32, err := FromFloat32s([]float32{1.5}, 1)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if v, err := FloatAt(f32, 0); err != nil || v != 1.5 {
|
|
t.Fatalf("FloatAt float32: %v %v", v, err)
|
|
}
|
|
if _, err := ComplexAt(i8, 0); err == nil || !strings.Contains(err.Error(), "not complex") {
|
|
t.Fatalf("ComplexAt int8: %v", err)
|
|
}
|
|
|
|
// WithInt writes the integer class with the implicit-store cast.
|
|
w, err := WithInt(i8, 300, 1)
|
|
if err != nil {
|
|
t.Fatalf("WithInt int8: %v", err)
|
|
}
|
|
stored := int64(300)
|
|
if got := w.RawInt8s()[1]; got != int8(stored) {
|
|
t.Fatalf("WithInt int8 = %d, want the wrapped %d", got, int8(stored))
|
|
}
|
|
if got := i8.RawInt8s()[1]; got != 3 {
|
|
t.Fatalf("WithInt int8 mutated its receiver: %d", got)
|
|
}
|
|
wb, err := WithInt(bl, 2, 0)
|
|
if err != nil || !wb.RawBools()[0] {
|
|
t.Fatalf("WithInt bool(2): %v %v", wb.RawBools(), err)
|
|
}
|
|
wb, err = WithInt(bl, 0, 1)
|
|
if err != nil || wb.RawBools()[1] {
|
|
t.Fatalf("WithInt bool(0): %v %v", wb.RawBools(), err)
|
|
}
|
|
// The kind gates the suite pins stay loud for the other receivers.
|
|
if _, err := WithFloat(i8, 1, 0); err == nil || !strings.Contains(err.Error(), "not float") {
|
|
t.Fatalf("WithFloat int8: %v", err)
|
|
}
|
|
if _, err := WithComplex(i8, 1, 0); err == nil || !strings.Contains(err.Error(), "not complex") {
|
|
t.Fatalf("WithComplex int8: %v", err)
|
|
}
|
|
|
|
// Slice carries the narrow dtypes: the view path rebases the narrow
|
|
// payload and the copy path walks setFrom.
|
|
m8, err := FromInt8s([]int8{1, 2, 3, 4, 5, 6}, 2, 3)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
view, err := Slice(m8, 0, 0, 1)
|
|
if err != nil {
|
|
t.Fatalf("Slice int8 view: %v", err)
|
|
}
|
|
if v, err := IntAt(view, 0, 2); err != nil || v != 3 {
|
|
t.Fatalf("Slice int8 view element: %d %v", v, err)
|
|
}
|
|
cols, err := Slice(m8, 1, 0, 2)
|
|
if err != nil {
|
|
t.Fatalf("Slice int8 copy: %v", err)
|
|
}
|
|
if want := []int8{1, 2, 4, 5}; !slices.Equal(cols.RawInt8s(), want) {
|
|
t.Fatalf("Slice int8 copy = %v, want %v", cols.RawInt8s(), want)
|
|
}
|
|
}
|
|
|
|
// TestAdvancedIndexingNarrowDtypes pins the advanced-indexing contract:
|
|
// Scatter builds its output through cloneArray so narrow and bool
|
|
// destinations carry payloads, Gather and Take read and write the narrow
|
|
// payloads, and Nonzero sees zeros in them.
|
|
func TestAdvancedIndexingNarrowDtypes(t *testing.T) {
|
|
idx, err := FromInts([]int64{0, 1}, 2, 1)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
selfB, err := FromBools([]bool{false, true, false, true}, 2, 2)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
srcB, err := FromBools([]bool{true, false}, 2, 1)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
out, err := Scatter(selfB, 1, idx, srcB)
|
|
if err != nil {
|
|
t.Fatalf("Scatter bool: %v", err)
|
|
}
|
|
if want := []bool{true, true, false, false}; !slices.Equal(out.RawBools(), want) {
|
|
t.Fatalf("Scatter bool = %v, want %v", out.RawBools(), want)
|
|
}
|
|
|
|
// An int8 destination from an int source stores with the
|
|
// implicit-store cast, exactly as setConverted stores it.
|
|
self8, err := FromInt8s([]int8{1, 2, 3, 4}, 2, 2)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
srcI, err := FromInts([]int64{300, 9}, 2, 1)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
out8, err := Scatter(self8, 1, idx, srcI)
|
|
if err != nil {
|
|
t.Fatalf("Scatter int8 from int: %v", err)
|
|
}
|
|
wrap := srcI.RawInts()[0]
|
|
if want := []int8{int8(wrap), 2, 3, 9}; !slices.Equal(out8.RawInt8s(), want) {
|
|
t.Fatalf("Scatter int8 from int = %v, want %v", out8.RawInt8s(), want)
|
|
}
|
|
|
|
src8, err := FromInt8s([]int8{5, 6, 7, 8}, 2, 2)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
gidx, err := FromInts([]int64{1, 0}, 2, 1)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
g, err := Gather(src8, 1, gidx)
|
|
if err != nil {
|
|
t.Fatalf("Gather int8: %v", err)
|
|
}
|
|
if want := []int8{6, 7}; !slices.Equal(g.RawInt8s(), want) {
|
|
t.Fatalf("Gather int8 = %v, want %v", g.RawInt8s(), want)
|
|
}
|
|
|
|
tb, err := FromBools([]bool{true, false, true}, 3)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
tidx, err := FromInts([]int64{2, 0}, 2)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
tk, err := Take(tb, tidx)
|
|
if err != nil {
|
|
t.Fatalf("Take bool: %v", err)
|
|
}
|
|
if want := []bool{true, true}; !slices.Equal(tk.RawBools(), want) {
|
|
t.Fatalf("Take bool = %v, want %v", tk.RawBools(), want)
|
|
}
|
|
|
|
nz8, err := FromInt8s([]int8{0, 2, 0, 0}, 4)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
nz, err := Nonzero(nz8)
|
|
if err != nil {
|
|
t.Fatalf("Nonzero int8: %v", err)
|
|
}
|
|
if want := []int{1}; !slices.Equal(nz[0], want) {
|
|
t.Fatalf("Nonzero int8 = %v, want %v", nz[0], want)
|
|
}
|
|
}
|