194 lines
4.8 KiB
Go
194 lines
4.8 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|||
|
|
// SPDX-License-Identifier: MIT
|
||
|
|
|
||
|
|
package core
|
||
|
|
|
||
|
|
import (
|
||
|
|
"strings"
|
||
|
|
"testing"
|
||
|
|
)
|
||
|
|
|
||
|
|
func mustFromInts(t *testing.T, vals []int64, shape ...int) *Array {
|
||
|
|
t.Helper()
|
||
|
|
a, err := FromInts(vals, shape...)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("FromInts(%v, %v): %v", vals, shape, err)
|
||
|
|
}
|
||
|
|
return a
|
||
|
|
}
|
||
|
|
|
||
|
|
func mustFromFloats(t *testing.T, vals []float64, shape ...int) *Array {
|
||
|
|
t.Helper()
|
||
|
|
a, err := FromFloats(vals, shape...)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("FromFloats(%v, %v): %v", vals, shape, err)
|
||
|
|
}
|
||
|
|
return a
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestConstructors(t *testing.T) {
|
||
|
|
a := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 2, 3)
|
||
|
|
if a.NDim() != 2 || a.Len() != 6 || a.Dtype() != Int {
|
||
|
|
t.Fatalf("accessors: ndim %d len %d dtype %s", a.NDim(), a.Len(), a.Dtype())
|
||
|
|
}
|
||
|
|
shape := a.Shape()
|
||
|
|
if len(shape) != 2 || shape[0] != 2 || shape[1] != 3 {
|
||
|
|
t.Fatalf("shape: %v", shape)
|
||
|
|
}
|
||
|
|
|
||
|
|
f := mustFromFloats(t, []float64{1.5, 2.5}, 2)
|
||
|
|
if f.Dtype() != Float || f.Len() != 2 {
|
||
|
|
t.Fatalf("float array: %s len %d", f.Dtype(), f.Len())
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestConstructorErrors(t *testing.T) {
|
||
|
|
if _, err := FromInts(nil); err == nil || !strings.Contains(err.Error(), "at least one dimension") {
|
||
|
|
t.Fatalf("no shape: %v", err)
|
||
|
|
}
|
||
|
|
if _, err := FromInts([]int64{1}, 2); err == nil || !strings.Contains(err.Error(), "do not fill the shape") {
|
||
|
|
t.Fatalf("wrong fill: %v", err)
|
||
|
|
}
|
||
|
|
if _, err := FromInts([]int64{1}, -1); err == nil || !strings.Contains(err.Error(), "zero or greater") {
|
||
|
|
t.Fatalf("negative dim: %v", err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestFilledConstructors(t *testing.T) {
|
||
|
|
z, err := Zeros(Int, 2, 2)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("Zeros: %v", err)
|
||
|
|
}
|
||
|
|
if v, _ := IntAt(z, 1, 1); v != 0 {
|
||
|
|
t.Fatalf("Zeros value: %d", v)
|
||
|
|
}
|
||
|
|
|
||
|
|
o, err := Ones(Float, 3)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("Ones: %v", err)
|
||
|
|
}
|
||
|
|
if v, _ := FloatAt(o, 2); v != 1.0 {
|
||
|
|
t.Fatalf("Ones value: %v", v)
|
||
|
|
}
|
||
|
|
|
||
|
|
fi, err := FullI(7, 2)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("FullI: %v", err)
|
||
|
|
}
|
||
|
|
if v, _ := IntAt(fi, 0); v != 7 {
|
||
|
|
t.Fatalf("FullI value: %d", v)
|
||
|
|
}
|
||
|
|
|
||
|
|
ff, err := FullF(2.5, 2, 2)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("FullF: %v", err)
|
||
|
|
}
|
||
|
|
if v, _ := FloatAt(ff, 1, 0); v != 2.5 {
|
||
|
|
t.Fatalf("FullF value: %v", v)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestRange(t *testing.T) {
|
||
|
|
r := mustFromInts(t, []int64{0, 1, 2, 3}, 4)
|
||
|
|
got, err := Range(0, 4)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("Range: %v", err)
|
||
|
|
}
|
||
|
|
if !Equal(r, got) {
|
||
|
|
t.Fatalf("Range: %s", got)
|
||
|
|
}
|
||
|
|
|
||
|
|
by, err := RangeBy(10, 0, -3)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("RangeBy: %v", err)
|
||
|
|
}
|
||
|
|
want := mustFromInts(t, []int64{10, 7, 4, 1}, 4)
|
||
|
|
if !Equal(want, by) {
|
||
|
|
t.Fatalf("RangeBy negative: %s", by)
|
||
|
|
}
|
||
|
|
|
||
|
|
empty, err := Range(5, 5)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("Range empty: %v", err)
|
||
|
|
}
|
||
|
|
if empty.Len() != 0 {
|
||
|
|
t.Fatalf("Range empty len: %d", empty.Len())
|
||
|
|
}
|
||
|
|
|
||
|
|
if _, err := RangeBy(0, 5, 0); err == nil || !strings.Contains(err.Error(), "step cannot be zero") {
|
||
|
|
t.Fatalf("RangeBy zero step: %v", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
// The walk stops instead of looping when the next value would wrap.
|
||
|
|
big, err := RangeBy(9223372036854775805, 9223372036854775807, 2)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("RangeBy wrap: %v", err)
|
||
|
|
}
|
||
|
|
if big.Len() != 1 {
|
||
|
|
t.Fatalf("RangeBy wrap len: %d", big.Len())
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestEqual(t *testing.T) {
|
||
|
|
a := mustFromInts(t, []int64{1, 2}, 2)
|
||
|
|
b := mustFromInts(t, []int64{1, 2}, 2)
|
||
|
|
c := mustFromInts(t, []int64{1, 2, 3}, 3)
|
||
|
|
f := mustFromFloats(t, []float64{1, 2}, 2)
|
||
|
|
|
||
|
|
if !Equal(a, b) {
|
||
|
|
t.Fatalf("equal arrays must compare equal")
|
||
|
|
}
|
||
|
|
if Equal(a, c) {
|
||
|
|
t.Fatalf("different shapes must not compare equal")
|
||
|
|
}
|
||
|
|
if Equal(a, f) {
|
||
|
|
t.Fatalf("int 1 must not equal float 1.0")
|
||
|
|
}
|
||
|
|
if Equal(a, nil) {
|
||
|
|
t.Fatalf("an array never equals nil")
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestImmutability(t *testing.T) {
|
||
|
|
vals := []int64{1, 2, 3}
|
||
|
|
a := mustFromInts(t, vals, 3)
|
||
|
|
|
||
|
|
// The constructor copies: later changes to vals never reach the array.
|
||
|
|
vals[0] = 99
|
||
|
|
if v, _ := IntAt(a, 0); v != 1 {
|
||
|
|
t.Fatalf("FromInts must copy: got %d", v)
|
||
|
|
}
|
||
|
|
|
||
|
|
// Shape returns a copy, so mutating it cannot corrupt the array.
|
||
|
|
shape := a.Shape()
|
||
|
|
shape[0] = 100
|
||
|
|
if a.Shape()[0] != 3 {
|
||
|
|
t.Fatalf("Shape must copy: got %d", a.Shape()[0])
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestString(t *testing.T) {
|
||
|
|
a := mustFromInts(t, []int64{1, 2, 3, 4}, 2, 2)
|
||
|
|
if got := a.String(); got != "int (2, 2) [1, 2, 3, 4]" {
|
||
|
|
t.Fatalf("String: %q", got)
|
||
|
|
}
|
||
|
|
f := mustFromFloats(t, []float64{1.5, 2}, 2)
|
||
|
|
if got := f.String(); got != "float (2) [1.5, 2]" {
|
||
|
|
t.Fatalf("String float: %q", got)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestStringMultidim(t *testing.T) {
|
||
|
|
// Ranks 1 and 2 stay flat; rank 3+ wraps per trailing dimension,
|
||
|
|
// with size-1 dimensions transparent.
|
||
|
|
c, _ := FromInts([]int64{1, 2, 3, 4, 5, 6, 7, 8}, 2, 2, 2)
|
||
|
|
want := "int (2, 2, 2) [[[1, 2], [3, 4]], [[5, 6], [7, 8]]]"
|
||
|
|
if got := c.String(); got != want {
|
||
|
|
t.Errorf("3-D:\n got %q\nwant %q", got, want)
|
||
|
|
}
|
||
|
|
s, _ := FromFloats([]float64{7, 9, 11}, 1, 1, 3)
|
||
|
|
if got := s.String(); got != "float (1, 1, 3) [7, 9, 11]" {
|
||
|
|
t.Errorf("size-1 dims: %q", got)
|
||
|
|
}
|
||
|
|
}
|