Files
tensor/internal/core/axis_test.go
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

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)
}
}