Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 8 additions & 0 deletions encodings/runend/src/trace_tests.rs
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,14 @@ fn trace_compare_on_runend() -> VortexResult<()> {
iter 0 current=vortex.runend(bool, len=9) builder_active=false
execute_until target=AnyCanonical root=vortex.binary(bool, len=3)
iter 0 current=vortex.binary(bool, len=3) builder_active=false
optimize root=vortex.slice(i32, len=1) session=false
reduce_parent static:SliceReduceAdaptor(Constant) slot=0 parent=vortex.slice(i32, len=1) child=vortex.constant(i32, len=3) -> vortex.constant(i32, len=1)
done output=vortex.constant(i32, len=1)
execute_until target=AnyCanonical root=vortex.constant(i32, len=1)
iter 0 current=vortex.constant(i32, len=1) builder_active=false
Done array=vortex.primitive(i32, len=1)
iter 1 current=vortex.primitive(i32, len=1) builder_active=false
return output=vortex.primitive(i32, len=1)
Done array=vortex.bool(bool, len=3)
iter 1 current=vortex.bool(bool, len=3) builder_active=false
return output=vortex.bool(bool, len=3)
Expand Down
8 changes: 8 additions & 0 deletions vortex-array/benches/binary_ops.rs
Original file line number Diff line number Diff line change
Expand Up @@ -170,6 +170,14 @@ fn div_i64_nonnull(bencher: Bencher) {
bench_primitive(bencher, lhs, rhs, Operator::Div);
}

#[divan::bench]
fn div_i64_nullable(bencher: Bencher) {
let lhs = primitive_nullable(1_000_000, 7).into_array();
let rhs = primitive_nullable(17, 5).into_array();

bench_primitive(bencher, lhs, rhs, Operator::Div);
}

#[divan::bench]
fn sub_i64_constant(bencher: Bencher) {
let lhs = primitive_nonnull(0).into_array();
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,7 @@ use vortex_error::vortex_err;

use crate::scalar_fn::fns::operators::Operator;

#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
/// Binary element-wise operations.
pub enum NumericOperator {
/// Binary element-wise addition of two arrays or of two scalars.
Expand Down
2 changes: 1 addition & 1 deletion vortex-array/src/scalar_fn/fns/binary/compare/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -211,7 +211,7 @@ fn compare_arrays(
)
.into_array()),
DType::Bool(_) => boolean::compare_bool(lhs, rhs, op, nullability, ctx),
DType::Primitive(..) => primitive::compare_primitive(lhs, rhs, op, nullability, ctx),
DType::Primitive(..) => primitive::compare_primitive(lhs, rhs, op, ctx),
DType::Decimal(..) => decimal::compare_decimal(lhs, rhs, op, nullability, ctx),
DType::Utf8(_) | DType::Binary(_) => bytes::compare_bytes(lhs, rhs, op, nullability, ctx),
DType::Struct(..) | DType::List(..) | DType::FixedSizeList(..) => {
Expand Down
164 changes: 72 additions & 92 deletions vortex-array/src/scalar_fn/fns/binary/compare/primitive.rs
Original file line number Diff line number Diff line change
@@ -1,27 +1,28 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright the Vortex contributors

//! Native comparison of primitive arrays via bit-packing lane kernels.
//! Primitive comparison execution through [`RowFn`].

#[cfg(target_arch = "x86_64")]
mod columnar;
#[cfg(target_arch = "x86_64")]
mod operand;

use vortex_buffer::BitBuffer;
use vortex_error::VortexResult;
use vortex_error::vortex_bail;
use vortex_error::vortex_err;

use crate::ArrayRef;
use crate::ExecutionCtx;
use crate::IntoArray;
use crate::arrays::BoolArray;
use crate::arrays::ConstantArray;
use crate::dtype::DType;
use crate::dtype::NativePType;
use crate::dtype::Nullability;
use crate::dtype::PType;
use crate::match_each_native_ptype;
use crate::scalar::Scalar;
use crate::scalar_fn::fns::binary::compare::collect_bits;
use crate::scalar_fn::fns::binary::compare::collect_zip_bits;
use crate::scalar_fn::fns::binary::compare::compare_validity;
use crate::scalar_fn::fns::binary::primitive_operand::PrimitiveOperand;
use crate::scalar_fn::RowFn;
use crate::scalar_fn::RowVisitor;
use crate::scalar_fn::ScalarFnId;
use crate::scalar_fn::ScalarFnVTable;
use crate::scalar_fn::VecExecutionArgs;
use crate::scalar_fn::fns::binary::Binary;
use crate::scalar_fn::fns::operators::CompareOperator;

/// Compare two primitive arrays of the same [`PType`].
Expand All @@ -32,99 +33,78 @@ pub(super) fn compare_primitive(
lhs: &ArrayRef,
rhs: &ArrayRef,
op: CompareOperator,
nullability: Nullability,
ctx: &mut ExecutionCtx,
) -> VortexResult<ArrayRef> {
let ptype = PType::try_from(lhs.dtype())?;
match_each_native_ptype!(ptype, |T| {
compare_primitive_typed::<T>(lhs, rhs, op, nullability, ctx)
})
#[cfg(target_arch = "x86_64")]
if use_columnar_comparison(lhs, rhs, op)? {
return columnar::compare_primitive(lhs, rhs, op, ctx);
}

let args = VecExecutionArgs::new(vec![lhs.clone(), rhs.clone()], lhs.len());

ScalarFnVTable::execute(&PrimitiveCompare, &op, &args, ctx)
}

fn compare_primitive_typed<T: NativePType>(
lhs: &ArrayRef,
rhs: &ArrayRef,
op: CompareOperator,
nullability: Nullability,
ctx: &mut ExecutionCtx,
) -> VortexResult<ArrayRef> {
let len = lhs.len();
let lhs = PrimitiveOperand::<T>::try_new(lhs, ctx)?;
let rhs = PrimitiveOperand::<T>::try_new(rhs, ctx)?;
if lhs.len() != rhs.len() {
vortex_bail!(
"compare operator requires equal lengths, got {} and {}",
lhs.len(),
rhs.len()
);
/// Internal row execution for primitive comparison operators.
#[derive(Clone)]
struct PrimitiveCompare;

impl RowFn for PrimitiveCompare {
type Options = CompareOperator;

const ARG_NAMES: &'static [&'static str] = &["lhs", "rhs"];

fn id(&self) -> ScalarFnId {
ScalarFnVTable::id(&Binary)
}

let validity = compare_validity(lhs.validity(), rhs.validity(), nullability)?;

let bits = match (&lhs, &rhs) {
(
PrimitiveOperand::Array { values: lhs, .. },
PrimitiveOperand::Array { values: rhs, .. },
) => compare_slices(lhs, rhs, op),
(
PrimitiveOperand::Array { values: lhs, .. },
PrimitiveOperand::Constant { value: rhs, .. },
) => compare_slice_constant(lhs, *rhs, op),
(
PrimitiveOperand::Constant { value: lhs, .. },
PrimitiveOperand::Array { values: rhs, .. },
) => compare_slice_constant(rhs, *lhs, op.swap()),
(
PrimitiveOperand::Constant { value: lhs, .. },
PrimitiveOperand::Constant { value: rhs, .. },
) => {
// Unreachable through `execute_compare` (constant-constant is folded there), but
// cheap to answer anyway.
BitBuffer::full(apply_op(*lhs, *rhs, op), len)
}
(PrimitiveOperand::Null(_), _) | (_, PrimitiveOperand::Null(_)) => {
return Ok(
ConstantArray::new(Scalar::null(DType::Bool(Nullability::Nullable)), len)
.into_array(),
);
}
};

Ok(BoolArray::try_new(bits, validity)?.into_array())
}
fn dispatch<V: RowVisitor>(
&self,
op: &Self::Options,
args: &[DType],
visitor: V,
) -> VortexResult<V::VisitResult> {
let ptype =
PType::try_from(args.first().ok_or_else(|| {
vortex_err!("a comparison operator takes two operands, got none")
})?)?;

#[inline(always)]
fn apply_op<T: NativePType>(lhs: T, rhs: T, op: CompareOperator) -> bool {
match op {
CompareOperator::Eq => lhs.is_eq(rhs),
CompareOperator::NotEq => !lhs.is_eq(rhs),
CompareOperator::Gt => lhs.is_gt(rhs),
CompareOperator::Gte => lhs.is_ge(rhs),
CompareOperator::Lt => lhs.is_lt(rhs),
CompareOperator::Lte => lhs.is_le(rhs),
match_each_native_ptype!(ptype, |T| { visit_compare::<T, V>(*op, visitor) })
}
}

fn compare_slices<T: NativePType>(lhs: &[T], rhs: &[T], op: CompareOperator) -> BitBuffer {
// Dispatch the operator outside the lane loop so each instantiation vectorizes a single
// branch-free predicate.
match op {
CompareOperator::Eq => collect_zip_bits(lhs, rhs, |a: T, b: T| a.is_eq(b)),
CompareOperator::NotEq => collect_zip_bits(lhs, rhs, |a: T, b: T| !a.is_eq(b)),
CompareOperator::Gt => collect_zip_bits(lhs, rhs, T::is_gt),
CompareOperator::Gte => collect_zip_bits(lhs, rhs, T::is_ge),
CompareOperator::Lt => collect_zip_bits(lhs, rhs, T::is_lt),
CompareOperator::Lte => collect_zip_bits(lhs, rhs, T::is_le),
#[cfg(target_arch = "x86_64")]
fn use_columnar_comparison(
lhs: &ArrayRef,
rhs: &ArrayRef,
op: CompareOperator,
) -> VortexResult<bool> {
if matches!(op, CompareOperator::Eq | CompareOperator::NotEq) {
return Ok(false);
}

let ptype = PType::try_from(lhs.dtype())?;
Ok(match ptype {
// The fused comparison and bit-packing loop produces better x86 code for signed 64-bit
// integers and f64. The RowFn byte-output loop remains faster for narrower lanes.
PType::I64 | PType::F64 => true,
// LLVM vectorizes varying u64 inputs, but not the mixed-constant RowFn loop.
PType::U64 => lhs.as_constant().is_some() || rhs.as_constant().is_some(),
_ => false,
})
}

fn compare_slice_constant<T: NativePType>(lhs: &[T], rhs: T, op: CompareOperator) -> BitBuffer {
fn visit_compare<T, V>(op: CompareOperator, visitor: V) -> VortexResult<V::VisitResult>
where
T: NativePType,
V: RowVisitor,
{
match op {
CompareOperator::Eq => collect_bits(lhs, |a: T| a.is_eq(rhs)),
CompareOperator::NotEq => collect_bits(lhs, |a: T| !a.is_eq(rhs)),
CompareOperator::Gt => collect_bits(lhs, |a: T| a.is_gt(rhs)),
CompareOperator::Gte => collect_bits(lhs, |a: T| a.is_ge(rhs)),
CompareOperator::Lt => collect_bits(lhs, |a: T| a.is_lt(rhs)),
CompareOperator::Lte => collect_bits(lhs, |a: T| a.is_le(rhs)),
CompareOperator::Eq => visitor.visit::<(T, T), bool>(|(lhs, rhs)| lhs.is_eq(rhs)),
CompareOperator::NotEq => visitor.visit::<(T, T), bool>(|(lhs, rhs)| !lhs.is_eq(rhs)),
CompareOperator::Gt => visitor.visit::<(T, T), bool>(|(lhs, rhs)| lhs.is_gt(rhs)),
CompareOperator::Gte => visitor.visit::<(T, T), bool>(|(lhs, rhs)| lhs.is_ge(rhs)),
CompareOperator::Lt => visitor.visit::<(T, T), bool>(|(lhs, rhs)| lhs.is_lt(rhs)),
CompareOperator::Lte => visitor.visit::<(T, T), bool>(|(lhs, rhs)| lhs.is_le(rhs)),
}
}
120 changes: 120 additions & 0 deletions vortex-array/src/scalar_fn/fns/binary/compare/primitive/columnar.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,120 @@
// SPDX-License-Identifier: Apache-2.0
// SPDX-FileCopyrightText: Copyright the Vortex contributors

//! Fused comparison and bit-packing for wide x86 lanes.

use vortex_buffer::BitBuffer;
use vortex_error::VortexResult;
use vortex_error::vortex_bail;

use super::operand::PrimitiveOperand;
use crate::ArrayRef;
use crate::ExecutionCtx;
use crate::IntoArray;
use crate::arrays::BoolArray;
use crate::arrays::ConstantArray;
use crate::dtype::DType;
use crate::dtype::NativePType;
use crate::dtype::Nullability;
use crate::dtype::PType;
use crate::scalar::Scalar;
use crate::scalar_fn::fns::binary::compare::collect_bits;
use crate::scalar_fn::fns::binary::compare::collect_zip_bits;
use crate::scalar_fn::fns::binary::compare::compare_validity;
use crate::scalar_fn::fns::operators::CompareOperator;

/// Compare primitive operands with one fused comparison and bit-packing loop.
pub(super) fn compare_primitive(
lhs: &ArrayRef,
rhs: &ArrayRef,
op: CompareOperator,
ctx: &mut ExecutionCtx,
) -> VortexResult<ArrayRef> {
match PType::try_from(lhs.dtype())? {
PType::I64 => compare_primitive_typed::<i64>(lhs, rhs, op, ctx),
PType::U64 => compare_primitive_typed::<u64>(lhs, rhs, op, ctx),
PType::F64 => compare_primitive_typed::<f64>(lhs, rhs, op, ctx),
ptype => vortex_bail!("columnar comparison is not selected for {ptype}"),
}
}

fn compare_primitive_typed<T: NativePType>(
lhs: &ArrayRef,
rhs: &ArrayRef,
op: CompareOperator,
ctx: &mut ExecutionCtx,
) -> VortexResult<ArrayRef> {
let len = lhs.len();
let nullability = Nullability::from(lhs.dtype().is_nullable() || rhs.dtype().is_nullable());
let lhs = PrimitiveOperand::<T>::try_new(lhs, ctx)?;
let rhs = PrimitiveOperand::<T>::try_new(rhs, ctx)?;
if lhs.len() != rhs.len() {
vortex_bail!(
"compare operator requires equal lengths, got {} and {}",
lhs.len(),
rhs.len()
);
}

let validity = compare_validity(lhs.validity(), rhs.validity(), nullability)?;
let bits = match (&lhs, &rhs) {
(
PrimitiveOperand::Array { values: lhs, .. },
PrimitiveOperand::Array { values: rhs, .. },
) => compare_slices(lhs, rhs, op),
(
PrimitiveOperand::Array { values: lhs, .. },
PrimitiveOperand::Constant { value: rhs, .. },
) => compare_slice_constant(lhs, *rhs, op),
(
PrimitiveOperand::Constant { value: lhs, .. },
PrimitiveOperand::Array { values: rhs, .. },
) => compare_slice_constant(rhs, *lhs, op.swap()),
(
PrimitiveOperand::Constant { value: lhs, .. },
PrimitiveOperand::Constant { value: rhs, .. },
) => BitBuffer::full(apply_op(*lhs, *rhs, op), len),
(PrimitiveOperand::Null(_), _) | (_, PrimitiveOperand::Null(_)) => {
return Ok(
ConstantArray::new(Scalar::null(DType::Bool(Nullability::Nullable)), len)
.into_array(),
);
}
};

Ok(BoolArray::try_new(bits, validity)?.into_array())
}

#[inline(always)]
fn apply_op<T: NativePType>(lhs: T, rhs: T, op: CompareOperator) -> bool {
match op {
CompareOperator::Eq => lhs.is_eq(rhs),
CompareOperator::NotEq => !lhs.is_eq(rhs),
CompareOperator::Gt => lhs.is_gt(rhs),
CompareOperator::Gte => lhs.is_ge(rhs),
CompareOperator::Lt => lhs.is_lt(rhs),
CompareOperator::Lte => lhs.is_le(rhs),
}
}

fn compare_slices<T: NativePType>(lhs: &[T], rhs: &[T], op: CompareOperator) -> BitBuffer {
match op {
CompareOperator::Eq => collect_zip_bits(lhs, rhs, |lhs: T, rhs: T| lhs.is_eq(rhs)),
CompareOperator::NotEq => collect_zip_bits(lhs, rhs, |lhs: T, rhs: T| !lhs.is_eq(rhs)),
CompareOperator::Gt => collect_zip_bits(lhs, rhs, T::is_gt),
CompareOperator::Gte => collect_zip_bits(lhs, rhs, T::is_ge),
CompareOperator::Lt => collect_zip_bits(lhs, rhs, T::is_lt),
CompareOperator::Lte => collect_zip_bits(lhs, rhs, T::is_le),
}
}

fn compare_slice_constant<T: NativePType>(lhs: &[T], rhs: T, op: CompareOperator) -> BitBuffer {
match op {
CompareOperator::Eq => collect_bits(lhs, |lhs: T| lhs.is_eq(rhs)),
CompareOperator::NotEq => collect_bits(lhs, |lhs: T| !lhs.is_eq(rhs)),
CompareOperator::Gt => collect_bits(lhs, |lhs: T| lhs.is_gt(rhs)),
CompareOperator::Gte => collect_bits(lhs, |lhs: T| lhs.is_ge(rhs)),
CompareOperator::Lt => collect_bits(lhs, |lhs: T| lhs.is_lt(rhs)),
CompareOperator::Lte => collect_bits(lhs, |lhs: T| lhs.is_le(rhs)),
}
}
Loading
Loading