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