// Copyright (c) 2026 Petr Balvín (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) } } }