331 lines
9.9 KiB
Go
331 lines
9.9 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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)
|
|
}
|
|
}
|