Files

260 lines
6.9 KiB
Go
Raw Permalink Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package core
import (
"math"
"strings"
"testing"
)
func TestMathFuncs(t *testing.T) {
f := mustFromFloats(t, []float64{1.0, 4.0}, 2)
sqrt := mustOk(Sqrt(f))
if v, _ := FloatAt(sqrt, 0); v != 1.0 {
t.Fatalf("Sqrt: %v", v)
}
log2 := mustOk(Log2(f))
if v, _ := FloatAt(log2, 1); v != 2.0 {
t.Fatalf("Log2: %v", v)
}
exp := mustOk(Exp(mustFromFloats(t, []float64{0}, 1)))
if v, _ := FloatAt(exp, 0); math.Abs(v-1) > 1e-12 {
t.Fatalf("Exp(0) must be 1: %v", v)
}
log := mustOk(Log(mustFromFloats(t, []float64{1}, 1)))
if v, _ := FloatAt(log, 0); v != 0 {
t.Fatalf("Log(1): %v", v)
}
log10 := mustOk(Log10(mustFromFloats(t, []float64{10}, 1)))
if v, _ := FloatAt(log10, 0); v != 1 {
t.Fatalf("Log10(10): %v", v)
}
sin := mustOk(Sin(mustFromFloats(t, []float64{0}, 1)))
if v, _ := FloatAt(sin, 0); v != 0 {
t.Fatalf("Sin(0): %v", v)
}
cos := mustOk(Cos(mustFromFloats(t, []float64{0}, 1)))
if v, _ := FloatAt(cos, 0); v != 1 {
t.Fatalf("Cos(0): %v", v)
}
tan := mustOk(Tan(mustFromFloats(t, []float64{0}, 1)))
if v, _ := FloatAt(tan, 0); v != 0 {
t.Fatalf("Tan(0): %v", v)
}
// IEEE domain rules: negatives yield NaN, never an error.
neg := mustFromFloats(t, []float64{-1}, 1)
sq, err := Sqrt(neg)
if err != nil {
t.Fatalf("Sqrt(-1): %v", err)
}
if v, _ := FloatAt(sq, 0); !math.IsNaN(v) {
t.Fatalf("Sqrt(-1): %v", v)
}
lg, _ := Log(neg)
if v, _ := FloatAt(lg, 0); !math.IsNaN(v) {
t.Fatalf("Log(-1): %v", v)
}
// int arrays convert for the transcendental family.
i := mustFromInts(t, []int64{4}, 1)
isqrtArr := mustOk(Sqrt(i))
isqrt, _ := FloatAt(isqrtArr, 0)
if isqrt != 2 || isqrtArr.Dtype() != Float {
t.Fatalf("Sqrt on int: %s", isqrtArr)
}
// complex arrays error for the real family.
c := mustFromComplexes(t, []complex128{1}, 1)
if _, err := Exp(c); err == nil || !strings.Contains(err.Error(), "not supported") {
t.Fatalf("Exp complex: %v", err)
}
if _, err := Sin(c); err == nil {
t.Fatalf("Sin complex must error")
}
if _, err := Floor(c); err == nil || !strings.Contains(err.Error(), "no rounding") {
t.Fatalf("Floor complex: %v", err)
}
}
func TestAbs(t *testing.T) {
i := mustFromInts(t, []int64{-5, 3}, 2)
ai := Abs(i)
if ai.Dtype() != Int {
t.Fatalf("Abs int dtype: %s", ai.Dtype())
}
if v, _ := IntAt(ai, 0); v != 5 {
t.Fatalf("Abs int: %d", v)
}
f := mustFromFloats(t, []float64{-2.5}, 1)
if v, _ := FloatAt(Abs(f), 0); v != 2.5 {
t.Fatalf("Abs float: %v", v)
}
// Complex magnitude is real.
c := mustFromComplexes(t, []complex128{complex(3, 4)}, 1)
ac := Abs(c)
if ac.Dtype() != Float {
t.Fatalf("Abs complex dtype: %s", ac.Dtype())
}
if v, _ := FloatAt(ac, 0); v != 5 {
t.Fatalf("Abs complex: %v", v)
}
}
func TestRounding(t *testing.T) {
f := mustFromFloats(t, []float64{2.7, -2.7}, 2)
floor := mustOk(Floor(f))
if v, _ := FloatAt(floor, 0); v != 2 {
t.Fatalf("Floor: %v", v)
}
if v, _ := FloatAt(floor, 1); v != -3 {
t.Fatalf("Floor neg: %v", v)
}
ceil := mustOk(Ceil(f))
if v, _ := FloatAt(ceil, 0); v != 3 {
t.Fatalf("Ceil: %v", v)
}
round := mustOk(Round(f))
if v, _ := FloatAt(round, 0); v != 3 {
t.Fatalf("Round: %v", v)
}
if v, _ := FloatAt(round, 1); v != -3 {
t.Fatalf("Round half away: %v", v)
}
trunc := mustOk(Trunc(f))
if v, _ := FloatAt(trunc, 0); v != 2 {
t.Fatalf("Trunc: %v", v)
}
// Int arrays are identity copies.
i := mustFromInts(t, []int64{7}, 1)
ir, err := Floor(i)
if err != nil || ir.Dtype() != Int {
t.Fatalf("Floor on int: %s %v", ir, err)
}
if v, _ := IntAt(ir, 0); v != 7 {
t.Fatalf("Floor int identity: %d", v)
}
}
func TestPow(t *testing.T) {
a := mustFromInts(t, []int64{2, 3}, 2)
e := mustFromInts(t, []int64{10, 2}, 2)
p, err := Pow(a, e)
if err != nil {
t.Fatalf("Pow: %v", err)
}
if p.Dtype() != Int {
t.Fatalf("Pow dtype: %s", p.Dtype())
}
if v, _ := IntAt(p, 0); v != 1024 {
t.Fatalf("Pow 2^10: %d", v)
}
if v, _ := IntAt(p, 1); v != 9 {
t.Fatalf("Pow 3^2: %d", v)
}
// Float promotion.
pf, err := Pow(mustFromFloats(t, []float64{4.0, 9.0}, 2), e)
if err != nil {
t.Fatalf("Pow promote: %v", err)
}
if pf.Dtype() != Float {
t.Fatalf("Pow promote dtype: %s", pf.Dtype())
}
if v, _ := FloatAt(pf, 0); v != 1048576 {
t.Fatalf("Pow 4^10: %v", v)
}
if v, _ := FloatAt(pf, 1); v != 81 {
t.Fatalf("Pow 9^2: %v", v)
}
// Scalar exponent.
pi, err := PowI(a, 3)
if err != nil {
t.Fatalf("PowI: %v", err)
}
if v, _ := IntAt(pi, 0); v != 8 {
t.Fatalf("PowI 2^3: %d", v)
}
pif, err := PowI(mustFromFloats(t, []float64{9.0}, 1), 2)
if err != nil {
t.Fatalf("PowI float: %v", err)
}
if v, _ := FloatAt(pif, 0); v != 81 {
t.Fatalf("PowI float value: %v", v)
}
// Negative int exponent on an int array errors; on float it works.
if _, err := PowI(a, -1); err == nil || !strings.Contains(err.Error(), "negative exponent") {
t.Fatalf("PowI negative: %v", err)
}
if _, err := PowI(mustFromFloats(t, []float64{2}, 1), -2); err != nil {
t.Fatalf("PowI negative on float: %v", err)
}
// PowI carries no narrow kernel: the refusal names
// the dtype and the conversion the contract asks for.
if _, err := PowI(narrowInt8s(t, []int8{2, 3}, 2), 2); err == nil ||
!strings.Contains(err.Error(), "convert with Astype") ||
!strings.Contains(err.Error(), "int8") {
t.Fatalf("PowI int8: %v", err)
}
// Complex powers use exact repeated squaring on complex128.
c := mustFromComplexes(t, []complex128{2 + 3i}, 1)
if _, err := Pow(c, c); err != nil {
t.Fatalf("Pow complex: %v", err)
}
squared := mustOk(PowI(c, 2))
want := (2 + 3i) * (2 + 3i) // −5+12i
if squared.RawComplexes()[0] != want {
t.Fatalf("PowI complex = %v, want %v", squared.RawComplexes()[0], want)
}
}
// mustOk unwraps a call that must succeed; a panic here fails the test.
func mustOk(a *Array, err error) *Array {
if err != nil {
panic(err)
}
return a
}
// TestPowDenseFloatOperands pins the dense float base against a dense
// float exponent, the pair whose kernel holds both payloads as slices;
// an int exponent and a promoted pair take the closure walk instead.
func TestPowDenseFloatOperands(t *testing.T) {
base := mustFromFloats(t, []float64{4, 9, 2, 2}, 2, 2)
exp := mustFromFloats(t, []float64{10, 2, 3, -1}, 2, 2)
got, err := Pow(base, exp)
if err != nil {
t.Fatalf("Pow float: %v", err)
}
if got.Dtype() != Float {
t.Fatalf("Pow float dtype: %s", got.Dtype())
}
for i, want := range []float64{1048576, 81, 8, 0.5} {
if v := got.RawFloats()[i]; v != want {
t.Fatalf("Pow float [%d] = %v, want %v", i, v, want)
}
}
// The float32 pair takes the same kernel one width down.
b32 := mustFromFloat32s(t, []float32{4, 2}, 2)
e32 := mustFromFloat32s(t, []float32{2, -1}, 2)
p32, err := Pow(b32, e32)
if err != nil {
t.Fatalf("Pow float32: %v", err)
}
for i, want := range []float32{16, 0.5} {
if v := p32.RawFloat32s()[i]; v != want {
t.Fatalf("Pow float32 [%d] = %v, want %v", i, v, want)
}
}
}