Files
tensor/internal/core/ops_test.go
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

421 lines
12 KiB
Go

// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package core
import (
"math"
"slices"
"strings"
"testing"
)
func intsOf(t *testing.T, a *Array) []int64 {
t.Helper()
out := make([]int64, a.Len())
for i := range out {
out[i], _ = IntAt(a, i)
}
return out
}
func floatsOf(t *testing.T, a *Array) []float64 {
t.Helper()
out := make([]float64, a.Len())
for i := range out {
out[i], _ = FloatAt(a, i)
}
return out
}
func TestElementwiseIntStaysInt(t *testing.T) {
a := mustFromInts(t, []int64{1, 2, 3}, 3)
b := mustFromInts(t, []int64{4, 5, 6}, 3)
sum, err := Add(a, b)
if err != nil {
t.Fatalf("Add: %v", err)
}
if sum.Dtype() != Int {
t.Fatalf("int + int must stay int, got %s", sum.Dtype())
}
want := []int64{5, 7, 9}
got := intsOf(t, sum)
for i := range want {
if got[i] != want[i] {
t.Fatalf("Add: %v", got)
}
}
diff, _ := Sub(b, a)
if diff.Dtype() != Int || intsOf(t, diff)[0] != 3 {
t.Fatalf("Sub: %s %v", diff, intsOf(t, diff))
}
prod, _ := Mul(a, b)
if prod.Dtype() != Int || intsOf(t, prod)[2] != 18 {
t.Fatalf("Mul: %s %v", prod, intsOf(t, prod))
}
// int arithmetic wraps like Go's int64.
big := mustFromInts(t, []int64{math.MaxInt64}, 1)
one := mustFromInts(t, []int64{1}, 1)
wrapped, _ := Add(big, one)
if v, _ := IntAt(wrapped, 0); v != math.MinInt64 {
t.Fatalf("wrap: %d", v)
}
}
func TestElementwisePromotesToFloat(t *testing.T) {
i := mustFromInts(t, []int64{1, 2}, 2)
f := mustFromFloats(t, []float64{0.5, 1.5}, 2)
sum, err := Add(i, f)
if err != nil {
t.Fatalf("Add: %v", err)
}
if sum.Dtype() != Float {
t.Fatalf("int + float must promote, got %s", sum.Dtype())
}
got := floatsOf(t, sum)
if got[0] != 1.5 || got[1] != 3.5 {
t.Fatalf("Add promoted: %v", got)
}
}
func TestDivIsTrueDivision(t *testing.T) {
a := mustFromInts(t, []int64{1, 7, -7}, 3)
b := mustFromInts(t, []int64{2, 2, 2}, 3)
q, err := Div(a, b)
if err != nil {
t.Fatalf("Div: %v", err)
}
if q.Dtype() != Float {
t.Fatalf("Div must always yield float, got %s", q.Dtype())
}
got := floatsOf(t, q)
if got[0] != 0.5 || got[1] != 3.5 || got[2] != -3.5 {
t.Fatalf("Div: %v", got)
}
// Zero divisors are IEEE, never errors.
num := mustFromFloats(t, []float64{1.0, -1.0, 0.0}, 3)
zero := mustFromFloats(t, []float64{0.0, 0.0, 0.0}, 3)
zq, err := Div(num, zero)
if err != nil {
t.Fatalf("Div by zero: %v", err)
}
z := floatsOf(t, zq)
if !math.IsInf(z[0], 1) || !math.IsInf(z[1], -1) || !math.IsNaN(z[2]) {
t.Fatalf("IEEE: %v", z)
}
}
func TestQuo(t *testing.T) {
a := mustFromInts(t, []int64{7, -7, 9}, 3)
b := mustFromInts(t, []int64{2, 2, 3}, 3)
q, err := Quo(a, b)
if err != nil {
t.Fatalf("Quo: %v", err)
}
got := intsOf(t, q)
if got[0] != 3 || got[1] != -3 || got[2] != 3 {
t.Fatalf("Quo: %v", got)
}
f := mustFromFloats(t, []float64{1}, 1)
if _, err := Quo(f, f); err == nil || !strings.Contains(err.Error(), "needs int arrays") {
t.Fatalf("Quo float: %v", err)
}
z := mustFromInts(t, []int64{1, 0}, 2)
o := mustFromInts(t, []int64{1, 1}, 2)
if _, err := Quo(z, o); err != nil {
t.Fatalf("Quo zeros in dividend is fine: %v", err)
}
if _, err := Quo(o, z); err == nil || !strings.Contains(err.Error(), "division by zero") {
t.Fatalf("Quo zero divisor: %v", err)
}
}
func TestShapeMismatchIsLoud(t *testing.T) {
a := mustFromInts(t, []int64{1, 2, 3}, 3)
b := mustFromInts(t, []int64{1, 2}, 2)
_, err := Add(a, b)
if err == nil || !strings.Contains(err.Error(), "shape mismatch (3) vs (2)") {
t.Fatalf("shape mismatch: %v", err)
}
if _, err := Div(a, b); err == nil || !strings.Contains(err.Error(), "shape mismatch") {
t.Fatalf("Div shape mismatch: %v", err)
}
if _, err := Quo(a, b); err == nil || !strings.Contains(err.Error(), "shape mismatch") {
t.Fatalf("Quo shape mismatch: %v", err)
}
}
func TestScalarOps(t *testing.T) {
a := mustFromInts(t, []int64{1, 2}, 2)
if v, _ := IntAt(AddI(a, 5), 0); v != 6 {
t.Fatalf("AddI: %d", v)
}
if v, _ := IntAt(SubI(a, 1), 1); v != 1 {
t.Fatalf("SubI: %d", v)
}
if v, _ := IntAt(MulI(a, 3), 1); v != 6 {
t.Fatalf("MulI: %d", v)
}
f := mustFromFloats(t, []float64{1.0, 2.0}, 2)
if v, _ := FloatAt(AddI(f, 5), 0); v != 6.0 {
t.Fatalf("AddI on float: %v", v)
}
if v, _ := FloatAt(AddF(a, 0.5), 0); v != 1.5 {
t.Fatalf("AddF: %v", v)
}
if v, _ := FloatAt(SubF(a, 0.5), 1); v != 1.5 {
t.Fatalf("SubF: %v", v)
}
if v, _ := FloatAt(MulF(a, 1.5), 1); v != 3.0 {
t.Fatalf("MulF: %v", v)
}
// Scalar division is true division regardless of the scalar flavour.
div := DivI(a, 2)
if div.Dtype() != Float || floatsOf(t, div)[0] != 0.5 {
t.Fatalf("DivI: %s %v", div.Dtype(), floatsOf(t, div))
}
if v, _ := FloatAt(DivF(a, 4), 1); v != 0.5 {
t.Fatalf("DivF: %v", floatsOf(t, DivF(a, 4)))
}
q, err := QuoI(a, 2)
if err != nil {
t.Fatalf("QuoI: %v", err)
}
if v, _ := IntAt(q, 0); v != 0 {
t.Fatalf("QuoI: %d", v)
}
if _, err := QuoI(a, 0); err == nil || !strings.Contains(err.Error(), "division by zero") {
t.Fatalf("QuoI zero: %v", err)
}
if _, err := QuoI(f, 2); err == nil || !strings.Contains(err.Error(), "needs an int array") {
t.Fatalf("QuoI float: %v", err)
}
}
// narrowBools builds a bool array for the narrow-dtype contract tests.
func narrowBools(t *testing.T, vals []bool, shape ...int) *Array {
t.Helper()
a, err := FromBools(vals, shape...)
if err != nil {
t.Fatal(err)
}
return a
}
// narrowInt8s builds an int8 array for the narrow-dtype contract tests.
func narrowInt8s(t *testing.T, vals []int8, shape ...int) *Array {
t.Helper()
a, err := FromInt8s(vals, shape...)
if err != nil {
t.Fatal(err)
}
return a
}
// narrowUint8s builds a uint8 array for the narrow-dtype contract tests.
func narrowUint8s(t *testing.T, vals []uint8, shape ...int) *Array {
t.Helper()
a, err := FromUint8s(vals, shape...)
if err != nil {
t.Fatal(err)
}
return a
}
// narrowInt16s builds an int16 array for the narrow-dtype contract tests.
func narrowInt16s(t *testing.T, vals []int16, shape ...int) *Array {
t.Helper()
a, err := FromInt16s(vals, shape...)
if err != nil {
t.Fatal(err)
}
return a
}
// narrowUint32s builds a uint32 array for the narrow-dtype contract tests.
func narrowUint32s(t *testing.T, vals []uint32, shape ...int) *Array {
t.Helper()
a, err := FromUint32s(vals, shape...)
if err != nil {
t.Fatal(err)
}
return a
}
// TestBoolArithmeticIsALoudError pins the contract that bool arrays
// carry no arithmetic: every element-wise arithmetic entry whose
// promoted dtype is bool answers the named refusal, while comparisons
// over bool keep the Int mask.
func TestBoolArithmeticIsALoudError(t *testing.T) {
x := narrowBools(t, []bool{true, true, false}, 3)
y := narrowBools(t, []bool{true, false, false}, 3)
entries := []struct {
name string
fn func() error
}{
{"Add", func() error { _, err := Add(x, y); return err }},
{"Sub", func() error { _, err := Sub(x, y); return err }},
{"Mul", func() error { _, err := Mul(x, y); return err }},
{"Div", func() error { _, err := Div(x, y); return err }},
{"Pow", func() error { _, err := Pow(x, y); return err }},
{"Minimum", func() error { _, err := Minimum(x, y); return err }},
{"Maximum", func() error { _, err := Maximum(x, y); return err }},
{"Dot", func() error { _, err := Dot(x, y); return err }},
}
for _, tc := range entries {
err := tc.fn()
if err == nil || !strings.Contains(err.Error(), "bool arrays have no arithmetic") {
t.Errorf("%s on bool arrays: %v, want the named arithmetic refusal", tc.name, err)
}
}
// Comparisons are not arithmetic: the mask answers bool, true where
// the relation holds.
eq, err := Eq(x, y)
if err != nil {
t.Fatalf("Eq on bool: %v", err)
}
if eq.Dtype() != Bool {
t.Fatalf("Eq on bool answered dtype %s, want the bool mask", eq.Dtype())
}
if want := []bool{true, false, true}; !slices.Equal(eq.RawBools(), want) {
t.Fatalf("Eq on bool = %v, want %v", eq.RawBools(), want)
}
// A mixed bool pair promotes into the other operand's dtype, where
// arithmetic is defined on the widened values.
i8 := narrowInt8s(t, []int8{2, 3, 4}, 3)
sum, err := Add(x, i8)
if err != nil {
t.Fatalf("Add bool with int8: %v", err)
}
if sum.Dtype() != Int8 {
t.Fatalf("Add bool with int8 answered %s, want int8", sum.Dtype())
}
if want := []int8{3, 4, 4}; !slices.Equal(sum.RawInt8s(), want) {
t.Fatalf("Add bool with int8 = %v, want %v", sum.RawInt8s(), want)
}
}
// TestNarrowIntegerElementwiseContract pins the narrow element-type
// loops: same-width dense pairs wrap natively in their own type, mixed
// signedness promotes through the containment table, integer true
// division answers float64 exactly as the int path does, and the scalar
// maps keep their kind with the implicit-store cast.
func TestNarrowIntegerElementwiseContract(t *testing.T) {
a8 := narrowInt8s(t, []int8{100, -128, 5}, 3)
b8 := narrowInt8s(t, []int8{100, 2, 3}, 3)
s, err := Add(a8, b8)
if err != nil {
t.Fatalf("Add int8: %v", err)
}
if s.Dtype() != Int8 {
t.Fatalf("Add int8 answered %s, want int8", s.Dtype())
}
if want := []int8{-56, -126, 8}; !slices.Equal(s.RawInt8s(), want) {
t.Fatalf("Add int8 wrap = %v, want %v", s.RawInt8s(), want)
}
m, err := Mul(a8, b8)
if err != nil {
t.Fatalf("Mul int8: %v", err)
}
if want := []int8{16, 0, 15}; !slices.Equal(m.RawInt8s(), want) {
t.Fatalf("Mul int8 wrap = %v, want %v", m.RawInt8s(), want)
}
mx, err := Maximum(a8, b8)
if err != nil {
t.Fatalf("Maximum int8: %v", err)
}
if want := []int8{100, 2, 5}; !slices.Equal(mx.RawInt8s(), want) {
t.Fatalf("Maximum int8 = %v, want %v", mx.RawInt8s(), want)
}
// Mixed signedness: int8 with uint8 promotes to int16, where both
// value sets fit.
u8 := narrowUint8s(t, []uint8{200, 200, 200}, 3)
mix, err := Add(a8, u8)
if err != nil {
t.Fatalf("Add int8 with uint8: %v", err)
}
if mix.Dtype() != Int16 {
t.Fatalf("Add int8 with uint8 answered %s, want int16", mix.Dtype())
}
if want := []int16{300, 72, 205}; !slices.Equal(mix.RawInt16s(), want) {
t.Fatalf("Add int8 with uint8 = %v, want %v", mix.RawInt16s(), want)
}
// True division: an integer-class pair answers float64, the route
// the int pair has always taken.
d, err := Div(a8, b8)
if err != nil {
t.Fatalf("Div int8: %v", err)
}
if d.Dtype() != Float {
t.Fatalf("Div int8 answered %s, want float", d.Dtype())
}
got := d.RawFloats()
if got[0] != 1 || got[1] != -64 || math.Abs(got[2]-5.0/3.0) > 1e-12 {
t.Fatalf("Div int8 = %v, want [1 -64 1.666...]", got)
}
// Pow keeps the int contract's negative-exponent refusal on the
// narrow widths.
neg := narrowInt8s(t, []int8{-1}, 1)
pos := narrowInt8s(t, []int8{2}, 1)
if _, err := Pow(pos, neg); err == nil || !strings.Contains(err.Error(), "negative exponent") {
t.Fatalf("Pow int8 with a negative exponent: %v", err)
}
// The mixed pair whose promote() is Int: the exponent scan must
// refuse the negative narrow exponent instead of letting powInt's
// loop answer a silent 1.
pbase, perr := FromInts([]int64{2, 3}, 2)
if perr != nil {
t.Fatalf("FromInts: %v", perr)
}
if _, err := Pow(pbase, narrowInt8s(t, []int8{-3, 4}, 2)); err == nil ||
!strings.Contains(err.Error(), "negative exponent") {
t.Fatalf("Pow int with a negative int8 exponent: %v", err)
}
if _, err := Pow(narrowUint32s(t, []uint32{2, 3}, 2), narrowInt16s(t, []int16{-1, 2}, 2)); err == nil ||
!strings.Contains(err.Error(), "negative exponent") {
t.Fatalf("Pow uint32 with a negative int16 exponent: %v", err)
}
// The scalar maps: AddI keeps the kind with the implicit-store cast,
// AddF widens the whole integer class to float64, and AddI on a bool
// array has no error channel, so it answers nil.
si := AddI(a8, 200)
if si == nil || si.Dtype() != Int8 {
t.Fatalf("AddI int8: %v %v", si, err)
}
// The int64 sum narrows on store: int8(300) = 44, -128+200 = 72,
// int8(205) = -51.
if want := []int8{44, 72, -51}; !slices.Equal(si.RawInt8s(), want) {
t.Fatalf("AddI int8 = %v, want %v", si.RawInt8s(), want)
}
sf := AddF(a8, 0.5)
if sf.Dtype() != Float {
t.Fatalf("AddF int8 answered %s, want float", sf.Dtype())
}
if want := []float64{100.5, -127.5, 5.5}; !slices.Equal(sf.RawFloats(), want) {
t.Fatalf("AddF int8 = %v, want %v", sf.RawFloats(), want)
}
bl := narrowBools(t, []bool{true, false}, 2)
if got := AddI(bl, 1); got != nil {
t.Fatalf("AddI on a bool array = %v, want the nil refusal", got)
}
}