// Copyright (c) 2026 Petr BalvĂ­n (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) } }