Files
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

1095 lines
33 KiB
Go

// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package core
import "math"
// Element-wise arithmetic. All operations are
// package-level functions, not methods. Shapes must agree exactly
//: a mismatch is a loud error naming both shapes. The promotion
// ladder is int to float16 to float32 to float64 to complex128: int
// with int stays int (wrapping like Go's int64 arithmetic), float16
// keeps float16 against int,
// float32 keeps float32 against int or float16, any float64 operand
// promotes, any complex operand promotes. Div is true division (float
// or complex, IEEE behaviour on zero); Quo is integer division and
// rejects both float and complex operands.
// Add returns the element-wise sum of two arrays of the same shape.
func Add(a, b *Array) (*Array, error) { return elementwiseOp(a, b, "Add", pairAdd) }
// Sub returns the element-wise difference of two arrays of the same
// shape.
func Sub(a, b *Array) (*Array, error) { return elementwiseOp(a, b, "Sub", pairSub) }
// Mul returns the element-wise product of two arrays of the same shape.
func Mul(a, b *Array) (*Array, error) { return elementwiseOp(a, b, "Mul", pairMul) }
// Div returns the element-wise true division of two arrays of the same
// shape. The result is float for real operands and complex when a complex
// operand takes part; division by zero yields ±Inf or NaN per IEEE-754,
// never an error.
func Div(a, b *Array) (*Array, error) {
return elementwiseDiv(a, b)
}
// Quo returns the element-wise integer division of two int arrays of the
// same shape. A float or complex operand, or a zero divisor element, is
// an error.
func Quo(a, b *Array) (*Array, error) {
if a.dt != Int || b.dt != Int {
return nil, errf("Quo: needs int arrays, got %s and %s", a.dt, b.dt)
}
if !sameShape(a.shape, b.shape) {
return nil, errf("Quo: shape mismatch %s vs %s", shapeText(a.shape), shapeText(b.shape))
}
ints, _, _, _, _ := a.cloneData()
// The zero check rides the division walk: it visits the divisors in
// the same ascending order the separate scan did, so a zero reports
// the same element index, and a leading zero reports before any
// quotient is written.
for i := range ints {
d := b.ints[i]
if d == 0 {
return nil, errf("Quo: integer division by zero at element %d", i)
}
ints[i] /= d
}
return &Array{shape: a.Shape(), dt: Int, ints: ints}, nil
}
// scalarIntOp selects the int-class operation a mapKeeping walk applies.
type scalarIntOp uint8
const (
addScalarInt scalarIntOp = iota
subScalarInt
mulScalarInt
)
// AddI adds an int scalar element-wise; each dtype keeps its kind.
func AddI(a *Array, v int64) *Array {
return a.mapKeeping(v, addScalarInt,
func(x, s float64) float64 { return x + s },
func(x, s complex128) complex128 { return x + s })
}
// SubI subtracts an int scalar element-wise; each dtype keeps its kind.
func SubI(a *Array, v int64) *Array {
return a.mapKeeping(v, subScalarInt,
func(x, s float64) float64 { return x - s },
func(x, s complex128) complex128 { return x - s })
}
// MulI multiplies by an int scalar element-wise; each dtype keeps its
// kind.
func MulI(a *Array, v int64) *Array {
return a.mapKeeping(v, mulScalarInt,
func(x, s float64) float64 { return x * s },
func(x, s complex128) complex128 { return x * s })
}
// AddF adds a float scalar element-wise; integer arrays become
// float64, floating arrays keep their dtype, and complex arrays stay
// complex.
func AddF(a *Array, v float64) *Array {
return a.mapReal(v,
func(x float64) float64 { return x + v },
func(x complex128) complex128 { return x + complex(v, 0) })
}
// SubF subtracts a float scalar element-wise; integer arrays become
// float64, floating arrays keep their dtype, and complex arrays stay
// complex.
func SubF(a *Array, v float64) *Array {
return a.mapReal(v,
func(x float64) float64 { return x - v },
func(x complex128) complex128 { return x - complex(v, 0) })
}
// MulF multiplies by a float scalar element-wise; integer arrays
// become float64, floating arrays keep their dtype, and complex arrays
// stay complex.
func MulF(a *Array, v float64) *Array {
return a.mapReal(v,
func(x float64) float64 { return x * v },
func(x complex128) complex128 { return x * complex(v, 0) })
}
// DivI divides by an int scalar as true division; integer arrays
// become float64, floating arrays keep their dtype, and complex arrays
// stay complex, with IEEE behaviour when v is zero.
func DivI(a *Array, v int64) *Array {
return a.mapReal(float64(v),
func(x float64) float64 { return x / float64(v) },
func(x complex128) complex128 { return x / complex(float64(v), 0) })
}
// DivF divides by a float scalar as true division; integer arrays
// become float64, floating arrays keep their dtype, and complex arrays
// stay complex, with IEEE behaviour when v is zero.
func DivF(a *Array, v float64) *Array {
return a.mapReal(v,
func(x float64) float64 { return x / v },
func(x complex128) complex128 { return x / complex(v, 0) })
}
// AddC adds a complex scalar element-wise; the result is always complex.
func AddC(a *Array, v complex128) *Array {
return a.mapComplex(func(x complex128) complex128 { return x + v })
}
// SubC subtracts a complex scalar element-wise; the result is always
// complex.
func SubC(a *Array, v complex128) *Array {
return a.mapComplex(func(x complex128) complex128 { return x - v })
}
// MulC multiplies by a complex scalar element-wise; the result is always
// complex.
func MulC(a *Array, v complex128) *Array {
return a.mapComplex(func(x complex128) complex128 { return x * v })
}
// DivC divides by a complex scalar element-wise; the result is always
// complex.
func DivC(a *Array, v complex128) *Array {
return a.mapComplex(func(x complex128) complex128 { return x / v })
}
// QuoI divides an int array by an int scalar with integer division; a
// float or complex array, or a zero scalar, is an error.
func QuoI(a *Array, v int64) (*Array, error) {
if a.dt != Int {
return nil, errf("QuoI: needs an int array, got %s", a.dt)
}
if v == 0 {
return nil, errf("QuoI: integer division by zero")
}
ints, _, _, _, _ := a.cloneData()
parallelMin(len(ints), elementwiseMinPerWorker, func(s, e int) {
xs := ints[s:e]
for i := range xs {
xs[i] /= v
}
})
return &Array{shape: a.Shape(), dt: Int, ints: ints}, nil
}
// pairOp names an element-wise binary operation. The kernels write the
// operation as a direct expression per dtype, so the element loop holds
// an add or a multiply rather than a call through a func value, which
// neither inlines nor lets the compiler keep the payloads in registers
// across the loop. Only the mixed-dtype pair, whose walk pays an
// accessor call per element whatever the body does, keeps the closure
// form below.
type pairOp uint8
const (
pairAdd pairOp = iota
pairSub
pairMul
pairMin
pairMax
pairPow
)
// complexOK reports whether the operation accepts complex operands. The
// extrema have no ordering to compare, so they reject them; every other
// operation carries them.
func (o pairOp) complexOK() bool { return o != pairMin && o != pairMax }
// narrowIntClass reports whether dt is one of the small integer
// payloads: the dtypes whose same-width arithmetic runs in the element
// type itself.
func narrowIntClass(dt Dtype) bool {
switch dt {
case Int8, Uint8, Int16, Uint16, Int32, Uint32:
return true
}
return false
}
// opClosures returns op as the closure triple elementwise takes. It
// serves the mixed-dtype pair, where the two operands need different
// accessors and the walk rebuilds each element anyway; the expression is
// the one the specialised kernels write, so the results agree bit for
// bit.
func opClosures(op pairOp) (func(x, y int64) int64, func(x, y float64) float64, func(x, y complex128) complex128) {
switch op {
case pairAdd:
return func(x, y int64) int64 { return x + y },
func(x, y float64) float64 { return x + y },
func(x, y complex128) complex128 { return x + y }
case pairSub:
return func(x, y int64) int64 { return x - y },
func(x, y float64) float64 { return x - y },
func(x, y complex128) complex128 { return x - y }
case pairMul:
return func(x, y int64) int64 { return x * y },
func(x, y float64) float64 { return x * y },
func(x, y complex128) complex128 { return x * y }
case pairMin:
return func(x, y int64) int64 { return min(x, y) },
func(x, y float64) float64 { return min(x, y) }, nil
case pairMax:
return func(x, y int64) int64 { return max(x, y) },
func(x, y float64) float64 { return max(x, y) }, nil
default: // pairPow
return powInt,
func(x, y float64) float64 { return math.Pow(x, y) },
func(x, y complex128) complex128 { return mathPowComplex(x, y) }
}
}
// elementwiseOp applies op pairwise under the promotion ladder, with the
// operation written into each dtype's loop as a direct expression. The
// two same-width payloads are read as slices, so the loop is a load, an
// operation and a store per element; a mixed pair and a strided view
// fall back to elementwise, whose accessor walk sees the same values in
// the same order.
func elementwiseOp(a, b *Array, name string, op pairOp) (*Array, error) {
if !sameShape(a.shape, b.shape) {
return nil, errf("%s: shape mismatch %s vs %s", name, shapeText(a.shape), shapeText(b.shape))
}
dt := promote(a.dt, b.dt)
if dt == Bool {
// Bool carries no arithmetic: the loud refusal sits here because
// every arithmetic entry (Add, Sub, Mul, Div, Pow, the pairwise
// extrema) funnels through this throat with its own name.
return nil, errf("%s: bool arrays have no arithmetic", name)
}
if dt == Complex && !op.complexOK() {
return nil, errf("%s: complex arrays have no ordering", name)
}
if op == pairPow && intClass(dt) {
// A negative exponent has no integer result: the same refusal
// Pow makes for int operands, asked of the whole integer class
// through intAt's exact widenings. promote answers an
// integer-class dtype exactly when both operands are
// integer-class, so a gate on the narrow widths alone would
// miss the mixed pair whose promoted dtype is Int.
for i := range b.Len() {
if e := b.intAt(i); e < 0 {
return nil, errf("%s: negative exponent %d at element %d has no int result", name, e, i)
}
}
}
if a.dt != dt || b.dt != dt || !a.isContiguous() || !b.isContiguous() {
ii, ff, cc := opClosures(op)
return elementwise(a, b, name, ii, ff, cc)
}
out := &Array{shape: a.Shape(), dt: dt}
out.alloc(a.Len())
n := a.Len()
switch dt {
case Int8:
narrowPairRun(op, a.i8s, b.i8s, out.i8s)
case Uint8:
narrowPairRun(op, a.u8s, b.u8s, out.u8s)
case Int16:
narrowPairRun(op, a.i16s, b.i16s, out.i16s)
case Uint16:
narrowPairRun(op, a.u16s, b.u16s, out.u16s)
case Int32:
narrowPairRun(op, a.i32s, b.i32s, out.i32s)
case Uint32:
narrowPairRun(op, a.u32s, b.u32s, out.u32s)
case Int:
ai, bi, oi := a.ints, b.ints, out.ints
parallelMin(n, elementwiseMinPerWorker, func(s, e int) {
as, bs, os := ai[s:e], bi[s:e], oi[s:e]
switch op {
case pairAdd:
for i := range os {
os[i] = as[i] + bs[i]
}
case pairSub:
for i := range os {
os[i] = as[i] - bs[i]
}
case pairMul:
for i := range os {
os[i] = as[i] * bs[i]
}
case pairMin:
for i := range os {
os[i] = min(as[i], bs[i])
}
case pairMax:
for i := range os {
os[i] = max(as[i], bs[i])
}
default: // pairPow
for i := range os {
os[i] = powInt(as[i], bs[i])
}
}
})
case Float16:
// Float16 mirrors float32 one step down the ladder: the
// arithmetic runs in float64, where every half value is exact,
// and narrows once per element.
ai, bi, oi := a.halves, b.halves, out.halves
parallelMin(n, elementwiseMinPerWorker, func(s, e int) {
as, bs, os := ai[s:e], bi[s:e], oi[s:e]
switch op {
case pairAdd:
for i := range os {
os[i] = HalfFromFloat64(HalfToFloat64(as[i]) + HalfToFloat64(bs[i]))
}
case pairSub:
for i := range os {
os[i] = HalfFromFloat64(HalfToFloat64(as[i]) - HalfToFloat64(bs[i]))
}
case pairMul:
for i := range os {
os[i] = HalfFromFloat64(HalfToFloat64(as[i]) * HalfToFloat64(bs[i]))
}
case pairMin:
for i := range os {
os[i] = HalfFromFloat64(min(HalfToFloat64(as[i]), HalfToFloat64(bs[i])))
}
case pairMax:
for i := range os {
os[i] = HalfFromFloat64(max(HalfToFloat64(as[i]), HalfToFloat64(bs[i])))
}
default: // pairPow
for i := range os {
os[i] = HalfFromFloat64(math.Pow(HalfToFloat64(as[i]), HalfToFloat64(bs[i])))
}
}
})
case Float32:
// The widening to float64 is exact and the result narrows once,
// so the widened spelling is the accessor expression unchanged.
ai, bi, oi := a.floats32, b.floats32, out.floats32
parallelMin(n, elementwiseMinPerWorker, func(s, e int) {
as, bs, os := ai[s:e], bi[s:e], oi[s:e]
switch op {
case pairAdd:
for i := range os {
os[i] = float32(float64(as[i]) + float64(bs[i]))
}
case pairSub:
for i := range os {
os[i] = float32(float64(as[i]) - float64(bs[i]))
}
case pairMul:
for i := range os {
os[i] = float32(float64(as[i]) * float64(bs[i]))
}
case pairMin:
for i := range os {
os[i] = float32(min(float64(as[i]), float64(bs[i])))
}
case pairMax:
for i := range os {
os[i] = float32(max(float64(as[i]), float64(bs[i])))
}
default: // pairPow
for i := range os {
os[i] = float32(math.Pow(float64(as[i]), float64(bs[i])))
}
}
})
case Float:
ai, bi, oi := a.floats, b.floats, out.floats
parallelMin(n, elementwiseMinPerWorker, func(s, e int) {
as, bs, os := ai[s:e], bi[s:e], oi[s:e]
switch op {
case pairAdd:
for i := range os {
os[i] = as[i] + bs[i]
}
case pairSub:
for i := range os {
os[i] = as[i] - bs[i]
}
case pairMul:
for i := range os {
os[i] = as[i] * bs[i]
}
case pairMin:
for i := range os {
os[i] = min(as[i], bs[i])
}
case pairMax:
for i := range os {
os[i] = max(as[i], bs[i])
}
default: // pairPow
for i := range os {
os[i] = math.Pow(as[i], bs[i])
}
}
})
default:
ai, bi, oi := a.complexes, b.complexes, out.complexes
parallelMin(n, elementwiseMinPerWorker, func(s, e int) {
as, bs, os := ai[s:e], bi[s:e], oi[s:e]
switch op {
case pairAdd:
for i := range os {
os[i] = as[i] + bs[i]
}
case pairSub:
for i := range os {
os[i] = as[i] - bs[i]
}
case pairMul:
for i := range os {
os[i] = as[i] * bs[i]
}
default: // pairPow, the only complex operation left
for i := range os {
os[i] = mathPowComplex(as[i], bs[i])
}
}
})
}
return out, nil
}
// narrowPairRun applies op to two same-width integer payloads in the
// element type itself: Go's arithmetic wraps on overflow at every width,
// exactly the semantics the int64 kernels carry, and the pow widens
// exactly and narrows on store.
func narrowPairRun[T int8 | uint8 | int16 | uint16 | int32 | uint32](op pairOp, x, y, dst []T) {
parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) {
as, bs, os := x[s:e], y[s:e], dst[s:e]
switch op {
case pairAdd:
for i := range os {
os[i] = as[i] + bs[i]
}
case pairSub:
for i := range os {
os[i] = as[i] - bs[i]
}
case pairMul:
for i := range os {
os[i] = as[i] * bs[i]
}
case pairMin:
for i := range os {
os[i] = min(as[i], bs[i])
}
case pairMax:
for i := range os {
os[i] = max(as[i], bs[i])
}
default: // pairPow
for i := range os {
os[i] = T(powInt(int64(as[i]), int64(bs[i])))
}
}
})
}
// elementwise applies an operation pairwise under the promotion ladder.
// A nil cc marks an ordering, which complex operands reject. float16 and
// float32 operands compute in float64 and round once: their arithmetic
// is exact in float64. The operation arrives as closures, which
// the mixed-dtype pair of elementwiseOp and the extrema family use.
func elementwise(a, b *Array, name string,
ii func(x, y int64) int64, ff func(x, y float64) float64, cc func(x, y complex128) complex128,
) (*Array, error) {
if !sameShape(a.shape, b.shape) {
return nil, errf("%s: shape mismatch %s vs %s", name, shapeText(a.shape), shapeText(b.shape))
}
dt := promote(a.dt, b.dt)
if dt == Bool {
return nil, errf("%s: bool arrays have no arithmetic", name)
}
if dt == Complex && cc == nil {
return nil, errf("%s: complex arrays have no ordering", name)
}
out := &Array{shape: a.Shape(), dt: dt}
out.alloc(a.Len())
// A stride table is a mechanism no public constructor sets, but the
// accessor walk is the one reader that resolves it, so every
// raw-payload branch below is gated on both operands being dense.
dense := a.isContiguous() && b.isContiguous()
switch dt {
case Int8:
narrowElementwise(a, b, out, func(x *Array) []int8 { return x.i8s }, dense, ii)
case Uint8:
narrowElementwise(a, b, out, func(x *Array) []uint8 { return x.u8s }, dense, ii)
case Int16:
narrowElementwise(a, b, out, func(x *Array) []int16 { return x.i16s }, dense, ii)
case Uint16:
narrowElementwise(a, b, out, func(x *Array) []uint16 { return x.u16s }, dense, ii)
case Int32:
narrowElementwise(a, b, out, func(x *Array) []int32 { return x.i32s }, dense, ii)
case Uint32:
narrowElementwise(a, b, out, func(x *Array) []uint32 { return x.u32s }, dense, ii)
case Int:
// Dense same-width operands stream their raw payloads; a mixed
// integer pair promoted to Int (int32 with uint32, any narrow
// operand with Int) reads through intAt, which widens every
// integer-class payload exactly. This mirrors the float32
// branch's operand gates: the raw read is only sound when the
// operand really carries the Int payload.
aI, bI := a.dt == Int, b.dt == Int
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
os := out.ints[s:e]
switch {
case dense && aI && bI:
as, bs := a.ints[s:e], b.ints[s:e]
for i := range os {
os[i] = ii(as[i], bs[i])
}
case dense && aI:
as := a.ints[s:e]
for i := range os {
os[i] = ii(as[i], b.intAt(s+i))
}
case dense && bI:
bs := b.ints[s:e]
for i := range os {
os[i] = ii(a.intAt(s+i), bs[i])
}
default:
for i := s; i < e; i++ {
os[i-s] = ii(a.intAt(i), b.intAt(i))
}
}
})
case Float16:
// Float16 mirrors float32 one step down the ladder: the
// arithmetic runs in float64, where every half value is exact,
// and narrows once per element. Dense same-width operands stream
// their raw payloads; everything strided or mixed keeps the
// accessor fallback.
aH, bH := a.dt == Float16, b.dt == Float16
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
os := out.halves[s:e]
switch {
case dense && aH && bH:
as, bs := a.halves[s:e], b.halves[s:e]
for i := range os {
os[i] = HalfFromFloat64(ff(HalfToFloat64(as[i]), HalfToFloat64(bs[i])))
}
case dense && aH:
as := a.halves[s:e]
for i := range os {
os[i] = HalfFromFloat64(ff(HalfToFloat64(as[i]), b.floatAt(s+i)))
}
case dense && bH:
bs := b.halves[s:e]
for i := range os {
os[i] = HalfFromFloat64(ff(a.floatAt(s+i), HalfToFloat64(bs[i])))
}
default:
for i := s; i < e; i++ {
os[i-s] = HalfFromFloat64(ff(a.floatAt(i), b.floatAt(i)))
}
}
})
case Float32:
// Dense same-width operands stream their raw payloads: float32
// widens exactly, so the values and their order match accessor
// reads bit for bit. Everything strided or mixed keeps the
// accessor fallback.
a32, b32 := a.dt == Float32, b.dt == Float32
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
os := out.floats32[s:e]
switch {
case dense && a32 && b32:
as, bs := a.floats32[s:e], b.floats32[s:e]
for i := range os {
os[i] = float32(ff(float64(as[i]), float64(bs[i])))
}
case dense && a32:
as := a.floats32[s:e]
for i := range os {
os[i] = float32(ff(float64(as[i]), b.floatAt(s+i)))
}
case dense && b32:
bs := b.floats32[s:e]
for i := range os {
os[i] = float32(ff(a.floatAt(s+i), float64(bs[i])))
}
default:
for i := s; i < e; i++ {
os[i-s] = float32(ff(a.floatAt(i), b.floatAt(i)))
}
}
})
case Float:
aF, bF := a.dt == Float, b.dt == Float
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
os := out.floats[s:e]
switch {
case dense && aF && bF:
as, bs := a.floats[s:e], b.floats[s:e]
for i := range os {
os[i] = ff(as[i], bs[i])
}
case dense && aF:
as := a.floats[s:e]
for i := range os {
os[i] = ff(as[i], b.floatAt(s+i))
}
case dense && bF:
bs := b.floats[s:e]
for i := range os {
os[i] = ff(a.floatAt(s+i), bs[i])
}
default:
for i := s; i < e; i++ {
os[i-s] = ff(a.floatAt(i), b.floatAt(i))
}
}
})
default:
aC, bC := a.dt == Complex, b.dt == Complex
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
os := out.complexes[s:e]
switch {
case dense && aC && bC:
as, bs := a.complexes[s:e], b.complexes[s:e]
for i := range os {
os[i] = cc(as[i], bs[i])
}
case dense && aC:
as := a.complexes[s:e]
for i := range os {
os[i] = cc(as[i], b.complexAt(s+i))
}
case dense && bC:
bs := b.complexes[s:e]
for i := range os {
os[i] = cc(a.complexAt(s+i), bs[i])
}
default:
for i := s; i < e; i++ {
os[i-s] = cc(a.complexAt(i), b.complexAt(i))
}
}
})
}
return out, nil
}
// narrowElementwise is elementwise's walk for a narrow integer promoted
// result: dense same-width operands stream their own payloads, every
// mixed or strided read goes through intAt, and the int64 closure value
// narrows on store with Go's conversion, the implicit-store cast
// setConverted performs.
func narrowElementwise[T int8 | uint8 | int16 | uint16 | int32 | uint32](
a, b, out *Array, payload func(*Array) []T, dense bool, ii func(x, y int64) int64,
) {
aN, bN := a.dt == out.dt, b.dt == out.dt
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
os := payload(out)[s:e]
switch {
case dense && aN && bN:
as, bs := payload(a)[s:e], payload(b)[s:e]
for i := range os {
os[i] = T(ii(int64(as[i]), int64(bs[i])))
}
case dense && aN:
as := payload(a)[s:e]
for i := range os {
os[i] = T(ii(int64(as[i]), b.intAt(s+i)))
}
case dense && bN:
bs := payload(b)[s:e]
for i := range os {
os[i] = T(ii(a.intAt(s+i), int64(bs[i])))
}
default:
for i := s; i < e; i++ {
os[i-s] = T(ii(a.intAt(i), b.intAt(i)))
}
}
})
}
// elementwiseDiv applies true division; the result follows the promotion
// ladder, except that every integer-class pair divides into float64.
func elementwiseDiv(a, b *Array) (*Array, error) {
if !sameShape(a.shape, b.shape) {
return nil, errf("Div: shape mismatch %s vs %s", shapeText(a.shape), shapeText(b.shape))
}
dt := promote(a.dt, b.dt)
if dt == Bool {
return nil, errf("Div: bool arrays have no arithmetic")
}
if intClass(dt) {
// Integer true division has always answered float64 for an int
// pair; bool and the narrow integer widths take the same route,
// their widenings exact and the division run in float64.
dt = Float
}
out := &Array{shape: a.Shape(), dt: dt}
out.alloc(a.Len())
// The same contiguity gate the elementwise fallback carries: a
// strided operand reads through the accessors, never in payload
// order.
dense := a.isContiguous() && b.isContiguous()
switch dt {
case Float16:
// Dense same-width float16 operands stream their raw payloads;
// the division runs in float64 and narrows once per element, the
// mirror of the float32 path below.
aH, bH := a.dt == Float16, b.dt == Float16
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
os := out.halves[s:e]
switch {
case dense && aH && bH:
as, bs := a.halves[s:e], b.halves[s:e]
for i := range os {
os[i] = HalfFromFloat64(HalfToFloat64(as[i]) / HalfToFloat64(bs[i]))
}
case dense && aH:
as := a.halves[s:e]
for i := range os {
os[i] = HalfFromFloat64(HalfToFloat64(as[i]) / b.floatAt(s+i))
}
case dense && bH:
bs := b.halves[s:e]
for i := range os {
os[i] = HalfFromFloat64(a.floatAt(s+i) / HalfToFloat64(bs[i]))
}
default:
for i := s; i < e; i++ {
os[i-s] = HalfFromFloat64(a.floatAt(i) / b.floatAt(i))
}
}
})
case Float32:
// Dense same-width float32 operands stream their raw payloads;
// the widening is exact, so nothing shifts a bit. Everything
// strided or mixed keeps the accessor fallback.
a32, b32 := a.dt == Float32, b.dt == Float32
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
os := out.floats32[s:e]
switch {
case dense && a32 && b32:
as, bs := a.floats32[s:e], b.floats32[s:e]
for i := range os {
os[i] = float32(float64(as[i]) / float64(bs[i]))
}
case dense && a32:
as := a.floats32[s:e]
for i := range os {
os[i] = float32(float64(as[i]) / b.floatAt(s+i))
}
case dense && b32:
bs := b.floats32[s:e]
for i := range os {
os[i] = float32(a.floatAt(s+i) / float64(bs[i]))
}
default:
for i := s; i < e; i++ {
os[i-s] = float32(a.floatAt(i) / b.floatAt(i))
}
}
})
case Float:
aF, bF := a.dt == Float, b.dt == Float
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
os := out.floats[s:e]
switch {
case dense && aF && bF:
as, bs := a.floats[s:e], b.floats[s:e]
for i := range os {
os[i] = as[i] / bs[i]
}
case dense && aF:
as := a.floats[s:e]
for i := range os {
os[i] = as[i] / b.floatAt(s+i)
}
case dense && bF:
bs := b.floats[s:e]
for i := range os {
os[i] = a.floatAt(s+i) / bs[i]
}
default:
for i := s; i < e; i++ {
os[i-s] = a.floatAt(i) / b.floatAt(i)
}
}
})
default:
aC, bC := a.dt == Complex, b.dt == Complex
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
os := out.complexes[s:e]
switch {
case dense && aC && bC:
as, bs := a.complexes[s:e], b.complexes[s:e]
for i := range os {
os[i] = as[i] / bs[i]
}
case dense && aC:
as := a.complexes[s:e]
for i := range os {
os[i] = as[i] / b.complexAt(s+i)
}
case dense && bC:
bs := b.complexes[s:e]
for i := range os {
os[i] = a.complexAt(s+i) / bs[i]
}
default:
for i := s; i < e; i++ {
os[i-s] = a.complexAt(i) / b.complexAt(i)
}
}
})
}
return out, nil
}
// mapKeeping applies a scalar operation that keeps each real dtype and
// stays complex on complex arrays. The int-class operation arrives as an
// enum, so every element loop below holds the operation itself rather
// than a call through a func value.
func (a *Array) mapKeeping(v int64, op scalarIntOp,
ff func(x, s float64) float64, cc func(x, s complex128) complex128,
) *Array {
if a.dt == Bool {
// Bool carries no arithmetic, and this keeping-kind surface has
// no error channel; nil is the refusal shape the no-error
// constructors already use.
return nil
}
out := &Array{shape: a.Shape(), dt: a.dt}
out.alloc(a.Len())
switch a.dt {
case Int:
src, dst := a.ints, out.ints
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
xs, os := src[s:e], dst[s:e]
switch op {
case addScalarInt:
for i := range os {
os[i] = xs[i] + v
}
case subScalarInt:
for i := range os {
os[i] = xs[i] - v
}
default: // mulScalarInt
for i := range os {
os[i] = xs[i] * v
}
}
})
case Int8, Uint8, Int16, Uint16, Int32, Uint32:
// The keeping-kind walk for the narrow integer widths: the int64
// operation narrows on store with Go's conversion, the same
// implicit-store cast setConverted performs.
switch a.dt {
case Int8:
narrowMapKeeping(a, out, v, op, func(x *Array) []int8 { return x.i8s })
case Uint8:
narrowMapKeeping(a, out, v, op, func(x *Array) []uint8 { return x.u8s })
case Int16:
narrowMapKeeping(a, out, v, op, func(x *Array) []int16 { return x.i16s })
case Uint16:
narrowMapKeeping(a, out, v, op, func(x *Array) []uint16 { return x.u16s })
case Int32:
narrowMapKeeping(a, out, v, op, func(x *Array) []int32 { return x.i32s })
default:
narrowMapKeeping(a, out, v, op, func(x *Array) []uint32 { return x.u32s })
}
case Float16:
fv := float64(v)
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
as, os := a.halves[s:e], out.halves[s:e]
for i := range os {
os[i] = HalfFromFloat64(ff(HalfToFloat64(as[i]), fv))
}
})
case Float32:
fv := float64(v)
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(ff(float64(as[i]), fv))
}
})
case Float:
fv := float64(v)
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
as, os := a.floats[s:e], out.floats[s:e]
for i := range os {
os[i] = ff(as[i], fv)
}
})
default:
cv := complex(float64(v), 0)
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
as, os := a.complexes[s:e], out.complexes[s:e]
for i := range os {
os[i] = cc(as[i], cv)
}
})
}
return out
}
// narrowMapKeeping is mapKeeping's keeping-kind walk for a narrow
// integer payload: each element computes in int64 from the exact
// widening and narrows on store, the operation written into the loop.
func narrowMapKeeping[T int8 | uint8 | int16 | uint16 | int32 | uint32](
a, out *Array, v int64, op scalarIntOp, pick func(*Array) []T,
) {
src, dst := pick(a), pick(out)
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
xs, os := src[s:e], dst[s:e]
switch op {
case addScalarInt:
for i := range os {
os[i] = T(int64(xs[i]) + v)
}
case subScalarInt:
for i := range os {
os[i] = T(int64(xs[i]) - v)
}
default: // mulScalarInt
for i := range os {
os[i] = T(int64(xs[i]) * v)
}
}
})
}
// mapReal applies a scalar operation: integer arrays widen to float64,
// float16 computes in float64 and narrows, float32 arrays keep
// float32 (the scalar is weak and narrows), float64 stays
// float64, complex stays complex.
func (a *Array) mapReal(v float64, f func(float64) float64, c func(complex128) complex128) *Array {
if a.dt == Complex {
return a.mapComplex(c)
}
out := &Array{shape: a.Shape(), dt: a.dt}
if intClass(a.dt) {
// Every integer-class array, bool and the narrow widths included,
// widens to float64 under this surface's contract; the accessor
// walk below reads each widening exactly.
out.dt = Float
}
out.alloc(a.Len())
if a.dt == Float16 {
parallelMin(a.Len(), elementwiseMinPerWorker, 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
}
if a.dt == Float32 {
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(f(float64(as[i])))
}
})
return out
}
// A float64 payload reads directly; the integer-class payloads read
// their own slices when dense, each widening exactly the one the
// accessor hands over; an integer-class view keeps the accessor
// walk.
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
os := out.floats[s:e]
switch {
case a.dt == Float:
as := a.floats[s:e]
for i := range os {
os[i] = f(as[i])
}
case a.strides != nil:
for i := s; i < e; i++ {
os[i-s] = f(a.floatAt(i))
}
default:
intRealRun(a, os, s, f)
}
})
return out
}
// intRealRun is mapReal's widening walk for a dense integer-class
// payload: the dtype dispatch sits outside the element loop and the
// float64 operation f reads each exact widening.
func intRealRun(a *Array, os []float64, s int, f func(float64) float64) {
switch a.dt {
case Int:
as := a.ints[s : s+len(os)]
for i := range os {
os[i] = f(float64(as[i]))
}
case Bool:
as := a.bools[s : s+len(os)]
for i := range os {
w := 0.0
if as[i] {
w = 1
}
os[i] = f(w)
}
case Int8:
as := a.i8s[s : s+len(os)]
for i := range os {
os[i] = f(float64(as[i]))
}
case Uint8:
as := a.u8s[s : s+len(os)]
for i := range os {
os[i] = f(float64(as[i]))
}
case Int16:
as := a.i16s[s : s+len(os)]
for i := range os {
os[i] = f(float64(as[i]))
}
case Uint16:
as := a.u16s[s : s+len(os)]
for i := range os {
os[i] = f(float64(as[i]))
}
case Int32:
as := a.i32s[s : s+len(os)]
for i := range os {
os[i] = f(float64(as[i]))
}
default: // Uint32
as := a.u32s[s : s+len(os)]
for i := range os {
os[i] = f(float64(as[i]))
}
}
}
// mapComplex applies an operation whose result is always complex.
func (a *Array) mapComplex(c func(complex128) complex128) *Array {
out := &Array{shape: a.Shape(), dt: Complex}
out.complexes = make([]complex128, a.Len())
if a.dt == Complex {
ac := a.complexes
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
os := out.complexes[s:e]
as := ac[s:e]
for i := range os {
os[i] = c(as[i])
}
})
return out
}
parallelMin(a.Len(), elementwiseMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
out.complexes[i] = c(a.complexAt(i))
}
})
return out
}