Vendor dependencies

This commit is contained in:
2026-08-01 16:11:49 +03:00
parent 7f139a0241
commit 6b5e7f0f8b
29706 changed files with 9575646 additions and 0 deletions
+818
View File
@@ -0,0 +1,818 @@
// Licensed to the Apache Software Foundation (ASF) under one
// or more contributor license agreements. See the NOTICE file
// distributed with this work for additional information
// regarding copyright ownership. The ASF licenses this file
// to you under the Apache License, Version 2.0 (the
// "License"); you may not use this file except in compliance
// with the License. You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing,
// software distributed under the License is distributed on an
// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
// KIND, either express or implied. See the License for the
// specific language governing permissions and limitations
// under the License.
//! Comparison kernels for `Array`s.
//!
//! These kernels can leverage SIMD if available on your system. Currently no runtime
//! detection is provided, you should enable the specific SIMD intrinsics using
//! `RUSTFLAGS="-C target-feature=+avx2"` for example. See the documentation
//! [here](https://doc.rust-lang.org/stable/core/arch/) for more information.
//!
use arrow_array::cast::AsArray;
use arrow_array::types::{ByteArrayType, ByteViewType};
use arrow_array::{
AnyDictionaryArray, Array, ArrowNativeTypeOp, BooleanArray, Datum, FixedSizeBinaryArray,
GenericByteArray, GenericByteViewArray, downcast_primitive_array,
};
use arrow_buffer::bit_util::ceil;
use arrow_buffer::{BooleanBuffer, NullBuffer};
use arrow_schema::ArrowError;
use arrow_select::take::take;
use std::cmp::Ordering;
use std::ops::Not;
#[derive(Debug, Copy, Clone)]
enum Op {
Equal,
NotEqual,
Less,
LessEqual,
Greater,
GreaterEqual,
Distinct,
NotDistinct,
}
impl std::fmt::Display for Op {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Op::Equal => write!(f, "=="),
Op::NotEqual => write!(f, "!="),
Op::Less => write!(f, "<"),
Op::LessEqual => write!(f, "<="),
Op::Greater => write!(f, ">"),
Op::GreaterEqual => write!(f, ">="),
Op::Distinct => write!(f, "IS DISTINCT FROM"),
Op::NotDistinct => write!(f, "IS NOT DISTINCT FROM"),
}
}
}
/// Perform `left == right` operation on two [`Datum`].
///
/// Comparing null values on either side will yield a null in the corresponding
/// slot of the resulting [`BooleanArray`].
///
/// For floating values like f32 and f64, this comparison produces an ordering in accordance to
/// the totalOrder predicate as defined in the IEEE 754 (2008 revision) floating point standard.
/// Note that totalOrder treats positive and negative zeros as different. If it is necessary
/// to treat them as equal, please normalize zeros before calling this kernel. See
/// [`f32::total_cmp`] and [`f64::total_cmp`].
///
/// Nested types, such as lists, are not supported as the null semantics are not well-defined.
/// For comparisons involving nested types see [`crate::ord::make_comparator`]
pub fn eq(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<BooleanArray, ArrowError> {
compare_op(Op::Equal, lhs, rhs)
}
/// Perform `left != right` operation on two [`Datum`].
///
/// Comparing null values on either side will yield a null in the corresponding
/// slot of the resulting [`BooleanArray`].
///
/// For floating values like f32 and f64, this comparison produces an ordering in accordance to
/// the totalOrder predicate as defined in the IEEE 754 (2008 revision) floating point standard.
/// Note that totalOrder treats positive and negative zeros as different. If it is necessary
/// to treat them as equal, please normalize zeros before calling this kernel. See
/// [`f32::total_cmp`] and [`f64::total_cmp`].
///
/// Nested types, such as lists, are not supported as the null semantics are not well-defined.
/// For comparisons involving nested types see [`crate::ord::make_comparator`]
pub fn neq(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<BooleanArray, ArrowError> {
compare_op(Op::NotEqual, lhs, rhs)
}
/// Perform `left < right` operation on two [`Datum`].
///
/// Comparing null values on either side will yield a null in the corresponding
/// slot of the resulting [`BooleanArray`].
///
/// For floating values like f32 and f64, this comparison produces an ordering in accordance to
/// the totalOrder predicate as defined in the IEEE 754 (2008 revision) floating point standard.
/// Note that totalOrder treats positive and negative zeros as different. If it is necessary
/// to treat them as equal, please normalize zeros before calling this kernel. See
/// [`f32::total_cmp`] and [`f64::total_cmp`].
///
/// Nested types, such as lists, are not supported as the null semantics are not well-defined.
/// For comparisons involving nested types see [`crate::ord::make_comparator`]
pub fn lt(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<BooleanArray, ArrowError> {
compare_op(Op::Less, lhs, rhs)
}
/// Perform `left <= right` operation on two [`Datum`].
///
/// Comparing null values on either side will yield a null in the corresponding
/// slot of the resulting [`BooleanArray`].
///
/// For floating values like f32 and f64, this comparison produces an ordering in accordance to
/// the totalOrder predicate as defined in the IEEE 754 (2008 revision) floating point standard.
/// Note that totalOrder treats positive and negative zeros as different. If it is necessary
/// to treat them as equal, please normalize zeros before calling this kernel. See
/// [`f32::total_cmp`] and [`f64::total_cmp`].
///
/// Nested types, such as lists, are not supported as the null semantics are not well-defined.
/// For comparisons involving nested types see [`crate::ord::make_comparator`]
pub fn lt_eq(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<BooleanArray, ArrowError> {
compare_op(Op::LessEqual, lhs, rhs)
}
/// Perform `left > right` operation on two [`Datum`].
///
/// Comparing null values on either side will yield a null in the corresponding
/// slot of the resulting [`BooleanArray`].
///
/// For floating values like f32 and f64, this comparison produces an ordering in accordance to
/// the totalOrder predicate as defined in the IEEE 754 (2008 revision) floating point standard.
/// Note that totalOrder treats positive and negative zeros as different. If it is necessary
/// to treat them as equal, please normalize zeros before calling this kernel. See
/// [`f32::total_cmp`] and [`f64::total_cmp`].
///
/// Nested types, such as lists, are not supported as the null semantics are not well-defined.
/// For comparisons involving nested types see [`crate::ord::make_comparator`]
pub fn gt(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<BooleanArray, ArrowError> {
compare_op(Op::Greater, lhs, rhs)
}
/// Perform `left >= right` operation on two [`Datum`].
///
/// Comparing null values on either side will yield a null in the corresponding
/// slot of the resulting [`BooleanArray`].
///
/// For floating values like f32 and f64, this comparison produces an ordering in accordance to
/// the totalOrder predicate as defined in the IEEE 754 (2008 revision) floating point standard.
/// Note that totalOrder treats positive and negative zeros as different. If it is necessary
/// to treat them as equal, please normalize zeros before calling this kernel. See
/// [`f32::total_cmp`] and [`f64::total_cmp`].
///
/// Nested types, such as lists, are not supported as the null semantics are not well-defined.
/// For comparisons involving nested types see [`crate::ord::make_comparator`]
pub fn gt_eq(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<BooleanArray, ArrowError> {
compare_op(Op::GreaterEqual, lhs, rhs)
}
/// Perform `left IS DISTINCT FROM right` operation on two [`Datum`]
///
/// [`distinct`] is similar to [`neq`], only differing in null handling. In particular, two
/// operands are considered DISTINCT if they have a different value or if one of them is NULL
/// and the other isn't. The result of [`distinct`] is never NULL.
///
/// For floating values like f32 and f64, this comparison produces an ordering in accordance to
/// the totalOrder predicate as defined in the IEEE 754 (2008 revision) floating point standard.
/// Note that totalOrder treats positive and negative zeros as different. If it is necessary
/// to treat them as equal, please normalize zeros before calling this kernel. See
/// [`f32::total_cmp`] and [`f64::total_cmp`].
///
/// Nested types, such as lists, are not supported as the null semantics are not well-defined.
/// For comparisons involving nested types see [`crate::ord::make_comparator`]
pub fn distinct(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<BooleanArray, ArrowError> {
compare_op(Op::Distinct, lhs, rhs)
}
/// Perform `left IS NOT DISTINCT FROM right` operation on two [`Datum`]
///
/// [`not_distinct`] is similar to [`eq`], only differing in null handling. In particular, two
/// operands are considered `NOT DISTINCT` if they have the same value or if both of them
/// is NULL. The result of [`not_distinct`] is never NULL.
///
/// For floating values like f32 and f64, this comparison produces an ordering in accordance to
/// the totalOrder predicate as defined in the IEEE 754 (2008 revision) floating point standard.
/// Note that totalOrder treats positive and negative zeros as different. If it is necessary
/// to treat them as equal, please normalize zeros before calling this kernel. See
/// [`f32::total_cmp`] and [`f64::total_cmp`].
///
/// Nested types, such as lists, are not supported as the null semantics are not well-defined.
/// For comparisons involving nested types see [`crate::ord::make_comparator`]
pub fn not_distinct(lhs: &dyn Datum, rhs: &dyn Datum) -> Result<BooleanArray, ArrowError> {
compare_op(Op::NotDistinct, lhs, rhs)
}
/// Perform `op` on the provided `Datum`
#[inline(never)]
fn compare_op(op: Op, lhs: &dyn Datum, rhs: &dyn Datum) -> Result<BooleanArray, ArrowError> {
use arrow_schema::DataType::*;
let (l, l_s) = lhs.get();
let (r, r_s) = rhs.get();
let l_len = l.len();
let r_len = r.len();
if l_len != r_len && !l_s && !r_s {
return Err(ArrowError::InvalidArgumentError(format!(
"Cannot compare arrays of different lengths, got {l_len} vs {r_len}"
)));
}
let len = match l_s {
true => r_len,
false => l_len,
};
let l_nulls = l.logical_nulls();
let r_nulls = r.logical_nulls();
let l_v = l.as_any_dictionary_opt();
let l = l_v.map(|x| x.values().as_ref()).unwrap_or(l);
let l_t = l.data_type();
let r_v = r.as_any_dictionary_opt();
let r = r_v.map(|x| x.values().as_ref()).unwrap_or(r);
let r_t = r.data_type();
if r_t.is_nested() || l_t.is_nested() {
return Err(ArrowError::InvalidArgumentError(format!(
"Nested comparison: {l_t} {op} {r_t} (hint: use make_comparator instead)"
)));
} else if l_t != r_t {
return Err(ArrowError::InvalidArgumentError(format!(
"Invalid comparison operation: {l_t} {op} {r_t}"
)));
}
// Defer computation as may not be necessary
let values = || -> BooleanBuffer {
let d = downcast_primitive_array! {
(l, r) => apply(op, l.values().as_ref(), l_s, l_v, r.values().as_ref(), r_s, r_v),
(Boolean, Boolean) => apply(op, l.as_boolean(), l_s, l_v, r.as_boolean(), r_s, r_v),
(Utf8, Utf8) => apply(op, l.as_string::<i32>(), l_s, l_v, r.as_string::<i32>(), r_s, r_v),
(Utf8View, Utf8View) => apply(op, l.as_string_view(), l_s, l_v, r.as_string_view(), r_s, r_v),
(LargeUtf8, LargeUtf8) => apply(op, l.as_string::<i64>(), l_s, l_v, r.as_string::<i64>(), r_s, r_v),
(Binary, Binary) => apply(op, l.as_binary::<i32>(), l_s, l_v, r.as_binary::<i32>(), r_s, r_v),
(BinaryView, BinaryView) => apply(op, l.as_binary_view(), l_s, l_v, r.as_binary_view(), r_s, r_v),
(LargeBinary, LargeBinary) => apply(op, l.as_binary::<i64>(), l_s, l_v, r.as_binary::<i64>(), r_s, r_v),
(FixedSizeBinary(_), FixedSizeBinary(_)) => apply(op, l.as_fixed_size_binary(), l_s, l_v, r.as_fixed_size_binary(), r_s, r_v),
(Null, Null) => None,
_ => unreachable!(),
};
d.unwrap_or_else(|| BooleanBuffer::new_unset(len))
};
let l_nulls = l_nulls.filter(|n| n.null_count() > 0);
let r_nulls = r_nulls.filter(|n| n.null_count() > 0);
Ok(match (l_nulls, l_s, r_nulls, r_s) {
(Some(l), true, Some(r), true) | (Some(l), false, Some(r), false) => {
// Either both sides are scalar or neither side is scalar
match op {
Op::Distinct => {
let values = values();
let l = l.inner().bit_chunks().iter_padded();
let r = r.inner().bit_chunks().iter_padded();
let ne = values.bit_chunks().iter_padded();
let c = |((l, r), n)| (l ^ r) | (l & r & n);
let buffer = l.zip(r).zip(ne).map(c).collect();
BooleanBuffer::new(buffer, 0, len).into()
}
Op::NotDistinct => {
let values = values();
let l = l.inner().bit_chunks().iter_padded();
let r = r.inner().bit_chunks().iter_padded();
let e = values.bit_chunks().iter_padded();
let c = |((l, r), e)| u64::not(l | r) | (l & r & e);
let buffer = l.zip(r).zip(e).map(c).collect();
BooleanBuffer::new(buffer, 0, len).into()
}
_ => BooleanArray::new(values(), NullBuffer::union(Some(&l), Some(&r))),
}
}
(Some(_), true, Some(a), false) | (Some(a), false, Some(_), true) => {
// Scalar is null, other side is non-scalar and nullable
match op {
Op::Distinct => a.into_inner().into(),
Op::NotDistinct => a.into_inner().not().into(),
_ => BooleanArray::new_null(len),
}
}
(Some(nulls), is_scalar, None, _) | (None, _, Some(nulls), is_scalar) => {
// Only one side is nullable
match is_scalar {
true => match op {
// Scalar is null, other side is not nullable
Op::Distinct => BooleanBuffer::new_set(len).into(),
Op::NotDistinct => BooleanBuffer::new_unset(len).into(),
_ => BooleanArray::new_null(len),
},
false => match op {
Op::Distinct => {
let values = values();
let l = nulls.inner().bit_chunks().iter_padded();
let ne = values.bit_chunks().iter_padded();
let c = |(l, n)| u64::not(l) | n;
let buffer = l.zip(ne).map(c).collect();
BooleanBuffer::new(buffer, 0, len).into()
}
Op::NotDistinct => (nulls.inner() & &values()).into(),
_ => BooleanArray::new(values(), Some(nulls)),
},
}
}
// Neither side is nullable
(None, _, None, _) => BooleanArray::new(values(), None),
})
}
/// Perform a potentially vectored `op` on the provided `ArrayOrd`
fn apply<T: ArrayOrd>(
op: Op,
l: T,
l_s: bool,
l_v: Option<&dyn AnyDictionaryArray>,
r: T,
r_s: bool,
r_v: Option<&dyn AnyDictionaryArray>,
) -> Option<BooleanBuffer> {
if l.len() == 0 || r.len() == 0 {
return None; // Handle empty dictionaries
}
if !l_s && !r_s && (l_v.is_some() || r_v.is_some()) {
// Not scalar and at least one side has a dictionary, need to perform vectored comparison
let l_v = l_v
.map(|x| x.normalized_keys())
.unwrap_or_else(|| (0..l.len()).collect());
let r_v = r_v
.map(|x| x.normalized_keys())
.unwrap_or_else(|| (0..r.len()).collect());
assert_eq!(l_v.len(), r_v.len()); // Sanity check
Some(match op {
Op::Equal | Op::NotDistinct => apply_op_vectored(l, &l_v, r, &r_v, false, T::is_eq),
Op::NotEqual | Op::Distinct => apply_op_vectored(l, &l_v, r, &r_v, true, T::is_eq),
Op::Less => apply_op_vectored(l, &l_v, r, &r_v, false, T::is_lt),
Op::LessEqual => apply_op_vectored(r, &r_v, l, &l_v, true, T::is_lt),
Op::Greater => apply_op_vectored(r, &r_v, l, &l_v, false, T::is_lt),
Op::GreaterEqual => apply_op_vectored(l, &l_v, r, &r_v, true, T::is_lt),
})
} else {
let l_s = l_s.then(|| l_v.map(|x| x.normalized_keys()[0]).unwrap_or_default());
let r_s = r_s.then(|| r_v.map(|x| x.normalized_keys()[0]).unwrap_or_default());
let buffer = match op {
Op::Equal | Op::NotDistinct => apply_op(l, l_s, r, r_s, false, T::is_eq),
Op::NotEqual | Op::Distinct => apply_op(l, l_s, r, r_s, true, T::is_eq),
Op::Less => apply_op(l, l_s, r, r_s, false, T::is_lt),
Op::LessEqual => apply_op(r, r_s, l, l_s, true, T::is_lt),
Op::Greater => apply_op(r, r_s, l, l_s, false, T::is_lt),
Op::GreaterEqual => apply_op(l, l_s, r, r_s, true, T::is_lt),
};
// If a side had a dictionary, and was not scalar, we need to materialize this
Some(match (l_v, r_v) {
(Some(l_v), _) if l_s.is_none() => take_bits(l_v, buffer),
(_, Some(r_v)) if r_s.is_none() => take_bits(r_v, buffer),
_ => buffer,
})
}
}
/// Perform a take operation on `buffer` with the given dictionary
fn take_bits(v: &dyn AnyDictionaryArray, buffer: BooleanBuffer) -> BooleanBuffer {
let array = take(&BooleanArray::new(buffer, None), v.keys(), None).unwrap();
array.as_boolean().values().clone()
}
/// Invokes `f` with values `0..len` collecting the boolean results into a new `BooleanBuffer`
///
/// This is similar to [`arrow_buffer::MutableBuffer::collect_bool`] but with
/// the option to efficiently negate the result
fn collect_bool(len: usize, neg: bool, f: impl Fn(usize) -> bool) -> BooleanBuffer {
let mut buffer = Vec::with_capacity(ceil(len, 64));
let chunks = len / 64;
let remainder = len % 64;
buffer.extend((0..chunks).map(|chunk| {
let mut packed = 0;
for bit_idx in 0..64 {
let i = bit_idx + chunk * 64;
packed |= (f(i) as u64) << bit_idx;
}
if neg {
packed = !packed
}
packed
}));
if remainder != 0 {
let mut packed = 0;
for bit_idx in 0..remainder {
let i = bit_idx + chunks * 64;
packed |= (f(i) as u64) << bit_idx;
}
if neg {
packed = !packed
}
buffer.push(packed);
}
BooleanBuffer::new(buffer.into(), 0, len)
}
/// Applies `op` to possibly scalar `ArrayOrd`
///
/// If l is scalar `l_s` will be `Some(idx)` where `idx` is the index of the scalar value in `l`
/// If r is scalar `r_s` will be `Some(idx)` where `idx` is the index of the scalar value in `r`
///
/// If `neg` is true the result of `op` will be negated
fn apply_op<T: ArrayOrd>(
l: T,
l_s: Option<usize>,
r: T,
r_s: Option<usize>,
neg: bool,
op: impl Fn(T::Item, T::Item) -> bool,
) -> BooleanBuffer {
match (l_s, r_s) {
(None, None) => {
assert_eq!(l.len(), r.len());
collect_bool(l.len(), neg, |idx| unsafe {
op(l.value_unchecked(idx), r.value_unchecked(idx))
})
}
(Some(l_s), Some(r_s)) => {
let a = l.value(l_s);
let b = r.value(r_s);
std::iter::once(op(a, b) ^ neg).collect()
}
(Some(l_s), None) => {
let v = l.value(l_s);
collect_bool(r.len(), neg, |idx| op(v, unsafe { r.value_unchecked(idx) }))
}
(None, Some(r_s)) => {
let v = r.value(r_s);
collect_bool(l.len(), neg, |idx| op(unsafe { l.value_unchecked(idx) }, v))
}
}
}
/// Applies `op` to possibly scalar `ArrayOrd` with the given indices
fn apply_op_vectored<T: ArrayOrd>(
l: T,
l_v: &[usize],
r: T,
r_v: &[usize],
neg: bool,
op: impl Fn(T::Item, T::Item) -> bool,
) -> BooleanBuffer {
assert_eq!(l_v.len(), r_v.len());
collect_bool(l_v.len(), neg, |idx| unsafe {
let l_idx = *l_v.get_unchecked(idx);
let r_idx = *r_v.get_unchecked(idx);
op(l.value_unchecked(l_idx), r.value_unchecked(r_idx))
})
}
trait ArrayOrd {
type Item: Copy;
fn len(&self) -> usize;
fn value(&self, idx: usize) -> Self::Item {
assert!(idx < self.len());
unsafe { self.value_unchecked(idx) }
}
/// # Safety
///
/// Safe if `idx < self.len()`
unsafe fn value_unchecked(&self, idx: usize) -> Self::Item;
fn is_eq(l: Self::Item, r: Self::Item) -> bool;
fn is_lt(l: Self::Item, r: Self::Item) -> bool;
}
impl ArrayOrd for &BooleanArray {
type Item = bool;
fn len(&self) -> usize {
Array::len(self)
}
unsafe fn value_unchecked(&self, idx: usize) -> Self::Item {
unsafe { BooleanArray::value_unchecked(self, idx) }
}
fn is_eq(l: Self::Item, r: Self::Item) -> bool {
l == r
}
fn is_lt(l: Self::Item, r: Self::Item) -> bool {
!l & r
}
}
impl<T: ArrowNativeTypeOp> ArrayOrd for &[T] {
type Item = T;
fn len(&self) -> usize {
(*self).len()
}
unsafe fn value_unchecked(&self, idx: usize) -> Self::Item {
unsafe { *self.get_unchecked(idx) }
}
fn is_eq(l: Self::Item, r: Self::Item) -> bool {
l.is_eq(r)
}
fn is_lt(l: Self::Item, r: Self::Item) -> bool {
l.is_lt(r)
}
}
impl<'a, T: ByteArrayType> ArrayOrd for &'a GenericByteArray<T> {
type Item = &'a [u8];
fn len(&self) -> usize {
Array::len(self)
}
unsafe fn value_unchecked(&self, idx: usize) -> Self::Item {
unsafe { GenericByteArray::value_unchecked(self, idx).as_ref() }
}
fn is_eq(l: Self::Item, r: Self::Item) -> bool {
l == r
}
fn is_lt(l: Self::Item, r: Self::Item) -> bool {
l < r
}
}
impl<'a, T: ByteViewType> ArrayOrd for &'a GenericByteViewArray<T> {
/// This is the item type for the GenericByteViewArray::compare
/// Item.0 is the array, Item.1 is the index
type Item = (&'a GenericByteViewArray<T>, usize);
#[inline(always)]
fn is_eq(l: Self::Item, r: Self::Item) -> bool {
let l_view = unsafe { l.0.views().get_unchecked(l.1) };
let r_view = unsafe { r.0.views().get_unchecked(r.1) };
if l.0.data_buffers().is_empty() && r.0.data_buffers().is_empty() {
// For eq case, we can directly compare the inlined bytes
return l_view == r_view;
}
let l_len = *l_view as u32;
let r_len = *r_view as u32;
// This is a fast path for equality check.
// We don't need to look at the actual bytes to determine if they are equal.
if l_len != r_len {
return false;
}
if l_len == 0 && r_len == 0 {
return true;
}
// # Safety
// The index is within bounds as it is checked in value()
unsafe { GenericByteViewArray::compare_unchecked(l.0, l.1, r.0, r.1).is_eq() }
}
#[inline(always)]
fn is_lt(l: Self::Item, r: Self::Item) -> bool {
// If both arrays use only the inline buffer
if l.0.data_buffers().is_empty() && r.0.data_buffers().is_empty() {
let l_view = unsafe { l.0.views().get_unchecked(l.1) };
let r_view = unsafe { r.0.views().get_unchecked(r.1) };
return GenericByteViewArray::<T>::inline_key_fast(*l_view)
< GenericByteViewArray::<T>::inline_key_fast(*r_view);
}
// Fallback to the generic, unchecked comparison for non-inline cases
// # Safety
// The index is within bounds as it is checked in value()
unsafe { GenericByteViewArray::compare_unchecked(l.0, l.1, r.0, r.1).is_lt() }
}
fn len(&self) -> usize {
Array::len(self)
}
unsafe fn value_unchecked(&self, idx: usize) -> Self::Item {
(self, idx)
}
}
impl<'a> ArrayOrd for &'a FixedSizeBinaryArray {
type Item = &'a [u8];
fn len(&self) -> usize {
Array::len(self)
}
unsafe fn value_unchecked(&self, idx: usize) -> Self::Item {
unsafe { FixedSizeBinaryArray::value_unchecked(self, idx) }
}
fn is_eq(l: Self::Item, r: Self::Item) -> bool {
l == r
}
fn is_lt(l: Self::Item, r: Self::Item) -> bool {
l < r
}
}
/// Compares two [`GenericByteViewArray`] at index `left_idx` and `right_idx`
#[inline(always)]
pub fn compare_byte_view<T: ByteViewType>(
left: &GenericByteViewArray<T>,
left_idx: usize,
right: &GenericByteViewArray<T>,
right_idx: usize,
) -> Ordering {
assert!(left_idx < left.len());
assert!(right_idx < right.len());
if left.data_buffers().is_empty() && right.data_buffers().is_empty() {
let l_view = unsafe { left.views().get_unchecked(left_idx) };
let r_view = unsafe { right.views().get_unchecked(right_idx) };
return GenericByteViewArray::<T>::inline_key_fast(*l_view)
.cmp(&GenericByteViewArray::<T>::inline_key_fast(*r_view));
}
unsafe { GenericByteViewArray::compare_unchecked(left, left_idx, right, right_idx) }
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use arrow_array::{DictionaryArray, Int32Array, Scalar, StringArray};
use super::*;
#[test]
fn test_null_dict() {
let a = DictionaryArray::new(Int32Array::new_null(10), Arc::new(Int32Array::new_null(0)));
let r = eq(&a, &a).unwrap();
assert_eq!(r.null_count(), 10);
let a = DictionaryArray::new(
Int32Array::from(vec![1, 2, 3, 4, 5, 6]),
Arc::new(Int32Array::new_null(10)),
);
let r = eq(&a, &a).unwrap();
assert_eq!(r.null_count(), 6);
let scalar =
DictionaryArray::new(Int32Array::new_null(1), Arc::new(Int32Array::new_null(0)));
let r = eq(&a, &Scalar::new(&scalar)).unwrap();
assert_eq!(r.null_count(), 6);
let scalar =
DictionaryArray::new(Int32Array::new_null(1), Arc::new(Int32Array::new_null(0)));
let r = eq(&Scalar::new(&scalar), &Scalar::new(&scalar)).unwrap();
assert_eq!(r.null_count(), 1);
let a = DictionaryArray::new(
Int32Array::from(vec![0, 1, 2]),
Arc::new(Int32Array::from(vec![3, 2, 1])),
);
let r = eq(&a, &Scalar::new(&scalar)).unwrap();
assert_eq!(r.null_count(), 3);
}
#[test]
fn is_distinct_from_non_nulls() {
let left_int_array = Int32Array::from(vec![0, 1, 2, 3, 4]);
let right_int_array = Int32Array::from(vec![4, 3, 2, 1, 0]);
assert_eq!(
BooleanArray::from(vec![true, true, false, true, true,]),
distinct(&left_int_array, &right_int_array).unwrap()
);
assert_eq!(
BooleanArray::from(vec![false, false, true, false, false,]),
not_distinct(&left_int_array, &right_int_array).unwrap()
);
}
#[test]
fn is_distinct_from_nulls() {
// [0, 0, NULL, 0, 0, 0]
let left_int_array = Int32Array::new(
vec![0, 0, 1, 3, 0, 0].into(),
Some(NullBuffer::from(vec![true, true, false, true, true, true])),
);
// [0, NULL, NULL, NULL, 0, NULL]
let right_int_array = Int32Array::new(
vec![0; 6].into(),
Some(NullBuffer::from(vec![
true, false, false, false, true, false,
])),
);
assert_eq!(
BooleanArray::from(vec![false, true, false, true, false, true,]),
distinct(&left_int_array, &right_int_array).unwrap()
);
assert_eq!(
BooleanArray::from(vec![true, false, true, false, true, false,]),
not_distinct(&left_int_array, &right_int_array).unwrap()
);
}
#[test]
fn test_distinct_scalar() {
let a = Int32Array::new_scalar(12);
let b = Int32Array::new_scalar(12);
assert!(!distinct(&a, &b).unwrap().value(0));
assert!(not_distinct(&a, &b).unwrap().value(0));
let a = Int32Array::new_scalar(12);
let b = Int32Array::new_null(1);
assert!(distinct(&a, &b).unwrap().value(0));
assert!(!not_distinct(&a, &b).unwrap().value(0));
assert!(distinct(&b, &a).unwrap().value(0));
assert!(!not_distinct(&b, &a).unwrap().value(0));
let b = Scalar::new(b);
assert!(distinct(&a, &b).unwrap().value(0));
assert!(!not_distinct(&a, &b).unwrap().value(0));
assert!(!distinct(&b, &b).unwrap().value(0));
assert!(not_distinct(&b, &b).unwrap().value(0));
let a = Int32Array::new(
vec![0, 1, 2, 3].into(),
Some(vec![false, false, true, true].into()),
);
let expected = BooleanArray::from(vec![false, false, true, true]);
assert_eq!(distinct(&a, &b).unwrap(), expected);
assert_eq!(distinct(&b, &a).unwrap(), expected);
let expected = BooleanArray::from(vec![true, true, false, false]);
assert_eq!(not_distinct(&a, &b).unwrap(), expected);
assert_eq!(not_distinct(&b, &a).unwrap(), expected);
let b = Int32Array::new_scalar(1);
let expected = BooleanArray::from(vec![true; 4]);
assert_eq!(distinct(&a, &b).unwrap(), expected);
assert_eq!(distinct(&b, &a).unwrap(), expected);
let expected = BooleanArray::from(vec![false; 4]);
assert_eq!(not_distinct(&a, &b).unwrap(), expected);
assert_eq!(not_distinct(&b, &a).unwrap(), expected);
let b = Int32Array::new_scalar(3);
let expected = BooleanArray::from(vec![true, true, true, false]);
assert_eq!(distinct(&a, &b).unwrap(), expected);
assert_eq!(distinct(&b, &a).unwrap(), expected);
let expected = BooleanArray::from(vec![false, false, false, true]);
assert_eq!(not_distinct(&a, &b).unwrap(), expected);
assert_eq!(not_distinct(&b, &a).unwrap(), expected);
}
#[test]
fn test_scalar_negation() {
let a = Int32Array::new_scalar(54);
let b = Int32Array::new_scalar(54);
let r = eq(&a, &b).unwrap();
assert!(r.value(0));
let r = neq(&a, &b).unwrap();
assert!(!r.value(0))
}
#[test]
fn test_scalar_empty() {
let a = Int32Array::new_null(0);
let b = Int32Array::new_scalar(23);
let r = eq(&a, &b).unwrap();
assert_eq!(r.len(), 0);
let r = eq(&b, &a).unwrap();
assert_eq!(r.len(), 0);
}
#[test]
fn test_dictionary_nulls() {
let values = StringArray::from(vec![Some("us-west"), Some("us-east")]);
let nulls = NullBuffer::from(vec![false, true, true]);
let key_values = vec![100i32, 1i32, 0i32].into();
let keys = Int32Array::new(key_values, Some(nulls));
let col = DictionaryArray::try_new(keys, Arc::new(values)).unwrap();
neq(&col.slice(0, col.len() - 1), &col.slice(1, col.len() - 1)).unwrap();
}
}
File diff suppressed because it is too large Load Diff
+58
View File
@@ -0,0 +1,58 @@
// Licensed to the Apache Software Foundation (ASF) under one
// or more contributor license agreements. See the NOTICE file
// distributed with this work for additional information
// regarding copyright ownership. The ASF licenses this file
// to you under the Apache License, Version 2.0 (the
// "License"); you may not use this file except in compliance
// with the License. You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing,
// software distributed under the License is distributed on an
// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
// KIND, either express or implied. See the License for the
// specific language governing permissions and limitations
// under the License.
//! Arrow ordering kernels
//!
//! # Sort RecordBatch
//!
//! ```
//! # use std::sync::Arc;
//! # use arrow_array::*;
//! # use arrow_array::cast::AsArray;
//! # use arrow_array::types::Int32Type;
//! # use arrow_ord::sort::sort_to_indices;
//! # use arrow_select::take::take;
//! #
//! let a: ArrayRef = Arc::new(Int32Array::from(vec![1, 2, 3, 4]));
//! let b: ArrayRef = Arc::new(StringArray::from(vec!["b", "a", "e", "d"]));
//! let batch = RecordBatch::try_from_iter(vec![("a", a), ("b", b)]).unwrap();
//!
//! // Sort by column 1
//! let indices = sort_to_indices(batch.column(1), None, None).unwrap();
//!
//! // Apply indices to batch columns
//! let columns = batch.columns().iter().map(|c| take(&*c, &indices, None).unwrap()).collect();
//! let sorted = RecordBatch::try_new(batch.schema(), columns).unwrap();
//!
//! let col1 = sorted.column(0).as_primitive::<Int32Type>();
//! assert_eq!(col1.values(), &[2, 1, 4, 3]);
//! ```
//!
#![doc(
html_logo_url = "https://arrow.apache.org/img/arrow-logo_chevrons_black-txt_white-bg.svg",
html_favicon_url = "https://arrow.apache.org/img/arrow-logo_chevrons_black-txt_transparent-bg.svg"
)]
#![cfg_attr(docsrs, feature(doc_cfg))]
#![warn(missing_docs)]
pub mod cmp;
#[doc(hidden)]
pub mod comparison;
pub mod ord;
pub mod partition;
pub mod rank;
pub mod sort;
File diff suppressed because it is too large Load Diff
+335
View File
@@ -0,0 +1,335 @@
// Licensed to the Apache Software Foundation (ASF) under one
// or more contributor license agreements. See the NOTICE file
// distributed with this work for additional information
// regarding copyright ownership. The ASF licenses this file
// to you under the Apache License, Version 2.0 (the
// "License"); you may not use this file except in compliance
// with the License. You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing,
// software distributed under the License is distributed on an
// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
// KIND, either express or implied. See the License for the
// specific language governing permissions and limitations
// under the License.
//! Defines partition kernel for `ArrayRef`
use std::ops::Range;
use arrow_array::{Array, ArrayRef};
use arrow_buffer::BooleanBuffer;
use arrow_schema::{ArrowError, SortOptions};
use crate::cmp::distinct;
use crate::ord::make_comparator;
/// A computed set of partitions, see [`partition`]
#[derive(Debug, Clone)]
pub struct Partitions(Option<BooleanBuffer>);
impl Partitions {
/// Returns the range of each partition
///
/// Consecutive ranges will be contiguous: i.e [`(a, b)` and `(b, c)`], and
/// `start = 0` and `end = self.len()` for the first and last range respectively
pub fn ranges(&self) -> Vec<Range<usize>> {
let boundaries = match &self.0 {
Some(boundaries) => boundaries,
None => return vec![],
};
let mut out = vec![];
let mut current = 0;
for idx in boundaries.set_indices() {
let t = current;
current = idx + 1;
out.push(t..current)
}
let last = boundaries.len() + 1;
if current != last {
out.push(current..last)
}
out
}
/// Returns the number of partitions
pub fn len(&self) -> usize {
match &self.0 {
Some(b) => b.count_set_bits() + 1,
None => 0,
}
}
/// Returns true if this contains no partitions
pub fn is_empty(&self) -> bool {
self.0.is_none()
}
}
/// Given a list of lexicographically sorted columns, computes the [`Partitions`],
/// where a partition consists of the set of consecutive rows with equal values
///
/// Returns an error if no columns are specified or all columns do not
/// have the same number of rows.
///
/// # Example:
///
/// For example, given columns `x`, `y` and `z`, calling
/// [`partition`]`(values, (x, y))` will divide the
/// rows into ranges where the values of `(x, y)` are equal:
///
/// ```text
/// ┌ ─ ┬───┬ ─ ─┌───┐─ ─ ┬───┬ ─ ─ ┐
/// │ 1 │ │ 1 │ │ A │ Range: 0..1 (x=1, y=1)
/// ├ ─ ┼───┼ ─ ─├───┤─ ─ ┼───┼ ─ ─ ┤
/// │ 1 │ │ 2 │ │ B │
/// │ ├───┤ ├───┤ ├───┤ │
/// │ 1 │ │ 2 │ │ C │ Range: 1..4 (x=1, y=2)
/// │ ├───┤ ├───┤ ├───┤ │
/// │ 1 │ │ 2 │ │ D │
/// ├ ─ ┼───┼ ─ ─├───┤─ ─ ┼───┼ ─ ─ ┤
/// │ 2 │ │ 1 │ │ E │ Range: 4..5 (x=2, y=1)
/// ├ ─ ┼───┼ ─ ─├───┤─ ─ ┼───┼ ─ ─ ┤
/// │ 3 │ │ 1 │ │ F │ Range: 5..6 (x=3, y=1)
/// └ ─ ┴───┴ ─ ─└───┘─ ─ ┴───┴ ─ ─ ┘
///
/// x y z partition(&[x, y])
/// ```
///
/// # Example Code
///
/// ```
/// # use std::{sync::Arc, ops::Range};
/// # use arrow_array::{RecordBatch, Int64Array, StringArray, ArrayRef};
/// # use arrow_ord::sort::{SortColumn, SortOptions};
/// # use arrow_ord::partition::partition;
/// let batch = RecordBatch::try_from_iter(vec![
/// ("x", Arc::new(Int64Array::from(vec![1, 1, 1, 1, 2, 3])) as ArrayRef),
/// ("y", Arc::new(Int64Array::from(vec![1, 2, 2, 2, 1, 1])) as ArrayRef),
/// ("z", Arc::new(StringArray::from(vec!["A", "B", "C", "D", "E", "F"])) as ArrayRef),
/// ]).unwrap();
///
/// // Partition on first two columns
/// let ranges = partition(&batch.columns()[..2]).unwrap().ranges();
///
/// let expected = vec![
/// (0..1),
/// (1..4),
/// (4..5),
/// (5..6),
/// ];
///
/// assert_eq!(ranges, expected);
/// ```
pub fn partition(columns: &[ArrayRef]) -> Result<Partitions, ArrowError> {
if columns.is_empty() {
return Err(ArrowError::InvalidArgumentError(
"Partition requires at least one column".to_string(),
));
}
let num_rows = columns[0].len();
if columns.iter().any(|item| item.len() != num_rows) {
return Err(ArrowError::InvalidArgumentError(
"Partition columns have different row counts".to_string(),
));
};
match num_rows {
0 => return Ok(Partitions(None)),
1 => return Ok(Partitions(Some(BooleanBuffer::new_unset(0)))),
_ => {}
}
let acc = find_boundaries(&columns[0])?;
let acc = columns
.iter()
.skip(1)
.try_fold(acc, |acc, c| find_boundaries(c.as_ref()).map(|b| &acc | &b))?;
Ok(Partitions(Some(acc)))
}
/// Returns a mask with bits set whenever the value or nullability changes
fn find_boundaries(v: &dyn Array) -> Result<BooleanBuffer, ArrowError> {
let slice_len = v.len() - 1;
let v1 = v.slice(0, slice_len);
let v2 = v.slice(1, slice_len);
if !v.data_type().is_nested() {
return Ok(distinct(&v1, &v2)?.values().clone());
}
// Given that we're only comparing values, null ordering in the input or
// sort options do not matter.
let cmp = make_comparator(&v1, &v2, SortOptions::default())?;
Ok((0..slice_len).map(|i| !cmp(i, i).is_eq()).collect())
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use super::*;
use arrow_array::*;
use arrow_schema::DataType;
#[test]
fn test_partition_empty() {
let err = partition(&[]).unwrap_err();
assert_eq!(
err.to_string(),
"Invalid argument error: Partition requires at least one column"
);
}
#[test]
fn test_partition_unaligned_rows() {
let input = vec![
Arc::new(Int64Array::from(vec![None, Some(-1)])) as _,
Arc::new(StringArray::from(vec![Some("foo")])) as _,
];
let err = partition(&input).unwrap_err();
assert_eq!(
err.to_string(),
"Invalid argument error: Partition columns have different row counts"
)
}
#[test]
fn test_partition_small() {
let results = partition(&[
Arc::new(Int32Array::new(vec![].into(), None)) as _,
Arc::new(Int32Array::new(vec![].into(), None)) as _,
Arc::new(Int32Array::new(vec![].into(), None)) as _,
])
.unwrap();
assert_eq!(results.len(), 0);
assert!(results.is_empty());
let results = partition(&[
Arc::new(Int32Array::from(vec![1])) as _,
Arc::new(Int32Array::from(vec![1])) as _,
])
.unwrap()
.ranges();
assert_eq!(results.len(), 1);
assert_eq!(results[0], 0..1);
}
#[test]
fn test_partition_single_column() {
let a = Int64Array::from(vec![1, 2, 2, 2, 2, 2, 2, 2, 9]);
let input = vec![Arc::new(a) as _];
assert_eq!(
partition(&input).unwrap().ranges(),
vec![(0..1), (1..8), (8..9)],
);
}
#[test]
fn test_partition_all_equal_values() {
let a = Int64Array::from_value(1, 1000);
let input = vec![Arc::new(a) as _];
assert_eq!(partition(&input).unwrap().ranges(), vec![(0..1000)]);
}
#[test]
fn test_partition_all_null_values() {
let input = vec![
new_null_array(&DataType::Int8, 1000),
new_null_array(&DataType::UInt16, 1000),
];
assert_eq!(partition(&input).unwrap().ranges(), vec![(0..1000)]);
}
#[test]
fn test_partition_unique_column_1() {
let input = vec![
Arc::new(Int64Array::from(vec![None, Some(-1)])) as _,
Arc::new(StringArray::from(vec![Some("foo"), Some("bar")])) as _,
];
assert_eq!(partition(&input).unwrap().ranges(), vec![(0..1), (1..2)],);
}
#[test]
fn test_partition_unique_column_2() {
let input = vec![
Arc::new(Int64Array::from(vec![None, Some(-1), Some(-1)])) as _,
Arc::new(StringArray::from(vec![
Some("foo"),
Some("bar"),
Some("apple"),
])) as _,
];
assert_eq!(
partition(&input).unwrap().ranges(),
vec![(0..1), (1..2), (2..3),],
);
}
#[test]
fn test_partition_non_unique_column_1() {
let input = vec![
Arc::new(Int64Array::from(vec![None, Some(-1), Some(-1), Some(1)])) as _,
Arc::new(StringArray::from(vec![
Some("foo"),
Some("bar"),
Some("bar"),
Some("bar"),
])) as _,
];
assert_eq!(
partition(&input).unwrap().ranges(),
vec![(0..1), (1..3), (3..4),],
);
}
#[test]
fn test_partition_masked_nulls() {
let input = vec![
Arc::new(Int64Array::new(vec![1; 9].into(), None)) as _,
Arc::new(Int64Array::new(
vec![1, 1, 2, 2, 2, 3, 3, 3, 3].into(),
Some(vec![false, true, true, true, true, false, false, true, false].into()),
)) as _,
Arc::new(Int64Array::new(
vec![1, 1, 2, 2, 2, 2, 2, 3, 7].into(),
Some(vec![true, true, true, true, false, true, true, true, false].into()),
)) as _,
];
assert_eq!(
partition(&input).unwrap().ranges(),
vec![(0..1), (1..2), (2..4), (4..5), (5..7), (7..8), (8..9)],
);
}
#[test]
fn test_partition_nested() {
let input = vec![
Arc::new(
StructArray::try_from(vec![(
"f1",
Arc::new(Int64Array::from(vec![
None,
None,
Some(1),
Some(2),
Some(2),
Some(2),
Some(3),
Some(4),
])) as _,
)])
.unwrap(),
) as _,
Arc::new(Int64Array::from(vec![1, 1, 1, 2, 3, 3, 3, 4])) as _,
];
assert_eq!(
partition(&input).unwrap().ranges(),
vec![0..2, 2..3, 3..4, 4..6, 6..7, 7..8]
)
}
}
+361
View File
@@ -0,0 +1,361 @@
// Licensed to the Apache Software Foundation (ASF) under one
// or more contributor license agreements. See the NOTICE file
// distributed with this work for additional information
// regarding copyright ownership. The ASF licenses this file
// to you under the Apache License, Version 2.0 (the
// "License"); you may not use this file except in compliance
// with the License. You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing,
// software distributed under the License is distributed on an
// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
// KIND, either express or implied. See the License for the
// specific language governing permissions and limitations
// under the License.
//! Provides `rank` function to assign a rank to each value in an array
use arrow_array::cast::AsArray;
use arrow_array::types::*;
use arrow_array::{
Array, ArrowNativeTypeOp, BooleanArray, GenericByteArray, downcast_primitive_array,
};
use arrow_buffer::NullBuffer;
use arrow_schema::{ArrowError, DataType, SortOptions};
use std::cmp::Ordering;
/// Whether `arrow_ord::rank` can rank an array of given data type.
pub(crate) fn can_rank(data_type: &DataType) -> bool {
data_type.is_primitive()
|| matches!(
data_type,
DataType::Boolean
| DataType::Utf8
| DataType::LargeUtf8
| DataType::Binary
| DataType::LargeBinary
)
}
/// Assigns a rank to each value in `array` based on its position in the sorted order
///
/// Where values are equal, they will be assigned the highest of their ranks,
/// leaving gaps in the overall rank assignment
///
/// ```
/// # use arrow_array::StringArray;
/// # use arrow_ord::rank::rank;
/// let array = StringArray::from(vec![Some("foo"), None, Some("foo"), None, Some("bar")]);
/// let ranks = rank(&array, None).unwrap();
/// assert_eq!(ranks, &[5, 2, 5, 2, 3]);
/// ```
pub fn rank(array: &dyn Array, options: Option<SortOptions>) -> Result<Vec<u32>, ArrowError> {
let options = options.unwrap_or_default();
let ranks = downcast_primitive_array! {
array => primitive_rank(array.values(), array.nulls(), options),
DataType::Boolean => boolean_rank(array.as_boolean(), options),
DataType::Utf8 => bytes_rank(array.as_bytes::<Utf8Type>(), options),
DataType::LargeUtf8 => bytes_rank(array.as_bytes::<LargeUtf8Type>(), options),
DataType::Binary => bytes_rank(array.as_bytes::<BinaryType>(), options),
DataType::LargeBinary => bytes_rank(array.as_bytes::<LargeBinaryType>(), options),
d => return Err(ArrowError::ComputeError(format!("{d:?} not supported in rank")))
};
Ok(ranks)
}
#[inline(never)]
fn primitive_rank<T: ArrowNativeTypeOp>(
values: &[T],
nulls: Option<&NullBuffer>,
options: SortOptions,
) -> Vec<u32> {
let len: u32 = values.len().try_into().unwrap();
let to_sort = match nulls.filter(|n| n.null_count() > 0) {
Some(n) => n
.valid_indices()
.map(|idx| (values[idx], idx as u32))
.collect(),
None => values.iter().copied().zip(0..len).collect(),
};
rank_impl(values.len(), to_sort, options, T::compare, T::is_eq)
}
#[inline(never)]
fn bytes_rank<T: ByteArrayType>(array: &GenericByteArray<T>, options: SortOptions) -> Vec<u32> {
let to_sort: Vec<(&[u8], u32)> = match array.nulls().filter(|n| n.null_count() > 0) {
Some(n) => n
.valid_indices()
.map(|idx| (array.value(idx).as_ref(), idx as u32))
.collect(),
None => (0..array.len())
.map(|idx| (array.value(idx).as_ref(), idx as u32))
.collect(),
};
rank_impl(array.len(), to_sort, options, Ord::cmp, PartialEq::eq)
}
fn rank_impl<T, C, E>(
len: usize,
mut valid: Vec<(T, u32)>,
options: SortOptions,
compare: C,
eq: E,
) -> Vec<u32>
where
T: Copy,
C: Fn(T, T) -> Ordering,
E: Fn(T, T) -> bool,
{
// We can use an unstable sort as we combine equal values later
valid.sort_unstable_by(|a, b| compare(a.0, b.0));
if options.descending {
valid.reverse();
}
let (mut valid_rank, null_rank) = match options.nulls_first {
true => (len as u32, (len - valid.len()) as u32),
false => (valid.len() as u32, len as u32),
};
let mut out: Vec<_> = vec![null_rank; len];
if let Some(v) = valid.last() {
out[v.1 as usize] = valid_rank;
}
let mut count = 1; // Number of values in rank
for w in valid.windows(2).rev() {
match eq(w[0].0, w[1].0) {
true => {
count += 1;
out[w[0].1 as usize] = valid_rank;
}
false => {
valid_rank -= count;
count = 1;
out[w[0].1 as usize] = valid_rank
}
}
}
out
}
/// Return the index for the rank when ranking boolean array
///
/// The index is calculated as follows:
/// if is_null is true, the index is 2
/// if is_null is false and the value is true, the index is 1
/// otherwise, the index is 0
///
/// false is 0 and true is 1 because these are the value when cast to number
#[inline]
fn get_boolean_rank_index(value: bool, is_null: bool) -> usize {
let is_null_num = is_null as usize;
(is_null_num << 1) | (value as usize & !is_null_num)
}
#[inline(never)]
fn boolean_rank(array: &BooleanArray, options: SortOptions) -> Vec<u32> {
let null_count = array.null_count() as u32;
let true_count = array.true_count() as u32;
let false_count = array.len() as u32 - null_count - true_count;
// Rank values for [false, true, null] in that order
//
// The value for a rank is last value rank + own value count
// this means that if we have the following order: `false`, `true` and then `null`
// the ranks will be:
// - false: false_count
// - true: false_count + true_count
// - null: false_count + true_count + null_count
//
// If we have the following order: `null`, `false` and then `true`
// the ranks will be:
// - false: null_count + false_count
// - true: null_count + false_count + true_count
// - null: null_count
//
// You will notice that the last rank is always the total length of the array but we don't use it for readability on how the rank is calculated
let ranks_index: [u32; 3] = match (options.descending, options.nulls_first) {
// The order is null, true, false
(true, true) => [
null_count + true_count + false_count,
null_count + true_count,
null_count,
],
// The order is true, false, null
(true, false) => [
true_count + false_count,
true_count,
true_count + false_count + null_count,
],
// The order is null, false, true
(false, true) => [
null_count + false_count,
null_count + false_count + true_count,
null_count,
],
// The order is false, true, null
(false, false) => [
false_count,
false_count + true_count,
false_count + true_count + null_count,
],
};
match array.nulls().filter(|n| n.null_count() > 0) {
Some(n) => array
.values()
.iter()
.zip(n.iter())
.map(|(value, is_valid)| ranks_index[get_boolean_rank_index(value, !is_valid)])
.collect::<Vec<u32>>(),
None => array
.values()
.iter()
.map(|value| ranks_index[value as usize])
.collect::<Vec<u32>>(),
}
}
#[cfg(test)]
mod tests {
use super::*;
use arrow_array::*;
#[test]
fn test_primitive() {
let descending = SortOptions {
descending: true,
nulls_first: true,
};
let nulls_last = SortOptions {
descending: false,
nulls_first: false,
};
let nulls_last_descending = SortOptions {
descending: true,
nulls_first: false,
};
let a = Int32Array::from(vec![Some(1), Some(1), None, Some(3), Some(3), Some(4)]);
let res = rank(&a, None).unwrap();
assert_eq!(res, &[3, 3, 1, 5, 5, 6]);
let res = rank(&a, Some(descending)).unwrap();
assert_eq!(res, &[6, 6, 1, 4, 4, 2]);
let res = rank(&a, Some(nulls_last)).unwrap();
assert_eq!(res, &[2, 2, 6, 4, 4, 5]);
let res = rank(&a, Some(nulls_last_descending)).unwrap();
assert_eq!(res, &[5, 5, 6, 3, 3, 1]);
// Test with non-zero null values
let nulls = NullBuffer::from(vec![true, true, false, true, false, false]);
let a = Int32Array::new(vec![1, 4, 3, 4, 5, 5].into(), Some(nulls));
let res = rank(&a, None).unwrap();
assert_eq!(res, &[4, 6, 3, 6, 3, 3]);
}
#[test]
fn test_get_boolean_rank_index() {
assert_eq!(get_boolean_rank_index(true, true), 2);
assert_eq!(get_boolean_rank_index(false, true), 2);
assert_eq!(get_boolean_rank_index(true, false), 1);
assert_eq!(get_boolean_rank_index(false, false), 0);
}
#[test]
fn test_nullable_booleans() {
let descending = SortOptions {
descending: true,
nulls_first: true,
};
let nulls_last = SortOptions {
descending: false,
nulls_first: false,
};
let nulls_last_descending = SortOptions {
descending: true,
nulls_first: false,
};
let a = BooleanArray::from(vec![Some(true), Some(true), None, Some(false), Some(false)]);
let res = rank(&a, None).unwrap();
assert_eq!(res, &[5, 5, 1, 3, 3]);
let res = rank(&a, Some(descending)).unwrap();
assert_eq!(res, &[3, 3, 1, 5, 5]);
let res = rank(&a, Some(nulls_last)).unwrap();
assert_eq!(res, &[4, 4, 5, 2, 2]);
let res = rank(&a, Some(nulls_last_descending)).unwrap();
assert_eq!(res, &[2, 2, 5, 4, 4]);
// Test with non-zero null values
let nulls = NullBuffer::from(vec![true, true, false, true, true]);
let a = BooleanArray::new(vec![true, true, true, false, false].into(), Some(nulls));
let res = rank(&a, None).unwrap();
assert_eq!(res, &[5, 5, 1, 3, 3]);
}
#[test]
fn test_booleans() {
let descending = SortOptions {
descending: true,
nulls_first: true,
};
let nulls_last = SortOptions {
descending: false,
nulls_first: false,
};
let nulls_last_descending = SortOptions {
descending: true,
nulls_first: false,
};
let a = BooleanArray::from(vec![true, false, false, false, true]);
let res = rank(&a, None).unwrap();
assert_eq!(res, &[5, 3, 3, 3, 5]);
let res = rank(&a, Some(descending)).unwrap();
assert_eq!(res, &[2, 5, 5, 5, 2]);
let res = rank(&a, Some(nulls_last)).unwrap();
assert_eq!(res, &[5, 3, 3, 3, 5]);
let res = rank(&a, Some(nulls_last_descending)).unwrap();
assert_eq!(res, &[2, 5, 5, 5, 2]);
}
#[test]
fn test_bytes() {
let v = vec!["foo", "fo", "bar", "bar"];
let values = StringArray::from(v.clone());
let res = rank(&values, None).unwrap();
assert_eq!(res, &[4, 3, 2, 2]);
let values = LargeStringArray::from(v.clone());
let res = rank(&values, None).unwrap();
assert_eq!(res, &[4, 3, 2, 2]);
let v: Vec<&[u8]> = vec![&[1, 2], &[0], &[1, 2, 3], &[1, 2]];
let values = LargeBinaryArray::from(v.clone());
let res = rank(&values, None).unwrap();
assert_eq!(res, &[3, 1, 4, 3]);
let values = BinaryArray::from(v);
let res = rank(&values, None).unwrap();
assert_eq!(res, &[3, 1, 4, 3]);
}
}
File diff suppressed because it is too large Load Diff