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
14 changes: 14 additions & 0 deletions cpp/src/arrow/compute/kernels/codegen_internal.h
Original file line number Diff line number Diff line change
Expand Up @@ -1337,6 +1337,20 @@ auto GenerateTypeAgnosticVarBinaryBase(detail::GetTypeId get_id) {
}
}

// Similar to GenerateTypeAgnosticVarBinaryBase, but for binary view types
template <template <typename...> class Generator, typename... Args>
auto GenerateTypeAgnosticVarBinaryViewBase(detail::GetTypeId get_id) {
using KernelType = decltype(&Generator<BinaryViewType, Args...>::Exec);
switch (get_id.id) {
case Type::BINARY_VIEW:
case Type::STRING_VIEW:
return Generator<BinaryViewType, Args...>::Exec;
default:
ARROW_DCHECK(false);
return FailFunctor<KernelType>::Exec;
}
}

// Generate a kernel given a templated functor for base binary types. Generates
// a single kernel for binary/string and large binary/large string. If your kernel
// implementation needs access to the specific type at compile time, please use
Expand Down
133 changes: 94 additions & 39 deletions cpp/src/arrow/compute/kernels/scalar_string_ascii.cc
Original file line number Diff line number Diff line change
Expand Up @@ -885,6 +885,12 @@ void AddAsciiStringLength(FunctionRegistry* registry) {
ty);
DCHECK_OK(func->AddKernel({ty}, int64(), std::move(exec)));
}
// View element length is int32-sized, so both view types emit int32.
for (const auto& ty : BinaryViewTypes()) {
auto exec = GenerateVarBinaryViewBase<applicator::ScalarUnaryNotNull, Int32Type,
BinaryLength>(*ty);
DCHECK_OK(func->AddKernel({ty}, int32(), std::move(exec)));
}
DCHECK_OK(func->AddKernel({InputType(Type::FIXED_SIZE_BINARY)}, int32(),
BinaryLength::FixedSizeExec));
DCHECK_OK(registry->AddFunction(std::move(func)));
Expand Down Expand Up @@ -1198,24 +1204,6 @@ void AddAsciiStringPad(FunctionRegistry* registry) {
// ----------------------------------------------------------------------
// Exact pattern detection

using StrToBoolTransformFunc =
std::function<void(const void*, const uint8_t*, int64_t, int64_t, uint8_t*)>;

// Apply `transform` to input character data- this function cannot change the
// length
template <typename Type>
void StringBoolTransform(KernelContext* ctx, const ExecSpan& batch,
StrToBoolTransformFunc transform, ExecResult* out) {
using offset_type = typename Type::offset_type;
const ArraySpan& input = batch[0].array;
ArraySpan* out_arr = out->array_span_mutable();
if (input.length > 0) {
transform(reinterpret_cast<const offset_type*>(input.buffers[1].data) + input.offset,
input.buffers[2].data, input.length, out_arr->offset,
out_arr->buffers[1].data);
}
}

using MatchSubstringState = OptionsWrapper<MatchSubstringOptions>;

// This is an implementation of the Knuth-Morris-Pratt algorithm
Expand Down Expand Up @@ -1333,27 +1321,26 @@ struct RegexSubstringMatcher {

template <typename Type, typename Matcher>
struct MatchSubstringImpl {
using offset_type = typename Type::offset_type;

static Status Exec(KernelContext* ctx, const ExecSpan& batch, ExecResult* out,
const Matcher* matcher) {
StringBoolTransform<Type>(
ctx, batch,
[&matcher](const void* raw_offsets, const uint8_t* data, int64_t length,
int64_t output_offset, uint8_t* output) {
const offset_type* offsets = reinterpret_cast<const offset_type*>(raw_offsets);
FirstTimeBitmapWriter bitmap_writer(output, output_offset, length);
for (int64_t i = 0; i < length; ++i) {
const char* current_data = reinterpret_cast<const char*>(data + offsets[i]);
int64_t current_length = offsets[i + 1] - offsets[i];
if (matcher->Match(std::string_view(current_data, current_length))) {
bitmap_writer.Set();
}
bitmap_writer.Next();
// Evaluate per element and skip null slots. The matcher may be expensive
// (e.g. a regex), so running it on nulls is wasteful; additionally, for the
// view types a null slot's header is unvalidated and may carry a bogus
// buffer_index/offset that decoding would dereference.
const ArraySpan& input = batch[0].array;
ArraySpan* out_arr = out->array_span_mutable();
FirstTimeBitmapWriter bitmap_writer(out_arr->buffers[1].data, out_arr->offset,
input.length);
VisitArrayValuesInline<Type>(
input,
[&](std::string_view val) {
if (matcher->Match(val)) {
bitmap_writer.Set();
}
bitmap_writer.Finish();
bitmap_writer.Next();
},
out);
[&]() { bitmap_writer.Next(); });
bitmap_writer.Finish();
return Status::OK();
}
};
Expand Down Expand Up @@ -1645,6 +1632,13 @@ void AddAsciiStringMatchSubstring(FunctionRegistry* registry) {
DCHECK_OK(
func->AddKernel({ty}, boolean(), std::move(exec), MatchSubstringState::Init));
}
for (const auto& ty : BinaryViewTypes()) {
auto exec =
GenerateTypeAgnosticVarBinaryViewBase<MatchSubstring, PlainSubstringMatcher>(
ty);
DCHECK_OK(
func->AddKernel({ty}, boolean(), std::move(exec), MatchSubstringState::Init));
}
DCHECK_OK(registry->AddFunction(std::move(func)));
}
{
Expand All @@ -1656,6 +1650,13 @@ void AddAsciiStringMatchSubstring(FunctionRegistry* registry) {
DCHECK_OK(
func->AddKernel({ty}, boolean(), std::move(exec), MatchSubstringState::Init));
}
for (const auto& ty : BinaryViewTypes()) {
auto exec =
GenerateTypeAgnosticVarBinaryViewBase<MatchSubstring, PlainStartsWithMatcher>(
ty);
DCHECK_OK(
func->AddKernel({ty}, boolean(), std::move(exec), MatchSubstringState::Init));
}
DCHECK_OK(registry->AddFunction(std::move(func)));
}
{
Expand All @@ -1667,6 +1668,12 @@ void AddAsciiStringMatchSubstring(FunctionRegistry* registry) {
DCHECK_OK(
func->AddKernel({ty}, boolean(), std::move(exec), MatchSubstringState::Init));
}
for (const auto& ty : BinaryViewTypes()) {
auto exec =
GenerateTypeAgnosticVarBinaryViewBase<MatchSubstring, PlainEndsWithMatcher>(ty);
DCHECK_OK(
func->AddKernel({ty}, boolean(), std::move(exec), MatchSubstringState::Init));
}
DCHECK_OK(registry->AddFunction(std::move(func)));
}
#ifdef ARROW_WITH_RE2
Expand All @@ -1679,6 +1686,13 @@ void AddAsciiStringMatchSubstring(FunctionRegistry* registry) {
DCHECK_OK(
func->AddKernel({ty}, boolean(), std::move(exec), MatchSubstringState::Init));
}
for (const auto& ty : BinaryViewTypes()) {
auto exec =
GenerateTypeAgnosticVarBinaryViewBase<MatchSubstring, RegexSubstringMatcher>(
ty);
DCHECK_OK(
func->AddKernel({ty}, boolean(), std::move(exec), MatchSubstringState::Init));
}
DCHECK_OK(registry->AddFunction(std::move(func)));
}
{
Expand All @@ -1689,6 +1703,11 @@ void AddAsciiStringMatchSubstring(FunctionRegistry* registry) {
DCHECK_OK(
func->AddKernel({ty}, boolean(), std::move(exec), MatchSubstringState::Init));
}
for (const auto& ty : BinaryViewTypes()) {
auto exec = GenerateTypeAgnosticVarBinaryViewBase<MatchLike>(ty);
DCHECK_OK(
func->AddKernel({ty}, boolean(), std::move(exec), MatchSubstringState::Init));
}
DCHECK_OK(registry->AddFunction(std::move(func)));
}
#endif
Expand All @@ -1697,6 +1716,17 @@ void AddAsciiStringMatchSubstring(FunctionRegistry* registry) {
// ----------------------------------------------------------------------
// Substring find - lfind/index/etc.

// The output offset type for find/count kernels. Like TypeTraits<T>::OffsetType,
// but also defined for the view types, whose per-element length is int32-sized.
template <typename T, typename Enable = void>
struct StringOffsetType {
using type = typename TypeTraits<T>::OffsetType;
};
template <typename T>
struct StringOffsetType<T, enable_if_binary_view_like<T>> {
using type = Int32Type;
};

struct FindSubstring {
const PlainSubstringMatcher matcher_;

Expand Down Expand Up @@ -1743,7 +1773,9 @@ struct FindSubstringRegex {

template <typename InputPhysicalType>
struct FindSubstringExec {
using OffsetType = typename TypeTraits<InputPhysicalType>::OffsetType;
// StringOffsetType maps the view types (which have no TypeTraits::OffsetType)
// to int32; for the other physical types it forwards to TypeTraits::OffsetType.
using OffsetType = typename StringOffsetType<InputPhysicalType>::type;

static_assert(!is_string_or_string_view(InputPhysicalType::type_id),
"should only codegen on physical types");
Expand Down Expand Up @@ -1779,7 +1811,7 @@ const FunctionDoc find_substring_doc(
#ifdef ARROW_WITH_RE2
template <typename InputPhysicalType>
struct FindSubstringRegexExec {
using OffsetType = typename TypeTraits<InputPhysicalType>::OffsetType;
using OffsetType = typename StringOffsetType<InputPhysicalType>::type;

static_assert(!is_string_or_string_view(InputPhysicalType::type_id),
"should only codegen on physical types");
Expand Down Expand Up @@ -1813,6 +1845,12 @@ void AddAsciiStringFindSubstring(FunctionRegistry* registry) {
GenerateTypeAgnosticVarBinaryBase<FindSubstringExec>(ty),
MatchSubstringState::Init));
}
// View element length is int32-sized, so both view types emit int32.
for (const auto& ty : BinaryViewTypes()) {
DCHECK_OK(func->AddKernel(
{ty}, int32(), GenerateTypeAgnosticVarBinaryViewBase<FindSubstringExec>(ty),
MatchSubstringState::Init));
}
DCHECK_OK(func->AddKernel({InputType(Type::FIXED_SIZE_BINARY)}, int32(),
FindSubstringExec<FixedSizeBinaryType>::Exec,
MatchSubstringState::Init));
Expand All @@ -1829,6 +1867,12 @@ void AddAsciiStringFindSubstring(FunctionRegistry* registry) {
GenerateTypeAgnosticVarBinaryBase<FindSubstringRegexExec>(ty),
MatchSubstringState::Init));
}
for (const auto& ty : BinaryViewTypes()) {
DCHECK_OK(func->AddKernel(
{ty}, int32(),
GenerateTypeAgnosticVarBinaryViewBase<FindSubstringRegexExec>(ty),
MatchSubstringState::Init));
}
DCHECK_OK(func->AddKernel({InputType(Type::FIXED_SIZE_BINARY)}, int32(),
FindSubstringRegexExec<FixedSizeBinaryType>::Exec,
MatchSubstringState::Init));
Expand Down Expand Up @@ -1902,7 +1946,7 @@ struct CountSubstringRegex {

template <typename InputType>
struct CountSubstringRegexExec {
using OffsetType = typename TypeTraits<InputType>::OffsetType;
using OffsetType = typename StringOffsetType<InputType>::type;
static Status Exec(KernelContext* ctx, const ExecSpan& batch, ExecResult* out) {
const MatchSubstringOptions& options = MatchSubstringState::Get(ctx);
const bool is_utf8 = is_string_or_string_view(batch[0].type()->id());
Expand All @@ -1916,7 +1960,7 @@ struct CountSubstringRegexExec {

template <typename InputType>
struct CountSubstringExec {
using OffsetType = typename TypeTraits<InputType>::OffsetType;
using OffsetType = typename StringOffsetType<InputType>::type;
static Status Exec(KernelContext* ctx, const ExecSpan& batch, ExecResult* out) {
const MatchSubstringOptions& options = MatchSubstringState::Get(ctx);
if (options.ignore_case) {
Expand Down Expand Up @@ -1963,6 +2007,11 @@ void AddAsciiStringCountSubstring(FunctionRegistry* registry) {
GenerateTypeAgnosticVarBinaryBase<CountSubstringExec>(ty),
MatchSubstringState::Init));
}
for (const auto& ty : BinaryViewTypes()) {
DCHECK_OK(func->AddKernel(
{ty}, int32(), GenerateTypeAgnosticVarBinaryViewBase<CountSubstringExec>(ty),
MatchSubstringState::Init));
}
DCHECK_OK(func->AddKernel({InputType(Type::FIXED_SIZE_BINARY)}, int32(),
CountSubstringExec<FixedSizeBinaryType>::Exec,
MatchSubstringState::Init));
Expand All @@ -1979,6 +2028,12 @@ void AddAsciiStringCountSubstring(FunctionRegistry* registry) {
GenerateTypeAgnosticVarBinaryBase<CountSubstringRegexExec>(ty),
MatchSubstringState::Init));
}
for (const auto& ty : BinaryViewTypes()) {
DCHECK_OK(func->AddKernel(
{ty}, int32(),
GenerateTypeAgnosticVarBinaryViewBase<CountSubstringRegexExec>(ty),
MatchSubstringState::Init));
}
DCHECK_OK(func->AddKernel({InputType(Type::FIXED_SIZE_BINARY)}, int32(),
CountSubstringRegexExec<FixedSizeBinaryType>::Exec,
MatchSubstringState::Init));
Expand Down
37 changes: 30 additions & 7 deletions cpp/src/arrow/compute/kernels/scalar_string_internal.h
Original file line number Diff line number Diff line change
Expand Up @@ -230,14 +230,33 @@ struct StringPredicateFunctor {
Status st = Status::OK();
EnsureUtf8LookupTablesFilled();
const ArraySpan& input = batch[0].array;
ArrayIterator<Type> input_it(input);
ArraySpan* out_arr = out->array_span_mutable();
::arrow::internal::GenerateBitsUnrolled(
out_arr->buffers[1].data, out_arr->offset, input.length, [&]() -> bool {
std::string_view val = input_it();
return Predicate::Call(ctx, reinterpret_cast<const uint8_t*>(val.data()),
val.size(), &st);
});
if constexpr (is_binary_view_like_type<Type>::value) {
// Skip null slots: a view's null header is not validated and may carry a
// bogus buffer_index/offset that decoding would dereference. (The base
// binary/string path below can safely read null slots.)
FirstTimeBitmapWriter writer(out_arr->buffers[1].data, out_arr->offset,
input.length);
VisitArrayValuesInline<Type>(
input,
[&](std::string_view val) {
if (Predicate::Call(ctx, reinterpret_cast<const uint8_t*>(val.data()),
val.size(), &st)) {
writer.Set();
}
writer.Next();
},
[&]() { writer.Next(); });
writer.Finish();
} else {
ArrayIterator<Type> input_it(input);
::arrow::internal::GenerateBitsUnrolled(
out_arr->buffers[1].data, out_arr->offset, input.length, [&]() -> bool {
std::string_view val = input_it();
return Predicate::Call(ctx, reinterpret_cast<const uint8_t*>(val.data()),
val.size(), &st);
});
}
return st;
}
};
Expand All @@ -250,6 +269,10 @@ void AddUnaryStringPredicate(std::string name, FunctionRegistry* registry,
auto exec = GenerateTypeAgnosticVarBinaryBase<StringPredicateFunctor, Predicate>(ty);
ARROW_DCHECK_OK(func->AddKernel({ty}, boolean(), std::move(exec)));
}
ARROW_DCHECK_OK(func->AddKernel(
{utf8_view()}, boolean(),
GenerateTypeAgnosticVarBinaryViewBase<StringPredicateFunctor, Predicate>(
utf8_view())));
ARROW_DCHECK_OK(registry->AddFunction(std::move(func)));
}

Expand Down
Loading
Loading