1095 lines
33 KiB
Go
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
|
||
|
|
}
|