feat: initial release
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s

Assisted-by: GLM 5.3 Flash
This commit is contained in:
2026-09-03 10:00:00 +02:00
commit af4ee19703
617 changed files with 191195 additions and 0 deletions
+215
View File
@@ -0,0 +1,215 @@
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package core
import (
"math"
"testing"
)
func TestArgMaxAxis(t *testing.T) {
// (3, 4) input, argmax along dim 1 drops the dim: shape (3,) with
// each row holding the column index of its maximum.
a := mustFromFloats(t, []float64{
1, 5, 3, 4, // row 0: max at index 1
9, 2, 7, 6, // row 1: max at index 0
8, 1, 4, 2, // row 2: max at index 0
}, 3, 4)
got, err := ArgMaxAxis(a, 1)
if err != nil {
t.Fatal(err)
}
if want := []int{3}; !sameShape(got.Shape(), want) {
t.Fatalf("ArgMaxAxis shape: got %v, want %v", got.Shape(), want)
}
for i, w := range []int64{1, 0, 0} {
v, _ := IntAt(got, i)
if v != w {
t.Errorf("ArgMaxAxis dim=1 [%d]: got %d, want %d", i, v, w)
}
}
// argmax along dim 0: each column holds the row index of its max.
// col 0: rows [1, 9, 8] give idx 1; col 1: [5, 2, 1] give 0;
// col 2: [3, 7, 4] give 1; col 3: [4, 6, 2] give 1.
got, err = ArgMaxAxis(a, 0)
if err != nil {
t.Fatal(err)
}
if want := []int{4}; !sameShape(got.Shape(), want) {
t.Fatalf("ArgMaxAxis shape: got %v, want %v", got.Shape(), want)
}
want := []int64{1, 0, 1, 1}
for i, w := range want {
v, _ := IntAt(got, i)
if v != w {
t.Errorf("ArgMaxAxis dim=0 [%d]: got %d, want %d", i, v, w)
}
}
}
func TestArgMinAxis(t *testing.T) {
a := mustFromFloats(t, []float64{
1, 5, 3, 4, // row 0: min at index 0 (1)
9, 2, 7, 6, // row 1: min at index 1 (2)
8, 1, 4, 2, // row 2: min at index 1 (1)
}, 3, 4)
got, err := ArgMinAxis(a, 1)
if err != nil {
t.Fatal(err)
}
if want := []int{3}; !sameShape(got.Shape(), want) {
t.Fatalf("ArgMinAxis shape: got %v, want %v", got.Shape(), want)
}
for i, w := range []int64{0, 1, 1} {
v, _ := IntAt(got, i)
if v != w {
t.Errorf("ArgMinAxis dim=1 [%d]: got %d, want %d", i, v, w)
}
}
}
func TestArgMaxAxisErrors(t *testing.T) {
a := mustFromFloats(t, []float64{1, 2, 3}, 3)
if _, err := ArgMaxAxis(a, 0); err == nil {
t.Error("ArgMaxAxis: expected error for 1-D input")
}
// Complex rejected.
c, _ := FromComplexes([]complex128{complex(1, 0)}, 1)
if _, err := ArgMaxAxis(c, 0); err == nil {
t.Error("ArgMaxAxis: expected error for complex input")
}
}
func TestTopK(t *testing.T) {
a := mustFromFloats(t, []float64{3, 1, 4, 1, 5, 9, 2, 6}, 8)
vals, idxs, err := TopK(a, 3, 0)
if err != nil {
t.Fatal(err)
}
if vals.Len() != 3 || idxs.Len() != 3 {
t.Errorf("TopK: wrong length (%d, %d)", vals.Len(), idxs.Len())
}
// Top 3: 9 (idx 5), 6 (idx 7), 5 (idx 4).
expectVals := []float64{9, 6, 5}
expectIdx := []int64{5, 7, 4}
for i := range 3 {
v, _ := FloatAt(vals, i)
j, _ := IntAt(idxs, i)
if v != expectVals[i] {
t.Errorf("TopK vals [%d]: got %v, want %v", i, v, expectVals[i])
}
if j != expectIdx[i] {
t.Errorf("TopK idxs [%d]: got %d, want %d", i, j, expectIdx[i])
}
}
}
func TestTopK2D(t *testing.T) {
a := mustFromFloats(t, []float64{
1, 5, 3, 4,
9, 2, 7, 6,
}, 2, 4)
vals, idxs, err := TopK(a, 2, 1)
if err != nil {
t.Fatal(err)
}
if vals.Shape()[0] != 2 || vals.Shape()[1] != 2 {
t.Errorf("TopK2D shape: %v", vals.Shape())
}
// Row 0 top-2: 5 (idx 1), 4 (idx 3). Row 1 top-2: 9 (idx 0), 7 (idx 2).
wantVals := []float64{5, 4, 9, 7}
wantIdx := []int64{1, 3, 0, 2}
for i := range 4 {
v, _ := FloatAt(vals, i/2, i%2)
j, _ := IntAt(idxs, i/2, i%2)
if v != wantVals[i] {
t.Errorf("TopK2D vals [%d]: got %v, want %v", i, v, wantVals[i])
}
if j != wantIdx[i] {
t.Errorf("TopK2D idxs [%d]: got %d, want %d", i, j, wantIdx[i])
}
}
}
// TestTopK2DDim0 pins the non-trailing-dimension layout: the output
// keeps (k, suffix) row-major order, which the write offset used to
// transpose into (suffix, k).
func TestTopK2DDim0(t *testing.T) {
a := mustFromFloats(t, []float64{
1, 4,
3, 2,
5, 0,
}, 3, 2)
vals, idxs, err := TopK(a, 2, 0)
if err != nil {
t.Fatal(err)
}
if vals.Shape()[0] != 2 || vals.Shape()[1] != 2 {
t.Errorf("TopK2DDim0 shape: %v", vals.Shape())
}
// Column 0 top-2: 5 (idx 2), 3 (idx 1). Column 1 top-2: 4 (idx 0), 2 (idx 1).
wantVals := []float64{5, 4, 3, 2}
wantIdx := []int64{2, 0, 1, 1}
for i := range 4 {
v, _ := FloatAt(vals, i/2, i%2)
j, _ := IntAt(idxs, i/2, i%2)
if v != wantVals[i] {
t.Errorf("TopK2DDim0 vals [%d]: got %v, want %v", i, v, wantVals[i])
}
if j != wantIdx[i] {
t.Errorf("TopK2DDim0 idxs [%d]: got %d, want %d", i, j, wantIdx[i])
}
}
}
func TestTopKErrors(t *testing.T) {
a := mustFromFloats(t, []float64{1, 2, 3}, 3)
if _, _, err := TopK(a, 5, 0); err == nil {
t.Error("TopK: expected error for k > dim size")
}
if _, _, err := TopK(a, -1, 0); err == nil {
t.Error("TopK: expected error for negative k")
}
}
func TestTopKNaN(t *testing.T) {
// NaN elements never rank: the finite values win and NaN fills the
// remaining slots of the reduced dimension when there are fewer than
// k finite values.
a := mustFromFloats(t, []float64{1, math.NaN(), 3, 2}, 4)
vals, idxs, err := TopK(a, 3, 0)
if err != nil {
t.Fatal(err)
}
// Three finite values exist, so all three slots are finite, in
// descending order.
wantVals := []float64{3, 2, 1}
wantIdx := []int64{2, 3, 0}
for i := range 3 {
v, _ := FloatAt(vals, i)
if v != wantVals[i] {
t.Errorf("TopKNaN vals [%d]: got %v, want %v", i, v, wantVals[i])
}
j, _ := IntAt(idxs, i)
if j != wantIdx[i] {
t.Errorf("TopKNaN idxs [%d]: got %d, want %d", i, j, wantIdx[i])
}
}
// With more NaN than the slots allow, NaN fills the tail: no finite
// candidate remains, so the slot reports NaN at the fill index 0.
b := mustFromFloats(t, []float64{1, math.NaN(), math.NaN()}, 3)
vb, ib, err := TopK(b, 2, 0)
if err != nil {
t.Fatal(err)
}
if v, _ := FloatAt(vb, 0); v != 1 {
t.Errorf("TopKNaN tail vals[0]: got %v, want 1", v)
}
if v, _ := FloatAt(vb, 1); !math.IsNaN(v) {
t.Errorf("TopKNaN tail vals[1]: got %v, want NaN", v)
}
if j, _ := IntAt(ib, 1); j != 0 {
t.Errorf("TopKNaN tail idxs[1]: got %d, want 0", j)
}
}