267 lines
7.9 KiB
Go
267 lines
7.9 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|
// SPDX-License-Identifier: MIT
|
|
|
|
package core
|
|
|
|
import (
|
|
"math"
|
|
"slices"
|
|
"strings"
|
|
"testing"
|
|
)
|
|
|
|
func TestSumAxis(t *testing.T) {
|
|
m := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 2, 3)
|
|
|
|
// Column sums: shape (2,) gives [6, 15].
|
|
cols, err := SumAxis(m, 1)
|
|
if err != nil {
|
|
t.Fatalf("SumAxis(1): %v", err)
|
|
}
|
|
want := mustFromInts(t, []int64{6, 15}, 2)
|
|
if !Equal(want, cols) {
|
|
t.Fatalf("SumAxis(1): %s", cols)
|
|
}
|
|
|
|
// Row sums: shape (3,) gives [5, 7, 9].
|
|
rows, err := SumAxis(m, 0)
|
|
if err != nil {
|
|
t.Fatalf("SumAxis(0): %v", err)
|
|
}
|
|
wantRows := mustFromInts(t, []int64{5, 7, 9}, 3)
|
|
if !Equal(wantRows, rows) {
|
|
t.Fatalf("SumAxis(0): %s", rows)
|
|
}
|
|
|
|
// Float and complex keep their dtypes.
|
|
f := mustFromFloats(t, []float64{0.5, 1.5, 2.5, 3.5}, 2, 2)
|
|
fs, _ := SumAxis(f, 0)
|
|
if fs.Dtype() != Float {
|
|
t.Fatalf("SumAxis float dtype: %s", fs.Dtype())
|
|
}
|
|
c := mustFromComplexes(t, []complex128{1, complex(0, 1), 0, 0}, 2, 2)
|
|
cs, err := SumAxis(c, 1)
|
|
if err != nil {
|
|
t.Fatalf("SumAxis complex: %v", err)
|
|
}
|
|
if v, _ := ComplexAt(cs, 0); v != complex(1, 1) {
|
|
t.Fatalf("SumAxis complex value: %v", v)
|
|
}
|
|
|
|
if _, err := SumAxis(m, 2); err == nil || !strings.Contains(err.Error(), "out of range") {
|
|
t.Fatalf("SumAxis dim: %v", err)
|
|
}
|
|
v := mustFromInts(t, []int64{1, 2}, 2)
|
|
if _, err := SumAxis(v, 0); err == nil || !strings.Contains(err.Error(), "global variant") {
|
|
t.Fatalf("SumAxis 1-D: %v", err)
|
|
}
|
|
}
|
|
|
|
func TestMinMaxAxis(t *testing.T) {
|
|
m := mustFromInts(t, []int64{1, 5, 2, 6, 4, 3}, 2, 3)
|
|
|
|
mn, err := MinAxis(m, 1)
|
|
if err != nil {
|
|
t.Fatalf("MinAxis: %v", err)
|
|
}
|
|
if !Equal(mustFromInts(t, []int64{1, 3}, 2), mn) {
|
|
t.Fatalf("MinAxis: %s", mn)
|
|
}
|
|
mx, err := MaxAxis(m, 1)
|
|
if err != nil {
|
|
t.Fatalf("MaxAxis: %v", err)
|
|
}
|
|
if !Equal(mustFromInts(t, []int64{5, 6}, 2), mx) {
|
|
t.Fatalf("MaxAxis: %s", mx)
|
|
}
|
|
|
|
// Along rows.
|
|
rowsMin, _ := MinAxis(m, 0)
|
|
if !Equal(mustFromInts(t, []int64{1, 4, 2}, 3), rowsMin) {
|
|
t.Fatalf("MinAxis(0): %s", rowsMin)
|
|
}
|
|
|
|
// A zero start must never win: all-negative values.
|
|
neg := mustFromInts(t, []int64{-5, -1, -9, -2}, 2, 2)
|
|
nmin, _ := MinAxis(neg, 1)
|
|
if !Equal(mustFromInts(t, []int64{-5, -9}, 2), nmin) {
|
|
t.Fatalf("MinAxis negatives: %s", nmin)
|
|
}
|
|
nmax, _ := MaxAxis(neg, 0)
|
|
if !Equal(mustFromInts(t, []int64{-5, -1}, 2), nmax) {
|
|
t.Fatalf("MaxAxis negatives: %s", nmax)
|
|
}
|
|
|
|
// A NaN element never wins: a line starting with NaN takes its
|
|
// first finite element, and a line that is all NaN has no extreme
|
|
// and reports NaN.
|
|
f := mustFromFloats(t, []float64{math.NaN(), 1.0, 3.0}, 1, 3)
|
|
fmin, err := MinAxis(f, 1)
|
|
if err != nil {
|
|
t.Fatalf("MinAxis NaN: %v", err)
|
|
}
|
|
if v, _ := FloatAt(fmin, 0); v != 1.0 {
|
|
t.Fatalf("MinAxis NaN never wins: got %v, want 1", v)
|
|
}
|
|
fmax, err := MaxAxis(f, 1)
|
|
if err != nil {
|
|
t.Fatalf("MaxAxis NaN: %v", err)
|
|
}
|
|
if v, _ := FloatAt(fmax, 0); v != 3.0 {
|
|
t.Fatalf("MaxAxis NaN never wins: got %v, want 3", v)
|
|
}
|
|
allNaN := mustFromFloats(t, []float64{math.NaN(), math.NaN(), math.NaN(), 2.0}, 2, 2)
|
|
anmin, _ := MinAxis(allNaN, 1)
|
|
if v, _ := FloatAt(anmin, 0); !math.IsNaN(v) {
|
|
t.Fatalf("MinAxis all-NaN line: got %v, want NaN", v)
|
|
}
|
|
if v, _ := FloatAt(anmin, 1); v != 2.0 {
|
|
t.Fatalf("MinAxis all-NaN neighbour: got %v, want 2", v)
|
|
}
|
|
|
|
c := mustFromComplexes(t, []complex128{1, 2, 3, 4}, 2, 2)
|
|
if _, err := MinAxis(c, 0); err == nil || !strings.Contains(err.Error(), "no ordering") {
|
|
t.Fatalf("MinAxis complex: %v", err)
|
|
}
|
|
}
|
|
|
|
func TestMeanAxis(t *testing.T) {
|
|
m := mustFromInts(t, []int64{1, 2, 3, 4, 5, 6}, 2, 3)
|
|
|
|
mean, err := MeanAxis(m, 1)
|
|
if err != nil {
|
|
t.Fatalf("MeanAxis: %v", err)
|
|
}
|
|
if mean.Dtype() != Float {
|
|
t.Fatalf("MeanAxis dtype: %s", mean.Dtype())
|
|
}
|
|
want := []float64{2.0, 5.0} // (1+2+3)/3, (4+5+6)/3
|
|
for i := range 2 {
|
|
if v, _ := FloatAt(mean, i); v != want[i] {
|
|
t.Fatalf("MeanAxis[%d]: %v", i, v)
|
|
}
|
|
}
|
|
|
|
byRows, _ := MeanAxis(m, 0)
|
|
wantRows := []float64{2.5, 3.5, 4.5} // column means of [[1,2,3],[4,5,6]]
|
|
for i := range 3 {
|
|
if v, _ := FloatAt(byRows, i); v != wantRows[i] {
|
|
t.Fatalf("MeanAxis(0)[%d]: %v", i, v)
|
|
}
|
|
}
|
|
|
|
c := mustFromComplexes(t, []complex128{1, 2, 3, 4}, 2, 2)
|
|
if _, err := MeanAxis(c, 0); err == nil || !strings.Contains(err.Error(), "no float mean") {
|
|
t.Fatalf("MeanAxis complex: %v", err)
|
|
}
|
|
}
|
|
|
|
func TestArgExtreme(t *testing.T) {
|
|
a := mustFromInts(t, []int64{3, 1, 2}, 3)
|
|
imax, err := ArgMax(a)
|
|
if err != nil || imax != 0 {
|
|
t.Fatalf("ArgMax: %d %v", imax, err)
|
|
}
|
|
imin, err := ArgMin(a)
|
|
if err != nil || imin != 1 {
|
|
t.Fatalf("ArgMin: %d %v", imin, err)
|
|
}
|
|
|
|
// NaN elements are skipped; a finite one still wins.
|
|
f := mustFromFloats(t, []float64{math.NaN(), 2.0, math.NaN(), 1.0}, 4)
|
|
fmin, err := ArgMin(f)
|
|
if err != nil || fmin != 3 {
|
|
t.Fatalf("ArgMin NaN skip: %d %v", fmin, err)
|
|
}
|
|
|
|
allNaN := mustFromFloats(t, []float64{math.NaN()}, 1)
|
|
if _, err := ArgMax(allNaN); err == nil || !strings.Contains(err.Error(), "every element is NaN") {
|
|
t.Fatalf("ArgMax all NaN: %v", err)
|
|
}
|
|
|
|
m := mustFromInts(t, []int64{1, 2}, 1, 2)
|
|
if _, err := ArgMax(m); err == nil || !strings.Contains(err.Error(), "needs a 1-D array") {
|
|
t.Fatalf("ArgMax 2-D: %v", err)
|
|
}
|
|
c := mustFromComplexes(t, []complex128{1}, 1)
|
|
if _, err := ArgMin(c); err == nil || !strings.Contains(err.Error(), "no ordering") {
|
|
t.Fatalf("ArgMin complex: %v", err)
|
|
}
|
|
}
|
|
|
|
// TestAxisReductionsNarrowDtypes pins the axis rules the scalar
|
|
// reductions carry: integer-class axis sums answer Int-axis results
|
|
// under exact widening, means answer float results, extrema compare
|
|
// natively in each payload type, and the arg extremes keep their Int
|
|
// index contract.
|
|
func TestAxisReductionsNarrowDtypes(t *testing.T) {
|
|
i8, err := FromInt8s([]int8{1, 2, 3, 4, 5, 6}, 2, 3)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
sum, err := SumAxis(i8, 1)
|
|
if err != nil {
|
|
t.Fatalf("SumAxis int8: %v", err)
|
|
}
|
|
if sum.Dtype() != Int {
|
|
t.Fatalf("SumAxis int8 answered %s, want the Int axis result", sum.Dtype())
|
|
}
|
|
if want := []int64{6, 15}; !slices.Equal(sum.RawInts(), want) {
|
|
t.Fatalf("SumAxis int8 = %v, want %v", sum.RawInts(), want)
|
|
}
|
|
mean, err := MeanAxis(i8, 1)
|
|
if err != nil || mean.Dtype() != Float {
|
|
t.Fatalf("MeanAxis int8: %s %v", mean.Dtype(), err)
|
|
}
|
|
if want := []float64{2, 5}; !slices.Equal(mean.RawFloats(), want) {
|
|
t.Fatalf("MeanAxis int8 = %v, want %v", mean.RawFloats(), want)
|
|
}
|
|
mn, err := MinAxis(i8, 1)
|
|
if err != nil || mn.Dtype() != Int || !slices.Equal(mn.RawInts(), []int64{1, 4}) {
|
|
t.Fatalf("MinAxis int8 = %s %v %v", mn.Dtype(), mn.RawInts(), err)
|
|
}
|
|
mx, err := MaxAxis(i8, 1)
|
|
if err != nil || mx.Dtype() != Int || !slices.Equal(mx.RawInts(), []int64{3, 6}) {
|
|
t.Fatalf("MaxAxis int8 = %s %v %v", mx.Dtype(), mx.RawInts(), err)
|
|
}
|
|
|
|
// Bool axis sums count trues per line into the Int axis result.
|
|
bl, err := FromBools([]bool{true, false, true, true}, 2, 2)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
bs, err := SumAxis(bl, 1)
|
|
if err != nil || bs.Dtype() != Int || !slices.Equal(bs.RawInts(), []int64{1, 2}) {
|
|
t.Fatalf("SumAxis bool = %s %v %v", bs.Dtype(), bs.RawInts(), err)
|
|
}
|
|
bmn, err := MinAxis(bl, 1)
|
|
if err != nil || !slices.Equal(bmn.RawInts(), []int64{0, 1}) {
|
|
t.Fatalf("MinAxis bool = %v %v", bmn.RawInts(), err)
|
|
}
|
|
|
|
// The arg extremes compare natively per payload dtype and keep the
|
|
// Int index contract.
|
|
i16, err := FromInt16s([]int16{1, 9, 3, 8, 2, 7}, 2, 3)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
amax, err := ArgMaxAxis(i16, 1)
|
|
if err != nil || amax.Dtype() != Int || !slices.Equal(amax.RawInts(), []int64{1, 0}) {
|
|
t.Fatalf("ArgMaxAxis int16 = %s %v %v", amax.Dtype(), amax.RawInts(), err)
|
|
}
|
|
u32, err := FromUint32s([]uint32{1, 5, 3}, 3)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if idx, err := ArgMax(u32); err != nil || idx != 1 {
|
|
t.Fatalf("ArgMax uint32 = %d %v, want 1", idx, err)
|
|
}
|
|
bl1, err := FromBools([]bool{true, false, true}, 3)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if idx, err := ArgMin(bl1); err != nil || idx != 1 {
|
|
t.Fatalf("ArgMin bool = %d %v, want the first false at 1", idx, err)
|
|
}
|
|
}
|