Skip to content
Merged
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
40 changes: 40 additions & 0 deletions core/benches/conversions.rs
Original file line number Diff line number Diff line change
Expand Up @@ -196,6 +196,24 @@ fn bench_f64_to_i32(c: &mut Criterion) {
.unwrap();
});
});

// No out_of_range (error on overflow) — all values in range
let config_none = FloatToIntConfig {
map_entries: vec![],
rounding: RoundingMode::NearestEven,
out_of_range: None,
};
group.bench_with_input(BenchmarkId::new("f64_to_i32/no_oor", n), &n, |b, &n| {
let mut dst = vec![0i32; n];
b.iter(|| {
convert_slice_float_to_int(
black_box(&src),
black_box(&mut dst),
black_box(&config_none),
)
.unwrap();
});
});
}
group.finish();
}
Expand Down Expand Up @@ -315,6 +333,28 @@ fn bench_f64_to_f32(c: &mut Criterion) {
.unwrap();
});
});

// Towards-zero rounding, no out_of_range (scalar fallback path)
let config_tz = FloatToFloatConfig {
map_entries: vec![],
rounding: RoundingMode::TowardsZero,
out_of_range: None,
};
group.bench_with_input(
BenchmarkId::new("f64_to_f32/towards_zero", n),
&n,
|b, &n| {
let mut dst = vec![0f32; n];
b.iter(|| {
convert_slice_float_to_float(
black_box(&src),
black_box(&mut dst),
black_box(&config_tz),
)
.unwrap();
});
},
);
}
group.finish();
}
Expand Down
47 changes: 45 additions & 2 deletions core/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -635,6 +635,27 @@ where
}
}

// SIMD fast path: empty scalar_map + no out_of_range (error mode) + supported rounding.
// Range-checking variant: errors if any value is out of range.
if config.map_entries.is_empty()
&& config.out_of_range.is_none()
&& config.rounding != RoundingMode::NearestAway
{
use std::any::TypeId;

// f64 → i32
if TypeId::of::<Src>() == TypeId::of::<f64>() && TypeId::of::<Dst>() == TypeId::of::<i32>()
{
let src_f64: &[f64] =
unsafe { std::slice::from_raw_parts(src.as_ptr() as *const f64, src.len()) };
let dst_i32: &mut [i32] =
unsafe { std::slice::from_raw_parts_mut(dst.as_mut_ptr() as *mut i32, dst.len()) };
if simd::try_f64_to_i32_check(src_f64, dst_i32, config.rounding)? {
return Ok(());
}
}
}

// Scalar fallback
for (in_val, out_slot) in src.iter().zip(dst.iter_mut()) {
*out_slot = convert_float_to_int(*in_val, config)?;
Expand All @@ -659,15 +680,37 @@ where
}

/// Convert a slice of float values to float values. Returns early on first error.
///
/// When the configuration allows it (empty scalar_map, nearest-even rounding),
/// uses SIMD-accelerated kernels for supported type pairs (f64->f32).
pub fn convert_slice_float_to_float<Src, Dst>(
src: &[Src],
dst: &mut [Dst],
config: &FloatToFloatConfig<Src, Dst>,
) -> Result<(), CastError>
where
Src: CastFloat + CastInto<Dst>,
Dst: CastFloat,
Src: CastFloat + CastInto<Dst> + 'static,
Dst: CastFloat + 'static,
{
// SIMD fast path: empty scalar_map + nearest-even rounding + f64→f32.
if config.map_entries.is_empty() && config.rounding == RoundingMode::NearestEven {
use std::any::TypeId;

if TypeId::of::<Src>() == TypeId::of::<f64>() && TypeId::of::<Dst>() == TypeId::of::<f32>()
{
// SAFETY: We just verified Src == f64 and Dst == f32 via TypeId.
let src_f64: &[f64] =
unsafe { std::slice::from_raw_parts(src.as_ptr() as *const f64, src.len()) };
let dst_f32: &mut [f32] =
unsafe { std::slice::from_raw_parts_mut(dst.as_mut_ptr() as *mut f32, dst.len()) };
let error_on_overflow = config.out_of_range != Some(OutOfRangeMode::Clamp);
if simd::try_f64_to_f32_nearest(src_f64, dst_f32, error_on_overflow)? {
return Ok(());
}
}
}

// Scalar fallback
for (in_val, out_slot) in src.iter().zip(dst.iter_mut()) {
*out_slot = convert_float_to_float(*in_val, config)?;
}
Expand Down
151 changes: 151 additions & 0 deletions core/src/simd/aarch64.rs
Original file line number Diff line number Diff line change
Expand Up @@ -284,6 +284,157 @@ pub(super) unsafe fn f32_to_u8_clamp(
Ok(())
}

/// Convert f64 slice to f32 slice using nearest-even rounding.
///
/// NEON's `vcvt_f32_f64` performs the narrowing with nearest-even rounding
/// (the default IEEE 754 mode), processing 2 f64 → 2 f32 per instruction.
///
/// Two-pass approach:
/// 1. Fast convert pass: just `vcvt_f32_f64` + store (no branching).
/// 2. If `error_on_overflow`, a second pass checks for finite→infinite overflow.
///
/// This keeps the hot convert loop branch-free for maximum throughput.
///
/// # Safety
///
/// Caller must ensure this runs on an AArch64 target.
pub(super) unsafe fn f64_to_f32_nearest(
src: &[f64],
dst: &mut [f32],
error_on_overflow: bool,
) -> Result<(), crate::CastError> {
let n = src.len();
let simd_len = n / 2 * 2;

// Pass 1: branch-free narrowing conversion
for i in (0..simd_len).step_by(2) {
let v = vld1q_f64(src.as_ptr().add(i));
let narrowed = vcvt_f32_f64(v);
vst1_f32(dst.as_mut_ptr().add(i), narrowed);
}
// Scalar tail
for i in simd_len..n {
dst[i] = src[i] as f32;
}

// Pass 2: overflow check (only when out_of_range is None)
if error_on_overflow {
let inf_f32 = vdup_n_f32(f32::INFINITY);
for i in (0..simd_len).step_by(2) {
// Check if result is ±Inf
let result = vld1_f32(dst.as_ptr().add(i));
let abs_result = vabs_f32(result);
let result_is_inf = vceq_f32(abs_result, inf_f32);
// Quick reject: if no Inf in result, no overflow possible
let inf_bytes: uint8x8_t = vreinterpret_u8_u32(result_is_inf);
if vmaxv_u8(inf_bytes) != 0 {
// At least one result is Inf — check if source was finite
for (&sv, &dv) in src[i..].iter().zip(dst[i..].iter()).take(2) {
if sv.is_finite() && dv.is_infinite() {
return Err(crate::CastError::OutOfRange {
value: sv,
lo: f32::MIN as f64,
hi: f32::MAX as f64,
});
}
}
}
}
// Tail check
for i in simd_len..n {
if src[i].is_finite() && dst[i].is_infinite() {
return Err(crate::CastError::OutOfRange {
value: src[i],
lo: f32::MIN as f64,
hi: f32::MAX as f64,
});
}
}
}

Ok(())
}

/// Convert f64 slice to i32 slice with rounding, returning an error if any
/// value is out of range (no clamping).
///
/// Same pipeline as `f64_to_i32_clamp` but instead of clamping, we
/// batch-check that all rounded values fall within [i32::MIN, i32::MAX]
/// and error if not.
///
/// # Safety
///
/// Caller must ensure this runs on an AArch64 target.
pub(super) unsafe fn f64_to_i32_check(
src: &[f64],
dst: &mut [i32],
rounding: RoundingMode,
) -> Result<(), crate::CastError> {
let n = src.len();
let simd_len = n / 2 * 2;

let lo = vdupq_n_f64(i32::MIN as f64);
let hi = vdupq_n_f64(i32::MAX as f64);

for i in (0..simd_len).step_by(2) {
let v = vld1q_f64(src.as_ptr().add(i));

// NaN check
if any_nan_f64x2(v) {
for &val in &src[i..std::cmp::min(i + 2, n)] {
if val.is_nan() {
return Err(crate::CastError::NanOrInf { value: val });
}
}
}

let r = round_f64x2(v, rounding);

// Range check: error if any value < lo or > hi
// vcltq_f64 returns all-ones for true, all-zeros for false
let below = vcltq_f64(r, lo);
let above = vcgtq_f64(r, hi);
let out_of_range = vorrq_u64(below, above);
let oor_bytes: uint8x16_t = vreinterpretq_u8_u64(out_of_range);
if vmaxvq_u8(oor_bytes) != 0 {
// Find exact offending element
for &val in src[i..].iter().take(2) {
let rounded = scalar_round_f64(val, rounding);
if rounded < i32::MIN as f64 || rounded > i32::MAX as f64 {
return Err(crate::CastError::OutOfRange {
value: val,
lo: i32::MIN as f64,
hi: i32::MAX as f64,
});
}
}
}

// Convert (values are in range, truncation after rounding is correct)
let i32_val = vmovn_s64(vcvtq_s64_f64(r));
vst1_s32(dst.as_mut_ptr().add(i), i32_val);
}

// Scalar tail
for i in simd_len..n {
let val = src[i];
if val.is_nan() {
return Err(crate::CastError::NanOrInf { value: val });
}
let rounded = scalar_round_f64(val, rounding);
if rounded < i32::MIN as f64 || rounded > i32::MAX as f64 {
return Err(crate::CastError::OutOfRange {
value: val,
lo: i32::MIN as f64,
hi: i32::MAX as f64,
});
}
dst[i] = rounded as i32;
}

Ok(())
}

// ---------------------------------------------------------------------------
// Scalar tail helpers (shared across kernels)
// ---------------------------------------------------------------------------
Expand Down
Loading
Loading