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))
|
|||
|
|
})
|
|||
|
|
}
|