Files

1284 lines
35 KiB
Go
Raw Permalink Normal View History

2026-09-03 10:00:00 +02:00
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
// SPDX-License-Identifier: MIT
package core
import (
"math"
"sync"
)
// Masks. Element-wise comparisons answer a bool mask: the
// payload is the predicate itself, one byte per element against the
// eight the 0/1 int mask of earlier releases wrote. Masks compose
// through the bool logic operations And, Or, Xor and Not at the foot of
// this file, and bool carries no arithmetic. Select picks the elements
// where a mask is set and Where chooses element-wise between two
// arrays; both take a bool mask, and both still take the int mask of
// 0/1 the comparisons used to answer, a nonzero element standing for
// true. NaN follows IEEE: every ordering and equality comparison is
// false, except Ne, which is true. Complex arrays take part in Eq and
// Ne only: ordering has no complex meaning.
// cmpOp selects the relation a comparison reports.
type cmpOp uint8
const (
opEq cmpOp = iota
opNe
opLt
opLe
opGt
opGe
)
// isOrdering reports whether op needs an order, which complex operands
// do not have.
func (o cmpOp) isOrdering() bool { return o == opLt || o == opLe || o == opGt || o == opGe }
// mirror returns op with its operands swapped: x < y is y > x.
func (o cmpOp) mirror() cmpOp {
switch o {
case opLt:
return opGt
case opGt:
return opLt
case opLe:
return opGe
case opGe:
return opLe
default:
return o
}
}
// applyInt applies op to two int64 values.
func applyInt(op cmpOp, x, y int64) bool {
switch op {
case opEq:
return x == y
case opNe:
return x != y
case opLt:
return x < y
case opLe:
return x <= y
case opGt:
return x > y
default:
return x >= y
}
}
// applyFloat applies op to two float64 values; NaN follows IEEE, so
// every ordering and equality comparison is false and Ne is true.
func applyFloat(op cmpOp, x, y float64) bool {
switch op {
case opEq:
return x == y
case opNe:
return x != y
case opLt:
return x < y
case opLe:
return x <= y
case opGt:
return x > y
default:
return x >= y
}
}
// applyComplex applies op to two complex values; complex takes part in
// Eq and Ne only, so an ordering op never reaches here.
func applyComplex(op cmpOp, x, y complex128) bool {
if op == opNe {
return x != y
}
return x == y
}
// twoPow63 is 2^63 as a float64: the first value above every int64.
const twoPow63 = 9223372036854775808.0
// applyIntFloat reports op(x, y) for an int64 x and a float64 y, exactly.
// Widening x to float64 rounds two neighbouring ints above 2^53 onto the
// same value, which is how Eq once answered 1 for 2^60+1 against the
// float 2^60. An integral y in int64 range compares natively instead; a
// fractional y, which no int can equal, is decided by its floor and
// ceiling, because an integer is below a fractional y exactly when it is
// at most floor(y) and above it at least ceil(y).
func applyIntFloat(op cmpOp, x int64, y float64) bool {
switch {
case math.IsNaN(y):
return op == opNe
case math.IsInf(y, 1):
return op == opNe || op == opLt || op == opLe
case math.IsInf(y, -1):
return op == opNe || op == opGt || op == opGe
case y >= twoPow63:
// Above every int64.
return op == opNe || op == opLt || op == opLe
case y < -twoPow63:
// Below every int64.
return op == opNe || op == opGt || op == opGe
}
trunc := int64(y)
if float64(trunc) == y {
return applyInt(op, x, trunc)
}
floor, ceil := trunc, trunc+1
if y < 0 {
floor, ceil = trunc-1, trunc
}
switch op {
case opLt, opLe:
return x <= floor
case opGt, opGe:
return x >= ceil
case opNe:
return true
default: // opEq
return false
}
}
// cmpArray runs an element-wise comparison, producing a bool mask. An
// ordering op rejects complex operands; mixed int and float operands
// compare exactly through applyIntFloat, while the int, float and
// complex branches widen exactly. Dense same-width payloads take the
// raw walk above, a mixed narrow-width pair takes its int64 widening
// walk; a strided view and a pair that puts a bool beside a narrow
// width keep the accessor walk, which reads the same values in the
// same order.
func (a *Array) cmpArray(b *Array, name string, op cmpOp) (*Array, error) {
if !sameShape(a.shape, b.shape) {
return nil, errf("%s: shape mismatch %s vs %s", name, shapeText(a.shape), shapeText(b.shape))
}
if op.isOrdering() && (a.dt == Complex || b.dt == Complex) {
return nil, errf("%s: complex arrays have no ordering", name)
}
n := a.Len()
out := &Array{shape: a.Shape(), dt: Bool, bools: make([]bool, n)}
switch {
case a.dt == Int && b.dt == Int:
if a.isContiguous() && b.isContiguous() {
intCmpRun(op, a.ints[:n], b.ints[:n], out.bools)
return out, nil
}
parallelMin(n, elementwiseMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
out.bools[i] = applyInt(op, a.ints[a.physIndex(i)], b.ints[b.physIndex(i)])
}
})
case promote(a.dt, b.dt) == Complex:
if a.dt == Complex && b.dt == Complex && a.isContiguous() && b.isContiguous() {
complexCmpRun(op, a.complexes[:n], b.complexes[:n], out.bools)
return out, nil
}
parallelMin(n, elementwiseMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
out.bools[i] = applyComplex(op, a.complexAt(i), b.complexAt(i))
}
})
case a.dt != Int && b.dt != Int:
if a.dt == b.dt && a.isContiguous() && b.isContiguous() {
switch a.dt {
case Float:
floatCmpRun(op, a.floats[:n], b.floats[:n], out.bools)
return out, nil
case Float32:
floatCmpRun(op, a.floats32[:n], b.floats32[:n], out.bools)
return out, nil
case Bool:
boolCmpRun(op, a.bools[:n], b.bools[:n], out.bools)
return out, nil
case Int8:
narrowCmpRun(op, a.i8s[:n], b.i8s[:n], out.bools)
return out, nil
case Uint8:
narrowCmpRun(op, a.u8s[:n], b.u8s[:n], out.bools)
return out, nil
case Int16:
narrowCmpRun(op, a.i16s[:n], b.i16s[:n], out.bools)
return out, nil
case Uint16:
narrowCmpRun(op, a.u16s[:n], b.u16s[:n], out.bools)
return out, nil
case Int32:
narrowCmpRun(op, a.i32s[:n], b.i32s[:n], out.bools)
return out, nil
case Uint32:
narrowCmpRun(op, a.u32s[:n], b.u32s[:n], out.bools)
return out, nil
}
}
if narrowIntClass(a.dt) && narrowIntClass(b.dt) && a.dt != b.dt && a.isContiguous() && b.isContiguous() {
// Two narrow widths: the mixed walk below widens both sides
// exactly, so the comparison runs over the int64 widenings
// in the raw kernel. Bool is an integer-class dtype but not
// a narrow payload, so it keeps the accessor walk.
narrowCmpMixed(op, a, b, n, out.bools)
return out, nil
}
parallelMin(n, elementwiseMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
out.bools[i] = applyFloat(op, a.floatAt(i), b.floatAt(i))
}
})
default:
// One side int, the other float: the int side stays exact, and
// the float side takes the comparison in its own order.
ints, floats, o := a, b, op
if b.dt == Int {
ints, floats, o = b, a, op.mirror()
}
parallelMin(n, elementwiseMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
out.bools[i] = applyIntFloat(o, ints.ints[ints.physIndex(i)], floats.floatAt(i))
}
})
}
return out, nil
}
// cmpScalarI compares against an int scalar. Int arrays compare
// natively, float arrays compare exactly against the int through
// applyIntFloat rather than through the rounded float64(v), and complex
// arrays compare a zero imaginary part plus the exact real part.
func (a *Array) cmpScalarI(v int64, name string, op cmpOp) (*Array, error) {
if op.isOrdering() && a.dt == Complex {
return nil, errf("%s: complex arrays have no ordering", name)
}
n := a.Len()
out := &Array{shape: a.Shape(), dt: Bool, bools: make([]bool, n)}
switch a.dt {
case Int:
if a.isContiguous() {
intCmpScalarRun(op, a.ints[:n], v, out.bools)
return out, nil
}
parallelMin(n, elementwiseMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
out.bools[i] = applyInt(op, a.ints[a.physIndex(i)], v)
}
})
case Complex:
// The scalar is an int, so the comparison is exact against the
// real part and asks whether the imaginary part is zero: the
// equality test's answer is inverted for every other relation.
want := op == opEq
parallelMin(n, elementwiseMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
z := a.complexes[a.physIndex(i)]
eq := imag(z) == 0 && applyIntFloat(opEq, v, real(z))
out.bools[i] = eq == want
}
})
default:
if a.isContiguous() {
// The scalar relation is a property of the call, so the
// mirror leaves the loop and the walk reads each payload
// slice directly, widening exactly as the accessor would.
scalarCmpIRun(op, a, v, out.bools)
return out, nil
}
// The mirror is a property of the call, so it leaves the loop.
mo := op.mirror()
parallelMin(n, elementwiseMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
out.bools[i] = applyIntFloat(mo, v, a.floatAt(i))
}
})
}
return out, nil
}
// scalarCmpIRun compares a dense real payload against an int scalar. The
// float sides keep the exact int-versus-float relation applyIntFloat
// gives; the integer-class sides widen exactly, so the int comparison is
// the accessor walk's answer, its float detour included.
func scalarCmpIRun(op cmpOp, a *Array, v int64, dst []bool) {
mo := op.mirror()
parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) {
zs := dst[s:e]
switch a.dt {
case Float:
xs := a.floats[s:e]
for i := range zs {
zs[i] = applyIntFloat(mo, v, xs[i])
}
case Float32:
xs := a.floats32[s:e]
for i := range zs {
zs[i] = applyIntFloat(mo, v, float64(xs[i]))
}
case Float16:
xs := a.halves[s:e]
for i := range zs {
zs[i] = applyIntFloat(mo, v, HalfToFloat64(xs[i]))
}
case Bool:
xs := a.bools[s:e]
for i := range zs {
w := int64(0)
if xs[i] {
w = 1
}
zs[i] = applyInt(mo, v, w)
}
case Int8:
xs := a.i8s[s:e]
for i := range zs {
zs[i] = applyInt(mo, v, int64(xs[i]))
}
case Uint8:
xs := a.u8s[s:e]
for i := range zs {
zs[i] = applyInt(mo, v, int64(xs[i]))
}
case Int16:
xs := a.i16s[s:e]
for i := range zs {
zs[i] = applyInt(mo, v, int64(xs[i]))
}
case Uint16:
xs := a.u16s[s:e]
for i := range zs {
zs[i] = applyInt(mo, v, int64(xs[i]))
}
case Int32:
xs := a.i32s[s:e]
for i := range zs {
zs[i] = applyInt(mo, v, int64(xs[i]))
}
default: // Uint32
xs := a.u32s[s:e]
for i := range zs {
zs[i] = applyInt(mo, v, int64(xs[i]))
}
}
})
}
// cmpScalarF compares against a float scalar. Int arrays compare
// exactly through applyIntFloat, real arrays compare in float64, and
// complex arrays compare a zero imaginary part plus the real part.
func (a *Array) cmpScalarF(v float64, name string, op cmpOp) (*Array, error) {
if op.isOrdering() && a.dt == Complex {
return nil, errf("%s: complex arrays have no ordering", name)
}
n := a.Len()
out := &Array{shape: a.Shape(), dt: Bool, bools: make([]bool, n)}
switch a.dt {
case Int:
if a.isContiguous() {
intCmpScalarFRun(op, a.ints[:n], v, out.bools)
return out, nil
}
parallelMin(n, elementwiseMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
out.bools[i] = applyIntFloat(op, a.ints[a.physIndex(i)], v)
}
})
case Complex:
z := complex(v, 0)
parallelMin(n, elementwiseMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
out.bools[i] = applyComplex(op, a.complexes[a.physIndex(i)], z)
}
})
default:
if a.isContiguous() {
switch a.dt {
case Float:
floatCmpScalarRun(op, a.floats[:n], v, out.bools)
return out, nil
case Float32:
floatCmpScalarRun(op, a.floats32[:n], v, out.bools)
return out, nil
default:
scalarCmpFRun(op, a, v, out.bools)
return out, nil
}
}
parallelMin(n, elementwiseMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
out.bools[i] = applyFloat(op, a.floatAt(i), v)
}
})
}
return out, nil
}
// intCmpScalarFRun compares a dense int payload against a float scalar
// through the exact int-versus-float relation.
func intCmpScalarFRun(op cmpOp, x []int64, v float64, dst []bool) {
parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) {
xs, zs := x[s:e], dst[s:e]
for i := range zs {
zs[i] = applyIntFloat(op, xs[i], v)
}
})
}
// scalarCmpFRun compares a dense half, bool or narrow-integer payload
// against a float scalar, widening each element exactly as the accessor
// read would. The float64 and float32 widths carry their own op-hoisted
// kernels and never reach here.
func scalarCmpFRun(op cmpOp, a *Array, v float64, dst []bool) {
parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) {
zs := dst[s:e]
switch a.dt {
case Float16:
xs := a.halves[s:e]
for i := range zs {
zs[i] = applyFloat(op, HalfToFloat64(xs[i]), v)
}
case Bool:
xs := a.bools[s:e]
for i := range zs {
w := 0.0
if xs[i] {
w = 1
}
zs[i] = applyFloat(op, w, v)
}
case Int8:
xs := a.i8s[s:e]
for i := range zs {
zs[i] = applyFloat(op, float64(xs[i]), v)
}
case Uint8:
xs := a.u8s[s:e]
for i := range zs {
zs[i] = applyFloat(op, float64(xs[i]), v)
}
case Int16:
xs := a.i16s[s:e]
for i := range zs {
zs[i] = applyFloat(op, float64(xs[i]), v)
}
case Uint16:
xs := a.u16s[s:e]
for i := range zs {
zs[i] = applyFloat(op, float64(xs[i]), v)
}
case Int32:
xs := a.i32s[s:e]
for i := range zs {
zs[i] = applyFloat(op, float64(xs[i]), v)
}
default: // Uint32
xs := a.u32s[s:e]
for i := range zs {
zs[i] = applyFloat(op, float64(xs[i]), v)
}
}
})
}
// The comparison kernels below hold one loop per relation. The relation
// is a property of the call, not of the element, so the switch leaves
// the loop and the body holds the comparison itself; the destination of
// every element is its own slot, so the walk splits across workers like
// the arithmetic maps. Each kernel writes the same predicate the
// accessor walk computed, from the same operands.
// intCmpRun compares two dense int payloads, writing true where op
// holds.
func intCmpRun(op cmpOp, x, y []int64, dst []bool) {
parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) {
xs, ys, zs := x[s:e], y[s:e], dst[s:e]
switch op {
case opEq:
for i := range zs {
zs[i] = xs[i] == ys[i]
}
case opNe:
for i := range zs {
zs[i] = xs[i] != ys[i]
}
case opLt:
for i := range zs {
zs[i] = xs[i] < ys[i]
}
case opLe:
for i := range zs {
zs[i] = xs[i] <= ys[i]
}
case opGt:
for i := range zs {
zs[i] = xs[i] > ys[i]
}
default: // opGe
for i := range zs {
zs[i] = xs[i] >= ys[i]
}
}
})
}
// floatCmpRun compares two dense real payloads of one width. Both
// widenings are exact for the comparison, matching the accessor walk's
// values exactly, and NaN follows IEEE because the comparison is the
// same one.
func floatCmpRun[T float32 | float64](op cmpOp, x, y []T, dst []bool) {
parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) {
xs, ys, zs := x[s:e], y[s:e], dst[s:e]
switch op {
case opEq:
for i := range zs {
zs[i] = float64(xs[i]) == float64(ys[i])
}
case opNe:
for i := range zs {
zs[i] = float64(xs[i]) != float64(ys[i])
}
case opLt:
for i := range zs {
zs[i] = float64(xs[i]) < float64(ys[i])
}
case opLe:
for i := range zs {
zs[i] = float64(xs[i]) <= float64(ys[i])
}
case opGt:
for i := range zs {
zs[i] = float64(xs[i]) > float64(ys[i])
}
default: // opGe
for i := range zs {
zs[i] = float64(xs[i]) >= float64(ys[i])
}
}
})
}
// complexCmpRun compares two dense complex payloads; an ordering op never
// reaches here, so the walk holds either equality or its negation.
func complexCmpRun(op cmpOp, x, y []complex128, dst []bool) {
parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) {
xs, ys, zs := x[s:e], y[s:e], dst[s:e]
if op == opNe {
for i := range zs {
zs[i] = xs[i] != ys[i]
}
return
}
for i := range zs {
zs[i] = xs[i] == ys[i]
}
})
}
// intCmpScalarRun compares a dense int payload against an int scalar.
func intCmpScalarRun(op cmpOp, x []int64, v int64, dst []bool) {
parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) {
xs, zs := x[s:e], dst[s:e]
switch op {
case opEq:
for i := range zs {
zs[i] = xs[i] == v
}
case opNe:
for i := range zs {
zs[i] = xs[i] != v
}
case opLt:
for i := range zs {
zs[i] = xs[i] < v
}
case opLe:
for i := range zs {
zs[i] = xs[i] <= v
}
case opGt:
for i := range zs {
zs[i] = xs[i] > v
}
default: // opGe
for i := range zs {
zs[i] = xs[i] >= v
}
}
})
}
// floatCmpScalarRun compares a dense real payload of one width against a
// float scalar, widening exactly.
func floatCmpScalarRun[T float32 | float64](op cmpOp, x []T, v float64, dst []bool) {
parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) {
xs, zs := x[s:e], dst[s:e]
switch op {
case opEq:
for i := range zs {
zs[i] = float64(xs[i]) == v
}
case opNe:
for i := range zs {
zs[i] = float64(xs[i]) != v
}
case opLt:
for i := range zs {
zs[i] = float64(xs[i]) < v
}
case opLe:
for i := range zs {
zs[i] = float64(xs[i]) <= v
}
case opGt:
for i := range zs {
zs[i] = float64(xs[i]) > v
}
default: // opGe
for i := range zs {
zs[i] = float64(xs[i]) >= v
}
}
})
}
// narrowCmpRun compares two dense same-width narrow integer payloads,
// writing true where op holds. Every widening the accessor walk takes is
// exact, so the comparison in the payload's own type sees the accessor
// walk's values and answers with its predicate.
func narrowCmpRun[T int8 | uint8 | int16 | uint16 | int32 | uint32](op cmpOp, x, y []T, dst []bool) {
parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) {
xs, ys, zs := x[s:e], y[s:e], dst[s:e]
switch op {
case opEq:
for i := range zs {
zs[i] = xs[i] == ys[i]
}
case opNe:
for i := range zs {
zs[i] = xs[i] != ys[i]
}
case opLt:
for i := range zs {
zs[i] = xs[i] < ys[i]
}
case opLe:
for i := range zs {
zs[i] = xs[i] <= ys[i]
}
case opGt:
for i := range zs {
zs[i] = xs[i] > ys[i]
}
default: // opGe
for i := range zs {
zs[i] = xs[i] >= ys[i]
}
}
})
}
// boolCmpRun compares two dense bool payloads: false orders below true,
// the order the accessor walk's 0/1 widening carried.
func boolCmpRun(op cmpOp, x, y, dst []bool) {
parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) {
xs, ys, zs := x[s:e], y[s:e], dst[s:e]
switch op {
case opEq:
for i := range zs {
zs[i] = xs[i] == ys[i]
}
case opNe:
for i := range zs {
zs[i] = xs[i] != ys[i]
}
case opLt:
for i := range zs {
zs[i] = !xs[i] && ys[i]
}
case opLe:
for i := range zs {
zs[i] = !xs[i] || ys[i]
}
case opGt:
for i := range zs {
zs[i] = xs[i] && !ys[i]
}
default: // opGe
for i := range zs {
zs[i] = xs[i] || !ys[i]
}
}
})
}
// narrowCmpMixed compares two dense narrow integer payloads of
// different widths. Both widenings are exact, so the int64 comparison
// is the accessor walk's float64 comparison unchanged: no narrow value
// ever rounds.
func narrowCmpMixed(op cmpOp, a, b *Array, n int, dst []bool) {
switch a.dt {
case Int8:
narrowCmpMixedRun(op, a.i8s[:n], b, dst)
case Uint8:
narrowCmpMixedRun(op, a.u8s[:n], b, dst)
case Int16:
narrowCmpMixedRun(op, a.i16s[:n], b, dst)
case Uint16:
narrowCmpMixedRun(op, a.u16s[:n], b, dst)
case Int32:
narrowCmpMixedRun(op, a.i32s[:n], b, dst)
default: // Uint32
narrowCmpMixedRun(op, a.u32s[:n], b, dst)
}
}
// narrowCmpMixedRun is narrowCmpMixed's kernel: the left payload's width
// is fixed by the caller, the right side dispatches per dtype inside.
func narrowCmpMixedRun[A int8 | uint8 | int16 | uint16 | int32 | uint32](op cmpOp, xs []A, b *Array, dst []bool) {
parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) {
as := xs[s:e]
switch b.dt {
case Int8:
narrowCmpMixedLoop(op, as, b.i8s[s:e], dst[s:e])
case Uint8:
narrowCmpMixedLoop(op, as, b.u8s[s:e], dst[s:e])
case Int16:
narrowCmpMixedLoop(op, as, b.i16s[s:e], dst[s:e])
case Uint16:
narrowCmpMixedLoop(op, as, b.u16s[s:e], dst[s:e])
case Int32:
narrowCmpMixedLoop(op, as, b.i32s[s:e], dst[s:e])
default: // Uint32
narrowCmpMixedLoop(op, as, b.u32s[s:e], dst[s:e])
}
})
}
// narrowCmpMixedLoop holds one comparison loop per relation over two
// differently typed narrow payloads, both widened exactly into int64.
func narrowCmpMixedLoop[A, B int8 | uint8 | int16 | uint16 | int32 | uint32](op cmpOp, xs []A, ys []B, zs []bool) {
switch op {
case opEq:
for i := range zs {
zs[i] = int64(xs[i]) == int64(ys[i])
}
case opNe:
for i := range zs {
zs[i] = int64(xs[i]) != int64(ys[i])
}
case opLt:
for i := range zs {
zs[i] = int64(xs[i]) < int64(ys[i])
}
case opLe:
for i := range zs {
zs[i] = int64(xs[i]) <= int64(ys[i])
}
case opGt:
for i := range zs {
zs[i] = int64(xs[i]) > int64(ys[i])
}
default: // opGe
for i := range zs {
zs[i] = int64(xs[i]) >= int64(ys[i])
}
}
}
func whereRun[T any](cond []int64, x, y, dst []T) {
parallelMin(len(cond), elementwiseMinPerWorker, func(s, e int) {
cs, xs, ys, zs := cond[s:e], x[s:e], y[s:e], dst[s:e]
for i := range zs {
if cs[i] != 0 {
zs[i] = xs[i]
} else {
zs[i] = ys[i]
}
}
})
}
// whereHalfRun is whereRun for float16, which narrows through
// HalfFromFloat64 as the accessor walk does: a half NaN canonicalises on
// the way out, so the payload cannot be copied across directly.
func whereHalfRun(cond []int64, x, y, dst []uint16) {
parallelMin(len(cond), elementwiseMinPerWorker, func(s, e int) {
cs, xs, ys, zs := cond[s:e], x[s:e], y[s:e], dst[s:e]
for i := range zs {
v := ys[i]
if cs[i] != 0 {
v = xs[i]
}
zs[i] = HalfFromFloat64(HalfToFloat64(v))
}
})
}
// Eq returns a bool mask that is true where a equals b element-wise.
func Eq(a, b *Array) (*Array, error) { return a.cmpArray(b, "Eq", opEq) }
// Ne returns a bool mask that is true where a differs from b
// element-wise.
func Ne(a, b *Array) (*Array, error) { return a.cmpArray(b, "Ne", opNe) }
// Lt returns a bool mask that is true where a is below b element-wise.
func Lt(a, b *Array) (*Array, error) { return a.cmpArray(b, "Lt", opLt) }
// Le returns a bool mask that is true where a is at most b element-wise.
func Le(a, b *Array) (*Array, error) { return a.cmpArray(b, "Le", opLe) }
// Gt returns a bool mask that is true where a is above b element-wise.
func Gt(a, b *Array) (*Array, error) { return a.cmpArray(b, "Gt", opGt) }
// Ge returns a bool mask that is true where a is at least b element-wise.
func Ge(a, b *Array) (*Array, error) { return a.cmpArray(b, "Ge", opGe) }
// EqI is Eq against an int scalar.
func EqI(a *Array, v int64) (*Array, error) { return a.cmpScalarI(v, "Eq", opEq) }
// NeI is Ne against an int scalar.
func NeI(a *Array, v int64) (*Array, error) { return a.cmpScalarI(v, "Ne", opNe) }
// LtI is Lt against an int scalar.
func LtI(a *Array, v int64) (*Array, error) { return a.cmpScalarI(v, "Lt", opLt) }
// LeI is Le against an int scalar.
func LeI(a *Array, v int64) (*Array, error) { return a.cmpScalarI(v, "Le", opLe) }
// GtI is Gt against an int scalar.
func GtI(a *Array, v int64) (*Array, error) { return a.cmpScalarI(v, "Gt", opGt) }
// GeI is Ge against an int scalar.
func GeI(a *Array, v int64) (*Array, error) { return a.cmpScalarI(v, "Ge", opGe) }
// EqF is Eq against a float scalar.
func EqF(a *Array, v float64) (*Array, error) { return a.cmpScalarF(v, "Eq", opEq) }
// NeF is Ne against a float scalar.
func NeF(a *Array, v float64) (*Array, error) { return a.cmpScalarF(v, "Ne", opNe) }
// LtF is Lt against a float scalar.
func LtF(a *Array, v float64) (*Array, error) { return a.cmpScalarF(v, "Lt", opLt) }
// LeF is Le against a float scalar.
func LeF(a *Array, v float64) (*Array, error) { return a.cmpScalarF(v, "Le", opLe) }
// GtF is Gt against a float scalar.
func GtF(a *Array, v float64) (*Array, error) { return a.cmpScalarF(v, "Gt", opGt) }
// GeF is Ge against a float scalar.
func GeF(a *Array, v float64) (*Array, error) { return a.cmpScalarF(v, "Ge", opGe) }
// Select returns a 1-D copy of the elements where the mask m, of the
// same shape, is set: a bool mask selects where it reads true, an int
// mask where it reads nonzero. Renamed from `Mask` because the
// operation is element selection, not bitmasking; the mask is just a
// selector array, bool natively or the int 0/1 the comparisons answered
// before the bool dtype carried them.
func Select(a, m *Array) (*Array, error) {
if m.dt != Bool && m.dt != Int {
return nil, errf("Select: the mask must be a bool or int array, got %s", m.dt)
}
if !sameShape(a.shape, m.shape) {
return nil, errf("Select: shape mismatch %s vs %s", shapeText(a.shape), shapeText(m.shape))
}
n := m.Len()
hits := 0
if m.dt == Bool {
for i := range n {
if m.bools[i] {
hits++
}
}
} else {
for i := range n {
if m.ints[i] != 0 {
hits++
}
}
}
out := &Array{shape: []int{hits}, dt: a.dt}
out.alloc(hits)
if hits == 0 {
return out, nil
}
if !a.isContiguous() {
// A view has no aligned payload, so the gather keeps the
// accessor per hit.
k := 0
if m.dt == Bool {
for i := range n {
if m.bools[i] {
out.setFrom(k, a, i)
k++
}
}
} else {
for i := range n {
if m.ints[i] != 0 {
out.setFrom(k, a, i)
k++
}
}
}
return out, nil
}
switch a.dt {
case Int:
gatherMasked(out.ints, a.ints[:n], m, n)
case Bool:
gatherMasked(out.bools, a.bools[:n], m, n)
case Int8:
gatherMasked(out.i8s, a.i8s[:n], m, n)
case Uint8:
gatherMasked(out.u8s, a.u8s[:n], m, n)
case Int16:
gatherMasked(out.i16s, a.i16s[:n], m, n)
case Uint16:
gatherMasked(out.u16s, a.u16s[:n], m, n)
case Int32:
gatherMasked(out.i32s, a.i32s[:n], m, n)
case Uint32:
gatherMasked(out.u32s, a.u32s[:n], m, n)
case Float16:
gatherMasked(out.halves, a.halves[:n], m, n)
case Float32:
gatherMasked(out.floats32, a.floats32[:n], m, n)
case Float:
gatherMasked(out.floats, a.floats[:n], m, n)
default:
gatherMasked(out.complexes, a.complexes[:n], m, n)
}
return out, nil
}
// gatherMasked runs the compact gather for one payload width, reading
// the mask as bool when it is one and as a nonzero int otherwise.
func gatherMasked[T any](dst, src []T, m *Array, n int) {
if m.dt == Bool {
selectCompactBool(dst, src, m.bools[:n])
return
}
selectCompact(dst, src, m.ints[:n])
}
// selectPlan splits n elements into per-worker chunks and folds each
// chunk's hit count into the running offsets the gathers write from,
// the schedule both compact walks share. The split answer is false when
// the walk is too small to split, and the caller gathers serially.
func selectPlan(n int, hits func(lo, hi int) int) (chunk int, base []int, split bool) {
w := workersFor(n)
chunk = (n + w - 1) / w
if w == 1 || chunk < elementwiseMinPerWorker {
return chunk, nil, false
}
nch := (n + chunk - 1) / chunk
base = make([]int, nch)
var wg sync.WaitGroup
for c := range nch {
wg.Go(func() {
base[c] = hits(c*chunk, min((c+1)*chunk, n))
})
}
wg.Wait()
off := 0
for c := range nch {
k := base[c]
base[c] = off
off += k
}
return chunk, base, true
}
// selectCompact gathers the elements of src whose int mask slot is
// nonzero into dst in ascending order. Each chunk counts its hits first,
// so the prefix sum hands every worker its own destination offset and
// the writes stay output-disjoint: the same count-then-gather shape the
// serial walk had, with the scan split across workers.
func selectCompact[T any](dst, src []T, mask []int64) {
chunk, base, split := selectPlan(len(mask), func(lo, hi int) int {
k := 0
for i := lo; i < hi; i++ {
if mask[i] != 0 {
k++
}
}
return k
})
if !split {
k := 0
for i := range mask {
if mask[i] != 0 {
dst[k] = src[i]
k++
}
}
return
}
n := len(mask)
var wg sync.WaitGroup
for c := range base {
wg.Go(func() {
lo, hi := c*chunk, min((c+1)*chunk, n)
k := base[c]
for i := lo; i < hi; i++ {
if mask[i] != 0 {
dst[k] = src[i]
k++
}
}
})
}
wg.Wait()
}
// selectCompactBool gathers the elements of src whose bool mask slot is
// true, with selectCompact's contract and schedule.
func selectCompactBool[T any](dst, src []T, mask []bool) {
chunk, base, split := selectPlan(len(mask), func(lo, hi int) int {
k := 0
for i := lo; i < hi; i++ {
if mask[i] {
k++
}
}
return k
})
if !split {
k := 0
for i := range mask {
if mask[i] {
dst[k] = src[i]
k++
}
}
return
}
n := len(mask)
var wg sync.WaitGroup
for c := range base {
wg.Go(func() {
lo, hi := c*chunk, min((c+1)*chunk, n)
k := base[c]
for i := lo; i < hi; i++ {
if mask[i] {
dst[k] = src[i]
k++
}
}
})
}
wg.Wait()
}
// Where selects element-wise between x and y: a set element of cond
// picks x, an unset one picks y. cond is a bool mask or an int array,
// whose nonzero elements stand for true; every other dtype is refused
// with the wording the tests pin. All three arrays must share a shape;
// the result dtype is promote(x, y) over the full ladder.
func Where(cond, x, y *Array) (*Array, error) {
if cond.dt != Int && cond.dt != Bool {
return nil, errf("Where: the condition must be a bool or int array, got %s", cond.dt)
}
if !sameShape(cond.shape, x.shape) || !sameShape(cond.shape, y.shape) {
return nil, errf("Where: shapes %s, %s and %s must agree",
shapeText(cond.shape), shapeText(x.shape), shapeText(y.shape))
}
n := cond.Len()
out := &Array{shape: x.Shape(), dt: promote(x.dt, y.dt)}
out.alloc(n)
// A dense int-masked result whose dtype is each operand's own dtype
// picks payload slots directly; a bool condition, a mixed pair, a
// narrow float16 result and a view keep the accessor walk below,
// which reads the same values.
if cond.dt == Int && cond.isContiguous() && x.isContiguous() && y.isContiguous() && x.dt == out.dt && y.dt == out.dt {
cs := cond.ints[:n]
switch out.dt {
case Int:
whereRun(cs, x.ints[:n], y.ints[:n], out.ints)
return out, nil
case Bool:
whereRun(cs, x.bools[:n], y.bools[:n], out.bools)
return out, nil
case Int8:
whereRun(cs, x.i8s[:n], y.i8s[:n], out.i8s)
return out, nil
case Uint8:
whereRun(cs, x.u8s[:n], y.u8s[:n], out.u8s)
return out, nil
case Int16:
whereRun(cs, x.i16s[:n], y.i16s[:n], out.i16s)
return out, nil
case Uint16:
whereRun(cs, x.u16s[:n], y.u16s[:n], out.u16s)
return out, nil
case Int32:
whereRun(cs, x.i32s[:n], y.i32s[:n], out.i32s)
return out, nil
case Uint32:
whereRun(cs, x.u32s[:n], y.u32s[:n], out.u32s)
return out, nil
case Float16:
// The half result narrows through HalfFromFloat64 as the
// accessor walk does, so a payload copy would not do.
whereHalfRun(cs, x.halves[:n], y.halves[:n], out.halves)
return out, nil
case Float32:
whereRun(cs, x.floats32[:n], y.floats32[:n], out.floats32)
return out, nil
case Float:
whereRun(cs, x.floats[:n], y.floats[:n], out.floats)
return out, nil
default:
whereRun(cs, x.complexes[:n], y.complexes[:n], out.complexes)
return out, nil
}
}
parallelMin(n, elementwiseMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
var nz bool
if cond.dt == Bool {
nz = cond.boolAt(i)
} else {
nz = cond.ints[i] != 0
}
src := y
if nz {
src = x
}
switch out.dt {
case Int:
// intAt widens a narrow or bool operand exactly: a mixed
// integer-class pair promotes to int and both sides
// convert on the way in.
out.ints[i] = src.intAt(i)
case Bool:
out.bools[i] = src.boolAt(i)
case Int8:
out.i8s[i] = int8(src.intAt(i))
case Uint8:
out.u8s[i] = uint8(src.intAt(i))
case Int16:
out.i16s[i] = int16(src.intAt(i))
case Uint16:
out.u16s[i] = uint16(src.intAt(i))
case Int32:
out.i32s[i] = int32(src.intAt(i))
case Uint32:
out.u32s[i] = uint32(src.intAt(i))
case Float16:
// A half result reads either operand through the exact
// widening and narrows once.
out.halves[i] = HalfFromFloat64(src.floatAt(i))
case Float32:
out.floats32[i] = src.float32At(i)
case Float:
out.floats[i] = src.floatAt(i)
default:
out.complexes[i] = src.complexAt(i)
}
}
})
return out, nil
}
// logicOp selects the boolean connective a logic operation applies.
type logicOp uint8
const (
logicAnd logicOp = iota
logicOr
logicXor
)
// And returns the element-wise conjunction of two bool arrays; both
// operands must be bool.
func And(a, b *Array) (*Array, error) { return logicArray(a, b, "And", logicAnd) }
// Or returns the element-wise disjunction of two bool arrays; both
// operands must be bool.
func Or(a, b *Array) (*Array, error) { return logicArray(a, b, "Or", logicOr) }
// Xor returns the element-wise exclusive disjunction of two bool
// arrays; both operands must be bool.
func Xor(a, b *Array) (*Array, error) { return logicArray(a, b, "Xor", logicXor) }
// Not returns the element-wise negation of a bool array; the operand
// must be bool.
func Not(a *Array) (*Array, error) {
if a.dt != Bool {
return nil, errf("Not: operands must be bool arrays, got %s", a.dt)
}
n := a.Len()
out := &Array{shape: a.Shape(), dt: Bool}
out.alloc(n)
if a.isContiguous() {
src := a.bools
parallelMin(n, elementwiseMinPerWorker, func(s, e int) {
xs, os := src[s:e], out.bools[s:e]
for i := range os {
os[i] = !xs[i]
}
})
return out, nil
}
parallelMin(n, elementwiseMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
out.bools[i] = !a.boolAt(i)
}
})
return out, nil
}
// logicArray runs the binary logic connectives: the only arithmetic-like
// surface the bool dtype carries, answered as a bool array.
func logicArray(a, b *Array, name string, op logicOp) (*Array, error) {
if a.dt != Bool || b.dt != Bool {
return nil, errf("%s: operands must be bool arrays, got %s and %s", name, a.dt, b.dt)
}
if !sameShape(a.shape, b.shape) {
return nil, errf("%s: shape mismatch %s vs %s", name, shapeText(a.shape), shapeText(b.shape))
}
n := a.Len()
out := &Array{shape: a.Shape(), dt: Bool}
out.alloc(n)
if a.isContiguous() && b.isContiguous() {
boolLogicRun(op, a.bools, b.bools, out.bools)
return out, nil
}
parallelMin(n, elementwiseMinPerWorker, func(s, e int) {
for i := s; i < e; i++ {
x, y := a.boolAt(i), b.boolAt(i)
switch op {
case logicAnd:
out.bools[i] = x && y
case logicOr:
out.bools[i] = x || y
default: // logicXor
out.bools[i] = x != y
}
}
})
return out, nil
}
// boolLogicRun applies one connective over two dense bool payloads, the
// relation a property of the call so the switch leaves the element loop.
func boolLogicRun(op logicOp, x, y, dst []bool) {
parallelMin(len(dst), elementwiseMinPerWorker, func(s, e int) {
xs, ys, os := x[s:e], y[s:e], dst[s:e]
switch op {
case logicAnd:
for i := range os {
os[i] = xs[i] && ys[i]
}
case logicOr:
for i := range os {
os[i] = xs[i] || ys[i]
}
default: // logicXor
for i := range os {
os[i] = xs[i] != ys[i]
}
}
})
}