// Copyright (c) 2026 Petr Balvín (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 }