// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package core import ( "math" "strings" "testing" ) func TestSum(t *testing.T) { i := mustFromInts(t, []int64{1, 2, 3}, 3) s := Sum(i) if s.IsFloat() || s.Int() != 6 { t.Fatalf("Sum int: %s", s) } f := mustFromFloats(t, []float64{0.5, 1.5}, 2) fs := Sum(f) if !fs.IsFloat() || fs.Float() != 2.0 { t.Fatalf("Sum float: %s", fs) } empty := mustFromInts(t, nil, 0) if Sum(empty).Int() != 0 { t.Fatalf("Sum of empty must be zero") } } func TestMinMeanMax(t *testing.T) { a := mustFromInts(t, []int64{3, 1, 2}, 3) mn, err := Min(a) if err != nil || mn.Int() != 1 { t.Fatalf("Min: %s %v", mn, err) } mx, err := Max(a) if err != nil || mx.Int() != 3 { t.Fatalf("Max: %s %v", mx, err) } f := mustFromFloats(t, []float64{2.5, -1.5}, 2) fmn, _ := Min(f) if !fmn.IsFloat() || fmn.Float() != -1.5 { t.Fatalf("Min float: %s", fmn) } mean, err := Mean(a) if err != nil || mean != 2.0 { t.Fatalf("Mean int: %v %v", mean, err) } fm, err := Mean(f) if err != nil || fm != 0.5 { t.Fatalf("Mean float: %v %v", fm, err) } empty := mustFromInts(t, nil, 0) if _, err := Min(empty); err == nil || !strings.Contains(err.Error(), "empty array") { t.Fatalf("Min empty: %v", err) } if _, err := Max(empty); err == nil || !strings.Contains(err.Error(), "empty array") { t.Fatalf("Max empty: %v", err) } if _, err := Mean(empty); err == nil || !strings.Contains(err.Error(), "empty array") { t.Fatalf("Mean empty: %v", err) } } // TestProdEmptyAxisIsUnitProduct pins the empty product along a reduced // axis: a line of zero elements multiplies to the multiplicative // identity in the array's own dtype, the same value the global // one-dimensional product answers for an empty array. The zeroed // allocation must never surface as the answer. func TestProdEmptyAxisIsUnitProduct(t *testing.T) { f := mustFromFloats(t, nil, 2, 0) got, err := Prod(f, 1, false) if err != nil { t.Fatalf("Prod float over an empty axis: %v", err) } if got.Dtype() != Float || got.Len() != 2 { t.Fatalf("Prod float over an empty axis: dtype %s shape %v", got.Dtype(), got.Shape()) } for i := range 2 { if v := got.RawFloats()[i]; v != 1 { t.Errorf("Prod float over an empty axis [%d] = %v, want 1", i, v) } } i := mustFromInts(t, nil, 2, 0) gotI, err := Prod(i, 1, false) if err != nil { t.Fatalf("Prod int over an empty axis: %v", err) } if gotI.Dtype() != Int { t.Fatalf("Prod int over an empty axis: dtype %s", gotI.Dtype()) } for k := range 2 { if v := gotI.RawInts()[k]; v != 1 { t.Errorf("Prod int over an empty axis [%d] = %v, want 1", k, v) } } h := mustFromFloat16s(t, nil, 2, 0) gotH, err := Prod(h, 1, false) if err != nil { t.Fatalf("Prod float16 over an empty axis: %v", err) } for k := range 2 { if bits := gotH.RawHalves()[k]; bits != halfOne { t.Errorf("Prod float16 over an empty axis [%d] = %#04x, want the half one %#04x", k, bits, halfOne) } } // keepDim keeps the unit answer under the reinserted size-1 axis. kept, err := Prod(f, 1, true) if err != nil { t.Fatalf("Prod keepDim over an empty axis: %v", err) } if kept.Shape()[0] != 2 || kept.Shape()[1] != 1 || kept.RawFloats()[1] != 1 { t.Fatalf("Prod keepDim over an empty axis: shape %v values %v", kept.Shape(), kept.RawFloats()) } // The global one-dimensional empty product agrees: both routes to a // product of zero factors answer the unit. global, err := Prod(mustFromFloats(t, nil, 0), 0, false) if err != nil { t.Fatalf("Prod of the empty array: %v", err) } if v := global.RawFloats()[0]; v != 1 { t.Errorf("Prod of the empty array = %v, want 1", v) } } func TestScalarBox(t *testing.T) { i := Scalar{i: 7} if i.IsFloat() || i.Int() != 7 || i.Float() != 7.0 || i.String() != "int 7" { t.Fatalf("int scalar: %s", i) } f := Scalar{isFloat: true, f: 2.5} if !f.IsFloat() || f.Float() != 2.5 || f.Int() != 2 || f.String() != "float 2.5" { t.Fatalf("float scalar: %s", f) } } func TestDot(t *testing.T) { a := mustFromInts(t, []int64{1, 2, 3}, 3) b := mustFromInts(t, []int64{4, 5, 6}, 3) d, err := Dot(a, b) if err != nil { t.Fatalf("Dot: %v", err) } if d.IsFloat() || d.Int() != 32 { t.Fatalf("Dot int: %s", d) } f := mustFromFloats(t, []float64{0.5, 0.5}, 2) fd, err := Dot(f, f) if err != nil || !fd.IsFloat() || fd.Float() != 0.5 { t.Fatalf("Dot float: %s %v", fd, err) } // Mixed dtypes promote to float. md, err := Dot(a, mustFromFloats(t, []float64{1, 1, 1}, 3)) if err != nil || !md.IsFloat() || md.Float() != 6.0 { t.Fatalf("Dot mixed: %s %v", md, err) } if _, err := Dot(a, mustFromInts(t, []int64{1}, 1)); err == nil || !strings.Contains(err.Error(), "length mismatch") { t.Fatalf("Dot length: %v", err) } m2 := mustFromInts(t, []int64{1, 2, 3, 4}, 2, 2) if _, err := Dot(m2, m2); err == nil || !strings.Contains(err.Error(), "needs 1-D arrays") { t.Fatalf("Dot 2-D: %v", err) } } // TestNormIntArms pins the magnitude walks over an int payload, whose // values are taken in int64 and widened exactly: p = 1 adds the // magnitudes, p = 2 squares them, and p = Inf takes the largest. func TestNormIntArms(t *testing.T) { a := mustFromInts(t, []int64{3, -4}, 2) for _, tc := range []struct { p float64 want float64 }{ {1, 7}, // |3| + |-4| {2, 5}, // √(3² + 4²), exact {math.Inf(1), 4}, // the largest magnitude } { got, err := Norm(a, tc.p, 0, false) if err != nil { t.Fatalf("Norm int p=%v: %v", tc.p, err) } if v := got.RawFloats()[0]; v != tc.want { t.Errorf("Norm int p=%v over [3 -4]: %v, want %v", tc.p, v, tc.want) } } // A second sample keeps the sum of squares off one Pythagorean // triple: 2² + 3² + 6² = 49. b := mustFromInts(t, []int64{2, -3, 6}, 3) l2, err := Norm(b, 2, 0, false) if err != nil { t.Fatalf("Norm int p=2: %v", err) } if v := l2.RawFloats()[0]; v != 7 { t.Errorf("Norm int p=2 over [2 -3 6]: %v, want 7", v) } } // TestNormFloat32Arms pins the same walks over a float32 payload: every // value widens to float64, and a negative component must contribute its // magnitude rather than its sign. func TestNormFloat32Arms(t *testing.T) { a := mustFromFloat32s(t, []float32{-3, 4}, 2) for _, tc := range []struct { p float64 want float64 }{ {1, 7}, {2, 5}, {math.Inf(1), 4}, } { got, err := Norm(a, tc.p, 0, false) if err != nil { t.Fatalf("Norm float32 p=%v: %v", tc.p, err) } if v := got.RawFloats()[0]; v != tc.want { t.Errorf("Norm float32 p=%v over [-3 4]: %v, want %v", tc.p, v, tc.want) } } b := mustFromFloat32s(t, []float32{-1.5, 2.5, -4}, 3) l1, err := Norm(b, 1, 0, false) if err != nil { t.Fatalf("Norm float32 p=1: %v", err) } if v := l1.RawFloats()[0]; v != 8 { t.Errorf("Norm float32 p=1 over [-1.5 2.5 -4]: %v, want 8", v) } } // TestMeanOverIntSample pins the mean of an int sample through the // magnitude pre-pass. The scaled branch cannot engage for an int // payload: maxAbs is at most 2^63 while the guard asks for a magnitude // above MaxFloat64/n, which no element count reaches, so the values // below are the plain sum divided by the count. func TestMeanOverIntSample(t *testing.T) { a := mustFromInts(t, []int64{-8, -4, 4, 8}, 4) got, err := Mean(a) if err != nil { t.Fatalf("Mean int: %v", err) } if got != 0 { t.Errorf("Mean int over [-8 -4 4 8]: %v, want 0", got) } // Both ends of the int64 range: the magnitudes are taken as float64 // (exact), the sum wraps to -1 as the int64 fold does, and the mean // is exact in float64. b := mustFromInts(t, []int64{math.MinInt64, math.MaxInt64}, 2) got, err = Mean(b) if err != nil { t.Fatalf("Mean int range: %v", err) } if got != -0.5 { t.Errorf("Mean int over [MinInt64 MaxInt64]: %v, want -0.5", got) } } // TestNarrowReductionsContract pins the integer-class reduction rules: // a bool sum counts its true elements, narrow sums widen exactly into an // Int scalar, Min/Max compare natively in each payload type and answer // Int scalars, Mean runs the float64 contract over every non-complex // dtype, and Dot refuses the bool pair by name. func TestNarrowReductionsContract(t *testing.T) { bl := narrowBools(t, []bool{true, false, true, true}, 4) u8 := narrowUint8s(t, []uint8{200, 200, 200}, 3) i8 := narrowInt8s(t, []int8{-100, 50}, 2) s := Sum(bl) if s.IsFloat() || s.IsComplex() || s.Int() != 3 { t.Fatalf("Sum bool = %v, want the Int scalar 3", s) } if s := Sum(u8); s.Int() != 600 { t.Fatalf("Sum uint8 = %v, want 600", s) } if s := Sum(i8); s.Int() != -50 { t.Fatalf("Sum int8 = %v, want -50", s) } mn, err := Min(bl) if err != nil || mn.Int() != 0 { t.Fatalf("Min bool = %v %v, want the Int scalar 0", mn, err) } mx, err := Max(bl) if err != nil || mx.Int() != 1 { t.Fatalf("Max bool = %v %v, want the Int scalar 1", mx, err) } if mn, err = Min(u8); err != nil || mn.Int() != 200 { t.Fatalf("Min uint8 = %v %v, want 200", mn, err) } if mx, err = Max(i8); err != nil || mx.Int() != 50 { t.Fatalf("Max int8 = %v %v, want 50", mx, err) } // Native comparison per payload type: the top of the uint32 range // keeps its exact int64 image. u32 := narrowUint32s(t, []uint32{math.MaxUint32, math.MaxUint32 - 1}, 2) if mx, err = Max(u32); err != nil || mx.Int() != int64(math.MaxUint32) { t.Fatalf("Max uint32 = %v %v, want %d", mx, err, uint64(math.MaxUint32)) } // Mean runs the float64 contract over every non-complex dtype. m, err := Mean(i8) if err != nil || m != -25 { t.Fatalf("Mean int8 = %v %v, want -25", m, err) } if m, err = Mean(bl); err != nil || m != 0.75 { t.Fatalf("Mean bool = %v %v, want 0.75", m, err) } // Dot: an integer-class pair answers an Int scalar accumulated in // int64; a bool pair is the arithmetic refusal. d, err := Dot(i8, i8) if err != nil || d.IsFloat() || d.Int() != 12500 { t.Fatalf("Dot int8 = %v %v, want the Int scalar 12500", d, err) } if _, err := Dot(bl, bl); err == nil || !strings.Contains(err.Error(), "bool arrays have no arithmetic") { t.Fatalf("Dot bool: %v", err) } }