Files
tensor/internal/core/mathfunc.go
T
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

428 lines
14 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
// 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))
})
}