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