1284 lines
35 KiB
Go
1284 lines
35 KiB
Go
// 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]
|
|
}
|
|
}
|
|
})
|
|
}
|