This is an automated email from the ASF dual-hosted git repository.
alamb pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/arrow-rs.git
The following commit(s) were added to refs/heads/main by this push:
new 46d5439abc perf: Optimize infallible numeric casts (#11219)
46d5439abc is described below
commit 46d5439abc0698847eab07181661d380d083d880
Author: Neil Conway <[email protected]>
AuthorDate: Wed Sep 30 09:23:41 2026 -0400
perf: Optimize infallible numeric casts (#11219)
# Which issue does this PR close?
- Closes: #10955
- Closes https://github.com/apache/arrow-rs/pull/10958
# Rationale for this change
`cast_numeric_arrays` performed a cast between two numeric types using
either `unary_opt` (in safe mode) or `try_unary` (in not-safe mode).
However, both kernels are considerably slower than `unary`: they don't
vectorize as well, and `unary_opt` also needs to allocate a null buffer.
For many common pairs of numeric types, casting is infallible, and
therefore the additional overhead of using `unary_opt` / `try_unary` is
unnecessary. Because `unary` applies the conversion to null slots as
well, the fast path is limited to type pairs for which `num_cast`
succeeds for every possible source value.
Benchmarks: (M4 Max)
- cast int32 to int64 512: 344.130 → 84.676 ns, −75.4%
- cast int32 to float32 512: 332.420 → 64.278 ns, −80.7%
- cast int32 to float64 512: 372.920 → 87.124 ns, −76.6%
- cast float64 to float32 512: 345.900 → 70.852 ns, −79.5%
- cast int32 to uint32 512: 564.230 → 540.660 ns, −4.2%
- cast int64 to int32 512: 655.640 → 635.060 ns, −3.1%
# What changes are included in this PR?
* Add `is_infallible_numeric_cast`, a compile-time predicate on pairs of
types to determine if a cast between them can fail
* Use `unary` for infallible numeric casts
* Refactor `cast_numeric_arrays` to get rid of two unnecessary helpers
(`try_numeric_cast`, `numeric_cast`), so that the reader can see all of
the numeric cast logic in a single place
* Add unit test
# Are these changes tested?
Yes; existing tests pass, new test added.
# Are there any user-facing changes?
No. The fast-path still uses the same underlying typecast code, the
change here is just in the machinery that invokes it.
# AI Usage
Developed with Claude Code (Fable 5.1), reviewed with Codex (Astra 6). I
have revised and understand the resulting code.
---
arrow-cast/src/cast/mod.rs | 168 +++++++++++++++++++++++++++++++++------------
1 file changed, 126 insertions(+), 42 deletions(-)
diff --git a/arrow-cast/src/cast/mod.rs b/arrow-cast/src/cast/mod.rs
index c02fb900b2..3e562e508d 100644
--- a/arrow-cast/src/cast/mod.rs
+++ b/arrow-cast/src/cast/mod.rs
@@ -2535,40 +2535,90 @@ fn cast_numeric_arrays<FROM, TO>(
where
FROM: ArrowPrimitiveType,
TO: ArrowPrimitiveType,
- FROM::Native: NumCast,
- TO::Native: NumCast,
+ FROM::Native: NumCast + NumericNative,
+ TO::Native: NumCast + NumericNative,
{
- if cast_options.safe {
- // If the value can't be casted to the `TO::Native`, return null
- Ok(Arc::new(numeric_cast::<FROM, TO>(
- from.as_primitive::<FROM>(),
- )))
+ let from = from.as_primitive::<FROM>();
+ let array: PrimitiveArray<TO> = if const {
+ is_infallible_numeric_cast(
+ <FROM::Native as NumericNative>::KIND,
+ <TO::Native as NumericNative>::KIND,
+ )
+ } {
+ // This cast cannot fail, so the fastest kernel, `unary`, can be used.
+ from.unary(|v| num_cast(v).expect("numeric cast is infallible"))
+ } else if cast_options.safe {
+ // If the value can't be cast to the `TO::Native`, return null
+ from.unary_opt(num_cast)
} else {
- // If the value can't be casted to the `TO::Native`, return error
- Ok(Arc::new(try_numeric_cast::<FROM, TO>(
- from.as_primitive::<FROM>(),
- )?))
- }
+ // If the value can't be cast to the `TO::Native`, return error
+ from.try_unary(|v| {
+ num_cast(v).ok_or_else(|| {
+ ArrowError::CastError(format!("Can't cast value {v:?} to type
{}", TO::DATA_TYPE))
+ })
+ })?
+ };
+ Ok(Arc::new(array))
}
-// Natural cast between numeric types
-// If the value of T can't be casted to R, will throw error
-fn try_numeric_cast<T, R>(from: &PrimitiveArray<T>) ->
Result<PrimitiveArray<R>, ArrowError>
-where
- T: ArrowPrimitiveType,
- R: ArrowPrimitiveType,
- T::Native: NumCast,
- R::Native: NumCast,
-{
- from.try_unary(|value| {
- num_cast::<T::Native, R::Native>(value).ok_or_else(|| {
- ArrowError::CastError(format!(
- "Can't cast value {:?} to type {}",
- value,
- R::DATA_TYPE
- ))
- })
- })
+/// Properties of a numeric native type that decide whether [`num_cast`] to
+/// another numeric type can fail.
+#[derive(Clone, Copy)]
+struct NumericKind {
+ float: bool,
+ signed: bool,
+ bits: u32,
+}
+
+/// A native type with a [`NumericKind`].
+trait NumericNative {
+ const KIND: NumericKind;
+}
+
+macro_rules! numeric_native {
+ ($($t:ty => $float:literal, $signed:literal;)*) => {$(
+ impl NumericNative for $t {
+ const KIND: NumericKind = NumericKind {
+ float: $float,
+ signed: $signed,
+ bits: 8 * std::mem::size_of::<$t>() as u32,
+ };
+ }
+ )*};
+}
+
+numeric_native! {
+ u8 => false, false;
+ u16 => false, false;
+ u32 => false, false;
+ u64 => false, false;
+ i8 => false, true;
+ i16 => false, true;
+ i32 => false, true;
+ i64 => false, true;
+ half::f16 => true, true;
+ f32 => true, true;
+ f64 => true, true;
+}
+
+/// Returns true if [`num_cast`] succeeds for every value of `from` when
casting
+/// to `to`: every cast to a float, and integer casts whose target holds every
+/// source value.
+///
+/// This must hold for every bit pattern of the source type, not only for the
+/// valid values of a particular array, because [`PrimitiveArray::unary`]
applies
+/// the conversion to null slots as well, and their contents are arbitrary.
+const fn is_infallible_numeric_cast(from: NumericKind, to: NumericKind) ->
bool {
+ if to.float {
+ true
+ } else if from.float {
+ false
+ } else if from.signed == to.signed {
+ to.bits >= from.bits
+ } else {
+ // Only an unsigned source fits a strictly wider signed target.
+ !from.signed && to.bits > from.bits
+ }
}
/// Natural cast between numeric types
@@ -2582,18 +2632,6 @@ where
num_traits::cast::cast::<I, O>(value)
}
-// Natural cast between numeric types
-// If the value of T can't be casted to R, it will be converted to null
-fn numeric_cast<T, R>(from: &PrimitiveArray<T>) -> PrimitiveArray<R>
-where
- T: ArrowPrimitiveType,
- R: ArrowPrimitiveType,
- T::Native: NumCast,
- R::Native: NumCast,
-{
- from.unary_opt::<_, R>(num_cast::<T::Native, R::Native>)
-}
-
fn cast_numeric_to_binary<FROM: ArrowPrimitiveType, O: OffsetSizeTrait>(
array: &dyn Array,
) -> Result<ArrayRef, ArrowError> {
@@ -4741,6 +4779,52 @@ mod tests {
);
}
+ /// Checks the fast-path classification against `num_cast` for every pair
of
+ /// numeric types. `num_cast` fails only on values outside the target's
range,
+ /// so each type's extremes decide, plus NaN for floats.
+ #[test]
+ fn test_is_infallible_numeric_cast() {
+ macro_rules! check {
+ ($from:ty => $values:expr, [$($to:ty),*]) => {$(
+ let infallible = is_infallible_numeric_cast(
+ <<$from as ArrowPrimitiveType>::Native as
NumericNative>::KIND,
+ <<$to as ArrowPrimitiveType>::Native as
NumericNative>::KIND,
+ );
+ let succeeds = $values
+ .iter()
+ .all(|&v| num_cast::<_, <$to as
ArrowPrimitiveType>::Native>(v).is_some());
+ assert_eq!(
+ infallible,
+ succeeds,
+ "{} to {}",
+ <$from as ArrowPrimitiveType>::DATA_TYPE,
+ <$to as ArrowPrimitiveType>::DATA_TYPE
+ );
+ )*};
+ }
+ macro_rules! check_all {
+ ($($from:ty => $values:expr),* $(,)?) => {$(
+ check!($from => $values, [
+ UInt8Type, UInt16Type, UInt32Type, UInt64Type, Int8Type,
Int16Type,
+ Int32Type, Int64Type, Float16Type, Float32Type, Float64Type
+ ]);
+ )*};
+ }
+ check_all! {
+ UInt8Type => [u8::MIN, u8::MAX],
+ UInt16Type => [u16::MIN, u16::MAX],
+ UInt32Type => [u32::MIN, u32::MAX],
+ UInt64Type => [u64::MIN, u64::MAX],
+ Int8Type => [i8::MIN, i8::MAX],
+ Int16Type => [i16::MIN, i16::MAX],
+ Int32Type => [i32::MIN, i32::MAX],
+ Int64Type => [i64::MIN, i64::MAX],
+ Float16Type => [f16::MIN, f16::MAX, f16::NAN],
+ Float32Type => [f32::MIN, f32::MAX, f32::NAN],
+ Float64Type => [f64::MIN, f64::MAX, f64::NAN],
+ }
+ }
+
#[test]
fn test_cast_i32_to_f64() {
let array = Int32Array::from(vec![5, 6, 7, 8, 9]);