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