Files
tensor/internal/core/reduce_test.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

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)
}
}