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
1 change: 1 addition & 0 deletions cpp/src/arrow/compute/kernels/scalar_cast_internal.cc
Original file line number Diff line number Diff line change
Expand Up @@ -113,6 +113,7 @@ struct CastPrimitive<HalfFloatType, InType, enable_if_integer<InType>> {
// Cast half float to int
template <typename OutType>
struct CastPrimitive<OutType, HalfFloatType, enable_if_integer<OutType>> {
ARROW_DISABLE_UBSAN("float-cast-overflow")
static void Exec(const ArraySpan& arr, ArraySpan* out) {
using OutT = typename OutType::c_type;
const uint16_t* in_values = arr.GetValues<uint16_t>(1);
Expand Down
86 changes: 63 additions & 23 deletions cpp/src/arrow/compute/kernels/scalar_cast_numeric.cc
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,8 @@

// Implementation of casting to integer, floating point, or decimal types

#include <cmath>

#include "arrow/array/builder_primitive.h"
#include "arrow/compute/kernels/common_internal.h"
#include "arrow/compute/kernels/scalar_cast_internal.h"
Expand Down Expand Up @@ -87,12 +89,48 @@ struct WasTruncated<HalfFloatType, OutType> {
}
};

// InType is a floating point type we are planning to cast to integer
template <typename InType, typename OutType, typename InT = typename InType::c_type,
typename OutT = typename OutType::c_type>
struct WasOutOfRange {
// Both bounds are zero or a power of two, so they are exact in InT
static constexpr auto kMin = static_cast<InT>(std::numeric_limits<OutT>::min());
static constexpr auto kMaxPlusOne =
static_cast<InT>(std::numeric_limits<OutT>::max() / 2 + 1) * 2;

static bool Check(OutT, InT in_val) {
return !(std::trunc(in_val) >= kMin) | !(in_val < kMaxPlusOne);
}

static bool CheckMaybeNull(OutT out_val, InT in_val, bool is_valid) {
return is_valid && Check(out_val, in_val);
}
};

template <typename OutType>
struct WasOutOfRange<HalfFloatType, OutType> {
using OutT = typename OutType::c_type;
static bool Check(OutT out_val, uint16_t in_val) {
return WasOutOfRange<FloatType, OutType>::Check(out_val,
Float16::FromBits(in_val).ToFloat());
}

static bool CheckMaybeNull(OutT out_val, uint16_t in_val, bool is_valid) {
return is_valid && Check(out_val, in_val);
}
};

// InType is a floating point type we are planning to cast to integer
template <template <typename...> class Checker, typename InType, typename OutType,
typename InT = typename InType::c_type,
typename OutT = typename OutType::c_type>
ARROW_DISABLE_UBSAN("float-cast-overflow")
Status CheckFloatTruncation(const ArraySpan& input, const ArraySpan& output) {
Status CheckFloatToIntValues(const ArraySpan& input, const ArraySpan& output) {
auto GetErrorMessage = [&](InT val) {
if constexpr (std::is_same_v<Checker<InType, OutType>,
WasOutOfRange<InType, OutType>>) {
return Status::Invalid("Float value ", val, " out of range converting to ",
*output.type);
}
return Status::Invalid("Float value ", val, " was truncated converting to ",
*output.type);
};
Expand All @@ -110,28 +148,27 @@ Status CheckFloatTruncation(const ArraySpan& input, const ArraySpan& output) {
if (block.popcount == block.length) {
// Fast path: branchless
for (int64_t i = 0; i < block.length; ++i) {
block_out_of_bounds |=
WasTruncated<InType, OutType>::Check(out_data[i], in_data[i]);
block_out_of_bounds |= Checker<InType, OutType>::Check(out_data[i], in_data[i]);
}
} else if (block.popcount > 0) {
// Indices have nulls, must only boundscheck non-null values
for (int64_t i = 0; i < block.length; ++i) {
block_out_of_bounds |= WasTruncated<InType, OutType>::CheckMaybeNull(
block_out_of_bounds |= Checker<InType, OutType>::CheckMaybeNull(
out_data[i], in_data[i], bit_util::GetBit(bitmap, offset_position + i));
}
}
if (ARROW_PREDICT_FALSE(block_out_of_bounds)) {
if (input.GetNullCount() > 0) {
for (int64_t i = 0; i < block.length; ++i) {
if (WasTruncated<InType, OutType>::CheckMaybeNull(
if (Checker<InType, OutType>::CheckMaybeNull(
out_data[i], in_data[i],
bit_util::GetBit(bitmap, offset_position + i))) {
return GetErrorMessage(in_data[i]);
}
}
} else {
for (int64_t i = 0; i < block.length; ++i) {
if (WasTruncated<InType, OutType>::Check(out_data[i], in_data[i])) {
if (Checker<InType, OutType>::Check(out_data[i], in_data[i])) {
return GetErrorMessage(in_data[i]);
}
}
Expand All @@ -145,41 +182,42 @@ Status CheckFloatTruncation(const ArraySpan& input, const ArraySpan& output) {
return Status::OK();
}

template <typename InType>
Status CheckFloatToIntTruncationImpl(const ArraySpan& input, const ArraySpan& output) {
template <template <typename...> class Checker, typename InType>
Status CheckFloatToIntImpl(const ArraySpan& input, const ArraySpan& output) {
switch (output.type->id()) {
case Type::INT8:
return CheckFloatTruncation<InType, Int8Type>(input, output);
return CheckFloatToIntValues<Checker, InType, Int8Type>(input, output);
case Type::INT16:
return CheckFloatTruncation<InType, Int16Type>(input, output);
return CheckFloatToIntValues<Checker, InType, Int16Type>(input, output);
case Type::INT32:
return CheckFloatTruncation<InType, Int32Type>(input, output);
return CheckFloatToIntValues<Checker, InType, Int32Type>(input, output);
case Type::INT64:
return CheckFloatTruncation<InType, Int64Type>(input, output);
return CheckFloatToIntValues<Checker, InType, Int64Type>(input, output);
case Type::UINT8:
return CheckFloatTruncation<InType, UInt8Type>(input, output);
return CheckFloatToIntValues<Checker, InType, UInt8Type>(input, output);
case Type::UINT16:
return CheckFloatTruncation<InType, UInt16Type>(input, output);
return CheckFloatToIntValues<Checker, InType, UInt16Type>(input, output);
case Type::UINT32:
return CheckFloatTruncation<InType, UInt32Type>(input, output);
return CheckFloatToIntValues<Checker, InType, UInt32Type>(input, output);
case Type::UINT64:
return CheckFloatTruncation<InType, UInt64Type>(input, output);
return CheckFloatToIntValues<Checker, InType, UInt64Type>(input, output);
default:
break;
}
DCHECK(false);
return Status::OK();
}

Status CheckFloatToIntTruncation(const ExecValue& input, const ExecResult& output) {
template <template <typename...> class Checker>
Status CheckFloatToInt(const ExecValue& input, const ExecResult& output) {
switch (input.type()->id()) {
case Type::FLOAT:
return CheckFloatToIntTruncationImpl<FloatType>(input.array, *output.array_span());
return CheckFloatToIntImpl<Checker, FloatType>(input.array, *output.array_span());
case Type::DOUBLE:
return CheckFloatToIntTruncationImpl<DoubleType>(input.array, *output.array_span());
return CheckFloatToIntImpl<Checker, DoubleType>(input.array, *output.array_span());
case Type::HALF_FLOAT:
return CheckFloatToIntTruncationImpl<HalfFloatType>(input.array,
*output.array_span());
return CheckFloatToIntImpl<Checker, HalfFloatType>(input.array,
*output.array_span());
default:
break;
}
Expand All @@ -192,7 +230,9 @@ Status CastFloatingToInteger(KernelContext* ctx, const ExecSpan& batch, ExecResu
CastNumberToNumberUnsafe(batch[0].type()->id(), out->type()->id(), batch[0].array,
out->array_span_mutable());
if (!options.allow_float_truncate) {
RETURN_NOT_OK(CheckFloatToIntTruncation(batch[0], *out));
RETURN_NOT_OK(CheckFloatToInt<WasTruncated>(batch[0], *out));
} else if (!options.allow_int_overflow) {
RETURN_NOT_OK(CheckFloatToInt<WasOutOfRange>(batch[0], *out));
}
return Status::OK();
}
Expand Down
36 changes: 36 additions & 0 deletions cpp/src/arrow/compute/kernels/scalar_cast_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -398,6 +398,42 @@ TEST(Cast, FloatingToInt) {
}
}

TEST(Cast, FloatingToIntOverflow) {
for (auto from : {float16(), float32(), float64()}) {
for (bool allow_float_truncate : {false, true}) {
auto opts = CastOptions::Safe(int8());
opts.allow_float_truncate = allow_float_truncate;
for (auto json : {"[128.0]", "[-129.0]", "[NaN]", "[Inf]", "[-Inf]"}) {
CheckCastFails(ArrayFromJSON(from, json), opts);
}
CheckCast(ArrayFromJSON(from, "[127.0, null, -128.0]"),
ArrayFromJSON(int8(), "[127, null, -128]"), opts);

opts.to_type = uint8();
CheckCastFails(ArrayFromJSON(from, "[256.0]"), opts);
CheckCastFails(ArrayFromJSON(from, "[-1.0]"), opts);
}

auto opts = CastOptions::Safe(int8());
opts.allow_float_truncate = true;
CheckCast(ArrayFromJSON(from, "[127.5, -128.5]"),
ArrayFromJSON(int8(), "[127, -128]"), opts);
opts.to_type = uint8();
CheckCast(ArrayFromJSON(from, "[255.5, -0.5]"), ArrayFromJSON(uint8(), "[255, 0]"),
opts);
}

for (auto from : {float32(), float64()}) {
auto opts = CastOptions::Safe(int64());
opts.allow_float_truncate = true;
CheckCastFails(ArrayFromJSON(from, "[9223372036854775808.0]"), opts);
CheckCast(ArrayFromJSON(from, "[-9223372036854775808.0]"),
ArrayFromJSON(int64(), "[-9223372036854775808]"), opts);
opts.to_type = uint64();
CheckCastFails(ArrayFromJSON(from, "[18446744073709551616.0]"), opts);
}
}

TEST(Cast, FloatingToFloating) {
for (auto from : {float16(), float32(), float64()}) {
for (auto to : {float16(), float32(), float64()}) {
Expand Down
2 changes: 1 addition & 1 deletion r/R/expression.R
Original file line number Diff line number Diff line change
Expand Up @@ -211,7 +211,7 @@ Expression$op <- function(FUN, ..., args = list(...)) {
"if_else",
Expression$op("==", args[[2]], 0L),
Scalar$create(NA_integer_, out_type),
cast(out, out_type, allow_float_truncate = TRUE)
cast(out, out_type, allow_float_truncate = TRUE, allow_int_overflow = TRUE)
)
}
return(out)
Expand Down
Loading