216 lines
6.0 KiB
Go
216 lines
6.0 KiB
Go
// 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)
|
||
|
|
}
|
||
|
|
}
|