// Copyright (c) 2026 Petr Balvín (https://petrbalvin.org) // SPDX-License-Identifier: MIT package core import "math" // Element-wise math functions. The transcendental family is // real-only for now: int arrays convert to float, complex arrays error // loudly. IEEE domain rules apply: log and sqrt of negatives yield NaN, // never an error. Rounding keeps each dtype: int arrays are identity // copies, float arrays map to float. // Abs returns the element-wise absolute value; a complex array yields its // float magnitudes. Bool and the narrow integer widths are refused: they // carry no kernel here, and this no-error entry answers nil, // the refusal shape the no-error constructors carry; the caller converts // with Astype first. func Abs(a *Array) *Array { if narrowRefused(a.dt) { return nil } out := &Array{shape: a.Shape()} switch a.dt { case Int: out.dt = Int out.ints = make([]int64, a.Len()) parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { as, os := a.ints[s:e], out.ints[s:e] for i := range os { v := as[i] if v < 0 { v = -v // wraps at the minimum value } os[i] = v } }) case Float16: out.dt = Float16 out.halves = make([]uint16, a.Len()) parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { as, os := a.halves[s:e], out.halves[s:e] for i := range os { // Abs of a half is exact: widen, clear the sign in // float64, narrow. os[i] = HalfFromFloat64(math.Abs(HalfToFloat64(as[i]))) } }) case Float32: out.dt = Float32 out.floats32 = make([]float32, a.Len()) parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { as, os := a.floats32[s:e], out.floats32[s:e] for i := range os { os[i] = float32(math.Abs(float64(as[i]))) } }) case Float: out.dt = Float out.floats = make([]float64, a.Len()) parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) { as, os := a.floats[s:e], out.floats[s:e] for i := range os { os[i] = math.Abs(as[i]) } }) default: // The magnitude of a complex element costs a math.Hypot, so this // walk amortises a spawn at the math-floor chunk, not the // arithmetic floor. out.dt = Float out.floats = make([]float64, a.Len()) parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { as, os := a.complexes[s:e], out.floats[s:e] for i := range os { os[i] = math.Hypot(real(as[i]), imag(as[i])) } }) } return out } // Exp returns e raised to each element. func Exp(a *Array) (*Array, error) { return a.realFunc("Exp", math.Exp) } // Log returns the natural logarithm of each element; negatives yield NaN. func Log(a *Array) (*Array, error) { return a.realFunc("Log", math.Log) } // Log2 returns the base-2 logarithm of each element; negatives yield NaN. func Log2(a *Array) (*Array, error) { return a.realFunc("Log2", math.Log2) } // Log10 returns the base-10 logarithm of each element; negatives yield // NaN. func Log10(a *Array) (*Array, error) { return a.realFunc("Log10", math.Log10) } // Sqrt returns the square root of each element; negatives yield NaN. // Specialised per dtype: math.Sqrt called directly lowers to the // SQRTSD intrinsic with no call, while a loop through realFunc's func // value pays an indirect call per element that dwarfs the operation // itself. The per-dtype loops read through realFunc's accessors // element for element, so the outputs, the error and the dtype // promotion are identical. func Sqrt(a *Array) (*Array, error) { if a.dt == Complex { return nil, errf("Sqrt: complex arrays are not supported") } if narrowRefused(a.dt) { return nil, errf("Sqrt: dtype %s is not supported; convert with Astype", a.dt) } out := &Array{shape: a.Shape(), dt: a.dt} if a.dt == Int { out.dt = Float } out.alloc(a.Len()) switch a.dt { case Int: parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { for i := s; i < e; i++ { out.floats[i] = math.Sqrt(a.floatAt(i)) } }) case Float16: // float16 keeps float16, computed in float64 and narrowed once: // the widening is exact, so the narrowed result is the correctly // rounded square root. parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { for i := s; i < e; i++ { out.halves[i] = HalfFromFloat64(math.Sqrt(HalfToFloat64(a.halves[i]))) } }) case Float32: // float32 keeps float32, computed in float64 and rounded once //: the widening is exact, so the narrowed result is the // correctly rounded square root. parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { for i := s; i < e; i++ { out.floats32[i] = float32(math.Sqrt(float64(a.floats32[i]))) } }) default: parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { os := out.floats[s:e] as := a.floats[s:e] for i := range os { os[i] = math.Sqrt(as[i]) } }) } return out, nil } // Sin returns the sine of each element (radians). func Sin(a *Array) (*Array, error) { return a.realFunc("Sin", math.Sin) } // Cos returns the cosine of each element (radians). func Cos(a *Array) (*Array, error) { return a.realFunc("Cos", math.Cos) } // Tan returns the tangent of each element (radians). func Tan(a *Array) (*Array, error) { return a.realFunc("Tan", math.Tan) } // Floor rounds each element down; an int array is an identity copy. func Floor(a *Array) (*Array, error) { return a.roundFunc("Floor", math.Floor) } // Ceil rounds each element up; an int array is an identity copy. func Ceil(a *Array) (*Array, error) { return a.roundFunc("Ceil", math.Ceil) } // Round rounds each element half away from zero; an int array is an // identity copy. func Round(a *Array) (*Array, error) { return a.roundFunc("Round", math.Round) } // Trunc rounds each element toward zero; an int array is an identity copy. func Trunc(a *Array) (*Array, error) { return a.roundFunc("Trunc", math.Trunc) } // Pow returns the element-wise power; integer-class operands promote // to an integer dtype and stay integer there (wrapping on overflow, // and a negative exponent is refused at every width: the reciprocal // has no integer result), any float or complex operand promotes per // the ladder (int operands convert), and negative bases with // fractional exponents follow math.Pow's IEEE NaN. func Pow(a, b *Array) (*Array, error) { if intClass(a.dt) && intClass(b.dt) { // An integer-class pair computes in integer space whatever // promote answers, and a negative exponent has no integer // result at any width. The refusal scans the exponent operand // through intAt, whose widening of the whole class is exact, // so no mixed pair (an Int base with a narrow exponent, a // Uint32 base with a negative narrow exponent) slips past the // gate into powInt, whose loop answers 1 for a negative // exponent. for i := range b.Len() { if e := b.intAt(i); e < 0 { return nil, errf("Pow: negative exponent %d at element %d has no int result", e, i) } } } return elementwiseOp(a, b, "Pow", pairPow) } // mathPowComplex raises a complex base to a complex exponent via // exp(y·log(x)): the principal branch. A zero base with a non-positive // exponent follows IEEE NaN/Inf conventions. func mathPowComplex(x, y complex128) complex128 { if x == 0 { if imag(y) == 0 && real(y) > 0 { return 0 } return complex(math.NaN(), math.NaN()) } return cexp(cmul(y, clog(x))) } // cexp, clog and cmul are the complex transcendals needed by Pow, // implemented here to keep the library dependency-free. func cexp(z complex128) complex128 { e := math.Exp(real(z)) return complex(e*math.Cos(imag(z)), e*math.Sin(imag(z))) } func clog(z complex128) complex128 { return complex(math.Log(math.Hypot(real(z), imag(z))), math.Atan2(imag(z), real(z))) } func cmul(a, b complex128) complex128 { return complex(real(a)*real(b)-imag(a)*imag(b), real(a)*imag(b)+imag(a)*real(b)) } // PowI raises each element to an int exponent; int arrays stay int with // wrapping on overflow (a negative exponent is an error), float32 arrays // keep float32 computed in float64, float64 arrays stay float64, // complex arrays use exact repeated squaring (a negative exponent takes // the reciprocal). func PowI(a *Array, n int64) (*Array, error) { switch a.dt { case Int: if n < 0 { return nil, errf("PowI: a negative exponent has no int result") } out := &Array{shape: a.Shape(), dt: Int, ints: make([]int64, a.Len())} // Bounded by Len, not by the payload: a rebased view carries a // payload longer than its extent. parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { as, os := a.ints[s:e], out.ints[s:e] for i := range os { os[i] = powInt(as[i], n) } }) return out, nil case Float16: out := &Array{shape: a.Shape(), dt: Float16, halves: make([]uint16, a.Len())} parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { as, os := a.halves[s:e], out.halves[s:e] for i := range os { os[i] = HalfFromFloat64(math.Pow(HalfToFloat64(as[i]), float64(n))) } }) return out, nil case Float32: out := &Array{shape: a.Shape(), dt: Float32, floats32: make([]float32, a.Len())} parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { as, os := a.floats32[s:e], out.floats32[s:e] for i := range os { os[i] = float32(math.Pow(float64(as[i]), float64(n))) } }) return out, nil case Float: return a.realFunc("PowI", func(x float64) float64 { return math.Pow(x, float64(n)) }) case Complex: if n < 0 { // A negative power is the reciprocal of the positive one. out := &Array{shape: a.Shape(), dt: Complex, complexes: make([]complex128, a.Len())} parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { as, os := a.complexes[s:e], out.complexes[s:e] for i := range os { p := powComplexI(as[i], -n) os[i] = complex(1, 0) / p } }) return out, nil } out := &Array{shape: a.Shape(), dt: Complex, complexes: make([]complex128, a.Len())} parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { as, os := a.complexes[s:e], out.complexes[s:e] for i := range os { os[i] = powComplexI(as[i], n) } }) return out, nil } return nil, errf("PowI: dtype %s is not supported; convert with Astype", a.dt) } // powComplexI raises z to a non-negative integer power by repeated // squaring: exact integer arithmetic in complex128. func powComplexI(z complex128, n int64) complex128 { result := complex(1, 0) base := z for n > 0 { if n&1 == 1 { result *= base } base *= base n >>= 1 } return result } // powInt raises x to a non-negative y by repeated squaring, wrapping on // overflow like every int operation. func powInt(x, y int64) int64 { result := int64(1) base := x for y > 0 { if y&1 == 1 { result *= base } base *= base y >>= 1 } return result } // realFunc applies a real function element-wise: int arrays convert to // float64, float16 and float32 arrays keep their dtype (computed in // float64, narrowed once), complex arrays error. func (a *Array) realFunc(name string, f func(float64) float64) (*Array, error) { if a.dt == Complex { return nil, errf("%s: complex arrays are not supported", name) } if narrowRefused(a.dt) { // The transcendental family widens integer arrays to float64; // bool and the narrow widths carry no kernel here. return nil, errf("%s: dtype %s is not supported; convert with Astype", name, a.dt) } out := &Array{shape: a.Shape(), dt: a.dt} if a.dt == Int { out.dt = Float } out.alloc(a.Len()) if out.dt == Float16 { parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { as, os := a.halves[s:e], out.halves[s:e] for i := range os { os[i] = HalfFromFloat64(f(HalfToFloat64(as[i]))) } }) return out, nil } if out.dt == Float32 { parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { as, os := a.floats32[s:e], out.floats32[s:e] for i := range os { os[i] = float32(f(float64(as[i]))) } }) return out, nil } // A float64 payload reads directly; int keeps the widening accessor. parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { os := out.floats[s:e] if a.dt == Float { as := a.floats[s:e] for i := range os { os[i] = f(as[i]) } return } for i := s; i < e; i++ { os[i-s] = f(a.floatAt(i)) } }) return out, nil } // roundFunc applies a rounding function element-wise: float arrays map to // their own width, int arrays are identity copies, complex arrays error. func (a *Array) roundFunc(name string, f func(float64) float64) (*Array, error) { if a.dt == Complex { return nil, errf("%s: complex arrays have no rounding", name) } if narrowRefused(a.dt) { // Integer arrays are identity copies under this family; bool and // the narrow widths carry no kernel here. return nil, errf("%s: dtype %s is not supported; convert with Astype", name, a.dt) } if a.dt == Int { ints, _, _, _, _ := a.cloneData() return &Array{shape: a.Shape(), dt: Int, ints: ints}, nil } out := &Array{shape: a.Shape(), dt: a.dt} out.alloc(a.Len()) if a.dt == Float16 { parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { as, os := a.halves[s:e], out.halves[s:e] for i := range os { os[i] = HalfFromFloat64(f(HalfToFloat64(as[i]))) } }) return out, nil } if a.dt == Float32 { parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { as, os := a.floats32[s:e], out.floats32[s:e] for i := range os { os[i] = float32(f(float64(as[i]))) } }) return out, nil } parallelMin(a.Len(), mathFuncMinPerWorker, func(s, e int) { as, os := a.floats[s:e], out.floats[s:e] for i := range os { os[i] = f(as[i]) } }) return out, nil } // Tanh returns the hyperbolic tangent of each element. func Tanh(a *Array) (*Array, error) { return a.realFunc("Tanh", math.Tanh) } // Sigmoid returns the logistic function 1/(1+e^−x) of each element. func Sigmoid(a *Array) (*Array, error) { return a.realFunc("Sigmoid", func(x float64) float64 { return 1 / (1 + math.Exp(-x)) }) }