feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
@@ -0,0 +1,427 @@
|
||||
// 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))
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user