feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,420 @@
|
||||
// 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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user