Files
tensor/internal/core/mathfunc_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

260 lines
6.9 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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)
}
}
}