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
236 changes: 236 additions & 0 deletions cpp/src/arrow/compute/kernels/scalar_nested.cc
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
#include <cmath>
#include "arrow/array/array_base.h"
#include "arrow/array/builder_nested.h"
#include "arrow/array/builder_primitive.h"
#include "arrow/compute/api_scalar.h"
#include "arrow/compute/kernels/common_internal.h"
#include "arrow/compute/registry_internal.h"
Expand All @@ -28,8 +29,11 @@
#include "arrow/util/bit_block_counter.h"
#include "arrow/util/bit_util.h"
#include "arrow/util/bitmap_generate.h"
#include "arrow/util/bitmap_ops.h"
#include "arrow/util/float16.h"
#include "arrow/util/logging_internal.h"
#include "arrow/util/string.h"
#include "arrow/util/ubsan.h"
#include "arrow/util/unreachable.h"

namespace arrow {
Expand Down Expand Up @@ -537,6 +541,233 @@ const FunctionDoc list_element_doc(
"is emitted. Null values emit a null in the output."),
{"lists", "index"});

bool IsNaN(const Scalar& value) {
switch (value.type->id()) {
case Type::HALF_FLOAT:
return util::Float16::FromBits(checked_cast<const HalfFloatScalar&>(value).value)
.is_nan();
case Type::FLOAT:
return std::isnan(checked_cast<const FloatScalar&>(value).value);
case Type::DOUBLE:
return std::isnan(checked_cast<const DoubleScalar&>(value).value);
default:
return false;
}
}

// Null values are matched without "equal", but must be comparable all the same
Status CheckNullValueType(KernelContext* ctx, const DataType& values_type,
const DataType& value_type) {
if (value_type.id() == Type::NA || value_type.Equals(values_type)) {
return Status::OK();
}
ARROW_ASSIGN_OR_RAISE(auto equal,
ctx->exec_context()->func_registry()->GetFunction("equal"));
std::vector<TypeHolder> types{&values_type, &value_type};
return equal->DispatchBest(&types).status();
}

// Match `values` with a scalar `value`, or element-wise with an array of them. Like
// "is_in", and unlike "equal", a null value matches null values and a NaN value
// matches NaN values.
Result<Datum> ListValuesMatch(KernelContext* ctx, const Datum& values,
const Datum& value) {
ExecContext* exec_ctx = ctx->exec_context();
if (value.null_count() == value.length()) {
RETURN_NOT_OK(CheckNullValueType(ctx, *values.type(), *value.type()));
return CallFunction("is_null", {values}, exec_ctx);
}
if (value.is_scalar()) {
if (is_floating(values.type()->id()) && IsNaN(*value.scalar())) {
return CallFunction("is_nan", {values}, exec_ctx);
}
return CallFunction("equal", {values, value}, exec_ctx);
}
ARROW_ASSIGN_OR_RAISE(Datum match, CallFunction("equal", {values, value}, exec_ctx));
ARROW_ASSIGN_OR_RAISE(Datum values_null, CallFunction("is_null", {values}, exec_ctx));
ARROW_ASSIGN_OR_RAISE(Datum value_null, CallFunction("is_null", {value}, exec_ctx));
ARROW_ASSIGN_OR_RAISE(Datum both_null,
CallFunction("and", {values_null, value_null}, exec_ctx));
ARROW_ASSIGN_OR_RAISE(match, CallFunction("or_kleene", {match, both_null}, exec_ctx));
if (is_floating(values.type()->id()) && is_floating(value.type()->id())) {
ARROW_ASSIGN_OR_RAISE(Datum values_nan, CallFunction("is_nan", {values}, exec_ctx));
ARROW_ASSIGN_OR_RAISE(Datum value_nan, CallFunction("is_nan", {value}, exec_ctx));
ARROW_ASSIGN_OR_RAISE(Datum both_nan,
CallFunction("and_kleene", {values_nan, value_nan}, exec_ctx));
ARROW_ASSIGN_OR_RAISE(match, CallFunction("or_kleene", {match, both_nan}, exec_ctx));
}
return match;
}

// Returns a function giving the number of child values of list i, with none for null
// list views
template <typename Type>
auto GetListLengths(const ArraySpan& list) {
if constexpr (std::is_same_v<Type, FixedSizeListType>) {
const int64_t width = checked_cast<const FixedSizeListType&>(*list.type).list_size();
return [width](int64_t) { return width; };
} else if constexpr (is_list_view_type<Type>::value) {
const auto* sizes = list.GetValues<typename Type::offset_type>(2);
return
[&list, sizes](int64_t i) -> int64_t { return list.IsValid(i) ? sizes[i] : 0; };
} else {
const auto* offsets = list.GetValues<typename Type::offset_type>(1);
return [offsets](int64_t i) -> int64_t { return offsets[i + 1] - offsets[i]; };
}
}

// Returns the child values of all lists, one list after the other
template <typename Type>
Result<std::shared_ptr<ArrayData>> GetListValues(KernelContext* ctx,
const ArraySpan& list) {
if constexpr (std::is_same_v<Type, FixedSizeListType>) {
const int64_t width = GetListLengths<Type>(list)(0);
return list.child_data[0].ToArrayData()->Slice(list.offset * width,
list.length * width);
} else if constexpr (is_list_view_type<Type>::value) {
// List views may reference child values in any order
typename TypeTraits<Type>::ArrayType list_view(list.ToArrayData());
ARROW_ASSIGN_OR_RAISE(auto values, list_view.Flatten(ctx->memory_pool()));
return values->data();
} else {
const auto* offsets = list.GetValues<typename Type::offset_type>(1);
return list.child_data[0].ToArrayData()->Slice(offsets[0],
offsets[list.length] - offsets[0]);
}
}

// Whether any of the `length` bits from `offset` is set in `bits`, a bitmap of `size`
// bytes
bool AnySet(const uint8_t* bits, int64_t size, int64_t offset, int64_t length) {
const int64_t byte_offset = offset / 8;
const int64_t bit_offset = offset % 8;
if (bit_offset + length <= 64 && byte_offset + 8 <= size) {
// Short ranges fit in a single word
const uint64_t word =
bit_util::FromLittleEndian(util::SafeLoadAs<uint64_t>(bits + byte_offset));
const uint64_t mask = length == 64 ? ~uint64_t{0} : (uint64_t{1} << length) - 1;
return ((word >> bit_offset) & mask) != 0;
}
arrow::internal::BitBlockCounter counter(bits, offset, length);
for (int64_t position = 0; position < length;) {
const auto block = counter.NextWord();
if (block.popcount > 0) {
return true;
}
position += block.length;
}
return false;
}

// Repeat the value of each list for each of its child values
template <typename Type>
Result<Datum> RepeatListValues(KernelContext* ctx, const ArraySpan& list,
int64_t values_length, const ArraySpan& value) {
const auto list_length = GetListLengths<Type>(list);
Int64Builder indices(ctx->memory_pool());
RETURN_NOT_OK(indices.Reserve(values_length));
for (int64_t i = 0; i < list.length; ++i) {
const int64_t length = list_length(i);
for (int64_t j = 0; j < length; ++j) {
indices.UnsafeAppend(i);
}
}
ARROW_ASSIGN_OR_RAISE(auto repeat_indices, indices.Finish());
return CallFunction("take", {value.ToArrayData(), repeat_indices}, ctx->exec_context());
}

// Match the child values of all lists with a single vectorized comparison, then check
// each list's range in the resulting bitmap.
template <typename Type>
Status ListContains(KernelContext* ctx, const ExecSpan& batch, ExecResult* out) {
if (batch.length == 0) {
return Status::OK();
}
if (batch[0].is_scalar()) {
// A scalar list is only paired with an array of values here, so broadcast it
ARROW_ASSIGN_OR_RAISE(auto lists, MakeArrayFromScalar(*batch[0].scalar, batch.length,
ctx->memory_pool()));
ExecSpan array_batch = batch;
array_batch.values[0].SetArray(*lists->data());
return ListContains<Type>(ctx, array_batch, out);
}

const ArraySpan& list = batch[0].array;
ArraySpan* out_arr = out->array_span_mutable();

// Only null lists emit a null
if (list.MayHaveNulls()) {
arrow::internal::CopyBitmap(list.buffers[0].data, list.offset, list.length,
out_arr->buffers[0].data, out_arr->offset);
} else {
bit_util::SetBitsTo(out_arr->buffers[0].data, out_arr->offset, out_arr->length, true);
}
out_arr->null_count = list.null_count;

ARROW_ASSIGN_OR_RAISE(auto values, GetListValues<Type>(ctx, list));
Datum value;
if (batch[1].is_scalar()) {
value = batch[1].scalar->GetSharedPtr();
} else if (batch[1].array.length == 1) {
// Also covers scalar inputs, which are promoted to arrays
ARROW_ASSIGN_OR_RAISE(value, batch[1].array.ToArray()->GetScalar(0));
} else {
ARROW_ASSIGN_OR_RAISE(
value, RepeatListValues<Type>(ctx, list, values->length, batch[1].array));
}
ARROW_ASSIGN_OR_RAISE(Datum match, ListValuesMatch(ctx, values, value));

// Null matches never count
const ArrayData& match_data = *match.array();
std::shared_ptr<Buffer> matches = match_data.buffers[1];
int64_t matches_offset = match_data.offset;
if (match_data.MayHaveNulls()) {
ARROW_ASSIGN_OR_RAISE(
matches, arrow::internal::BitmapAnd(
ctx->memory_pool(), matches->data(), match_data.offset,
match_data.buffers[0]->data(), match_data.offset, match_data.length,
/*out_offset=*/0));
matches_offset = 0;
}

// The output bits of null lists don't matter, so they are not special-cased
const auto list_length = GetListLengths<Type>(list);
int64_t start = matches_offset;
int64_t i = 0;
arrow::internal::GenerateBitsUnrolled(
out_arr->buffers[1].data, out_arr->offset, out_arr->length, [&] {
const int64_t length = list_length(i++);
const bool found = AnySet(matches->data(), matches->size(), start, length);
start += length;
return found;
});
return Status::OK();
}

void AddListContainsKernels(ScalarFunction* func) {
auto add_kernel = [&](Type::type list_type_id, ArrayKernelExec exec) {
ScalarKernel kernel({InputType(list_type_id), InputType::Any()}, boolean(), exec);
// A null value is searched for rather than propagated
kernel.null_handling = NullHandling::COMPUTED_PREALLOCATE;
DCHECK_OK(func->AddKernel(std::move(kernel)));
};
add_kernel(Type::LIST, ListContains<ListType>);
add_kernel(Type::LARGE_LIST, ListContains<LargeListType>);
add_kernel(Type::LIST_VIEW, ListContains<ListViewType>);
add_kernel(Type::LARGE_LIST_VIEW, ListContains<LargeListViewType>);
add_kernel(Type::FIXED_SIZE_LIST, ListContains<FixedSizeListType>);
}

const FunctionDoc list_contains_doc(
"Check whether lists contain a given value",
("`lists` must have a list-like type and `value` must be comparable\n"
"with the list value type.\n"
"For each list in `lists`, true is emitted if any of its values is equal\n"
"to the corresponding `value`, false otherwise. A null `value` matches\n"
"null list values, and a NaN `value` matches NaN list values; otherwise\n"
"null list values never match. Null lists emit a null in the output."),
{"lists", "value"});

struct StructFieldFunctor {
static Status Exec(KernelContext* ctx, const ExecSpan& batch, ExecResult* out) {
const auto& options = OptionsWrapper<StructFieldOptions>::Get(ctx);
Expand Down Expand Up @@ -960,6 +1191,11 @@ void RegisterScalarNested(FunctionRegistry* registry) {
AddListElementKernels(list_element.get());
DCHECK_OK(registry->AddFunction(std::move(list_element)));

auto list_contains = std::make_shared<ScalarFunction>("list_contains", Arity::Binary(),
list_contains_doc);
AddListContainsKernels(list_contains.get());
DCHECK_OK(registry->AddFunction(std::move(list_contains)));

auto list_slice =
std::make_shared<ScalarFunction>("list_slice", Arity::Unary(), list_slice_doc);
AddListSliceKernels(list_slice.get());
Expand Down
Loading
Loading