428 lines
14 KiB
Go
428 lines
14 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (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))
|
||
})
|
||
}
|