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
298 changes: 261 additions & 37 deletions library/core/src/num/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1589,6 +1589,160 @@ pub const fn can_not_overflow<T>(radix: u32, is_signed_ty: bool, digits: &[u8])
radix <= 16 && digits.len() <= size_of::<T>() * 2 - is_signed_ty as usize
}

/// SIMD-within-a-register helpers for radix-10 parsing.
///
/// Kept outside the integer-type macro so the SWAR implementation is not
/// duplicated for every `FromStr` instantiation, and so it is only emitted
/// for the types that actually dispatch here.
mod decimal_swar {
use super::{IntErrorKind, ParseIntError};

/// Checks if all bytes in `v` are ASCII decimal digits (`b'0'..=b'9'`).
///
/// The per-byte range test is turned into a high-bit test so that all
/// bytes are checked with a single branch:
///
/// - `c - b'0'` wraps around, setting the high bit, for every byte
/// below `'0'` (and for `0xb0..=0xff`, see below).
/// - For the upper bound we need a constant `k` with `'9' + k < 0x80`
/// and `':' + k >= 0x80`, so that a single bit separates the last
/// digit from the first non-digit above it. `0x7f - 0x39 = 0x46` is
/// the unique such constant: `'9' + 0x46 = 0x7f` and `':' + 0x46 =
/// 0x80`. (Adding `'9'` itself would put `':'` at 0x73 with the
/// high bit still clear, and the test would never fire.)
///
/// The addition flags `0x3a..=0xb9`; past that the sum wraps past
/// `0x100`, which leaves the high bit clear again, but those bytes
/// are caught by the subtraction (`c - b'0' >= 0x80`). Between the
/// two tests every non-digit byte is flagged and no digit ever is.
///
/// This is the same check `dec2flt`'s `is_8digits` performs.
#[inline]
pub(super) const fn is_digits(v: usize) -> bool {
let a = v.wrapping_add(usize::repeat_u8(0x46));
let b = v.wrapping_sub(usize::repeat_u8(0x30));
(a | b) & usize::repeat_u8(0x80) == 0
}

/// Parses 8 ASCII decimal digits packed in a `u64` into their numeric
/// value (little-endian: the first digit is the least significant
/// byte).
///
/// Three multiply-shift steps fold neighboring groups together. A
/// step that merges groups of `g` digits multiplies by
/// `10^g * 2^(8g) + 1`, which adds each group's value to the group
/// above it times `10^g`; the shift slides the results down, and the
/// masks keep only the cleanly merged groups (the multiplies also
/// leave overlapping garbage in between):
///
/// - `& 0x0f` strips the `0x3` high nibble of each ASCII digit,
/// leaving groups of one digit each,
/// - `* 2561 >> 8` (`2561 = 10 * 256 + 1`) leaves two-digit values,
/// - `* 6_553_601 >> 16` (`= 100 * 65_536 + 1`) leaves four-digit
/// values,
/// - `* 42_949_672_960_001 >> 32` (`= 10_000 * 2^32 + 1`) leaves the
/// final eight-digit value.
///
/// The caller must ensure all 8 bytes are ASCII digits, e.g. via
/// [`is_digits`].
#[cfg(not(target_pointer_width = "32"))]
#[inline]
pub(super) const fn parse_8digits(v: u64) -> u64 {
let mut v = v;
v = (v & 0x0f0f_0f0f_0f0f_0f0f).wrapping_mul(2561) >> 8;
v = (v & 0x00ff_00ff_00ff_00ff).wrapping_mul(6_553_601) >> 16;
v = (v & 0x0000_ffff_0000_ffff).wrapping_mul(42_949_672_960_001) >> 32;
v
}

/// Parses 4 ASCII decimal digits packed in a `u32` into their numeric
/// value (little-endian: the first digit is the least significant
/// byte).
///
/// The same folding scheme as [`parse_8digits`], stopped one step
/// early since only four digits are needed: `& 0x0f` strips the `0x3`
/// high nibble of each ASCII digit, `* 2561 >> 8` (`2561 =
/// 10 * 256 + 1`) leaves two-digit values, and `* 6_553_601 >> 16`
/// (`= 100 * 65_536 + 1`) leaves the four-digit value.
///
/// The caller must ensure all 4 bytes are ASCII digits, e.g. via
/// [`is_digits`].
#[cfg(target_pointer_width = "32")]
#[inline]
pub(super) const fn parse_4digits(v: u32) -> u32 {
let mut v = v;
v = (v & 0x0f0f_0f0f).wrapping_mul(2561) >> 8;
v = (v & 0x00ff_00ff).wrapping_mul(6_553_601) >> 16;
v
}

/// Parses up to 16 leading decimal digits in batches, returning the
/// accumulated result and the remaining digits.
///
/// The arithmetic is done in `i64` so the same code works for `u64`,
/// `i64`, `u128` and `i128`. 16 decimal digits can never overflow
/// `i64`/`u64`, so the batch arithmetic is safe to run unchecked even
/// in debug builds.
#[cfg(not(target_pointer_width = "32"))]
#[inline]
pub(super) const fn parse_decimal_i64(
is_positive: bool,
mut result: i64,
mut digits: &[u8],
) -> Result<(i64, &[u8]), ParseIntError> {
let mut remaining = 16;

while remaining >= 8 {
let [a, b, c, d, e, f, g, h, rest @ ..] = digits else { break };
let chunk = u64::from_le_bytes([*a, *b, *c, *d, *e, *f, *g, *h]);
if !is_digits(chunk as usize) {
return Err(ParseIntError { kind: IntErrorKind::InvalidDigit });
}
let parsed = parse_8digits(chunk) as i64;
result = result.wrapping_mul(100_000_000);
if is_positive {
result = result.wrapping_add(parsed);
} else {
result = result.wrapping_sub(parsed);
}
digits = rest;
remaining -= 8;
}

Ok((result, digits))
}

#[cfg(target_pointer_width = "32")]
#[inline]
pub(super) const fn parse_decimal_i64(
is_positive: bool,
mut result: i64,
mut digits: &[u8],
) -> Result<(i64, &[u8]), ParseIntError> {
// 32-bit platforms avoid 64-bit multiplication, so use 4-digit batches.
let mut remaining = 16;

while remaining >= 4 {
let [a, b, c, d, rest @ ..] = digits else { break };
let chunk = u32::from_le_bytes([*a, *b, *c, *d]);
if !is_digits(chunk as usize) {
return Err(ParseIntError { kind: IntErrorKind::InvalidDigit });
}
let parsed = parse_4digits(chunk) as i64;
result = result.wrapping_mul(10_000);
if is_positive {
result = result.wrapping_add(parsed);
} else {
result = result.wrapping_sub(parsed);
}
digits = rest;
remaining -= 4;
}

Ok((result, digits))
}
}

#[cfg_attr(not(panic = "immediate-abort"), inline(never))]
#[cfg_attr(panic = "immediate-abort", inline)]
#[cold]
Expand All @@ -1601,9 +1755,37 @@ const fn from_ascii_bytes_radix_panic(radix: u32) -> ! {
)
}

macro_rules! from_str_int_impl {
($signedness:ident $($int_ty:ty)+) => {$(
#[stable(feature = "rust1", since = "1.0.0")]
macro_rules! define_swar_min_len {
($int_ty:ty, true) => {
const SWAR_MIN_LEN: usize = 16;
};
($int_ty:ty, false) => {};
}

macro_rules! maybe_define_swar {
($int_ty:ty, true) => {
/// Parses up to 16 leading radix-10 digits in batches.
///
/// This thin wrapper dispatches to `decimal_swar`, which lives outside
/// the per-type macro so its body is not duplicated for types that
/// never reach it.
#[inline]
pub(super) const fn swar_parse_decimal(
is_positive: bool,
result: $int_ty,
digits: &[u8],
) -> Result<($int_ty, &[u8]), ParseIntError> {
match decimal_swar::parse_decimal_i64(is_positive, result as i64, digits) {
Ok((r, rest)) => Ok((r as $int_ty, rest)),
Err(e) => Err(e),
}
}
};
($int_ty:ty, false) => {};
}

macro_rules! from_str_int_impl_inner {
($signedness:ident $int_ty:ty, $swar:tt) => { #[stable(feature = "rust1", since = "1.0.0")]
#[rustc_const_unstable(feature = "const_convert", issue = "143773")]
const impl FromStr for $int_ty {
type Err = ParseIntError;
Expand Down Expand Up @@ -1784,7 +1966,7 @@ macro_rules! from_str_int_impl {
<$int_ty>::from_ascii_bytes_radix_impl(src.as_ref(), radix)
}

#[inline]
#[inline(always)]
pub(super) const fn from_ascii_bytes_radix_impl(src: &[u8], radix: u32) -> Result<$int_ty, ParseIntError> {
use self::IntErrorKind::*;
use self::ParseIntError as PIE;
Expand Down Expand Up @@ -1820,16 +2002,65 @@ macro_rules! from_str_int_impl {
};
}

// SWAR fast path: process leading decimal digits in
// batches using SIMD-within-a-register. Only radix 10 and
// only for types of at least 8 bytes; smaller types use the
// no-SWAR macro arm and never see this code.
//
// We never batch more than 16 leading digits: that many
// decimal digits can never overflow `u64`/`u128`, so the
// batch arithmetic is safe to run unchecked even in debug
// builds. The tail after the batches is validated by the
// checked loop, which still catches overflow for inputs
// longer than the type's range.
define_swar_min_len!($int_ty, $swar);

macro_rules! run_swar {
(true) => {{
if radix == 10 && digits.len() >= SWAR_MIN_LEN {
match Self::swar_parse_decimal(is_positive, result, digits) {
Ok((r, rest)) => {
result = r;
digits = rest;
}
Err(e) => return Err(e),
}
}
}};
(false) => {{}};
}

macro_rules! run_checked_loop {
($checked_additive_op:ident, $overflow_err:ident) => {{
while let [c, rest @ ..] = digits {
// When `radix` is passed in as a literal, rather than doing a slow `imul`
// the compiler can use shifts if `radix` can be expressed as a
// sum of powers of 2 (x*10 can be written as x*8 + x*2).
// When the compiler can't use these optimisations,
// the latency of the multiplication can be hidden by issuing it
// before the result is needed to improve performance on
// modern out-of-order CPU as multiplication here is slower
// than the other instructions, we can get the end result faster
// doing multiplication first and let the CPU spends other cycles
// doing other computation and get multiplication result later.
let mul = result.checked_mul(radix as $int_ty);
let x = unwrap_or_PIE!((*c as char).to_digit(radix), InvalidDigit) as $int_ty;
result = unwrap_or_PIE!(mul, $overflow_err);
result = unwrap_or_PIE!(<$int_ty>::$checked_additive_op(result, x), $overflow_err);
digits = rest;
}
}};
}

// If the len of the str is short compared to the range of the type
// we are parsing into, then we can be certain that an overflow will not occur.
// This bound is when `radix.pow(digits.len()) - 1 <= T::MAX` but the condition
// above is a faster (conservative) approximation of this.
//
// Consider radix 16 as it has the highest information density per digit and will thus overflow the earliest:
// `u8::MAX` is `ff` - any str of len 2 is guaranteed to not overflow.
// `i8::MAX` is `7f` - only a str of len 1 is guaranteed to not overflow.
if can_not_overflow::<$int_ty>(radix, is_signed_ty, digits) {
// If the len of the str is short compared to the range of the type
// we are parsing into, then we can be certain that an overflow will not occur.
// This bound is when `radix.pow(digits.len()) - 1 <= T::MAX` but the condition
// above is a faster (conservative) approximation of this.
//
// Consider radix 16 as it has the highest information density per digit and will thus overflow the earliest:
// `u8::MAX` is `ff` - any str of len 2 is guaranteed to not overflow.
// `i8::MAX` is `7f` - only a str of len 1 is guaranteed to not overflow.
//
// NOTE: We could use unchecked arithmetic here, but we don't, based on the observation
// that it produces the same assembly as wrapping ones. See #163099.
macro_rules! run_no_check_loop {
Expand All @@ -1848,27 +2079,7 @@ macro_rules! from_str_int_impl {
run_no_check_loop!(wrapping_sub)
};
} else {
macro_rules! run_checked_loop {
($checked_additive_op:ident, $overflow_err:ident) => {{
while let [c, rest @ ..] = digits {
// When `radix` is passed in as a literal, rather than doing a slow `imul`
// the compiler can use shifts if `radix` can be expressed as a
// sum of powers of 2 (x*10 can be written as x*8 + x*2).
// When the compiler can't use these optimisations,
// the latency of the multiplication can be hidden by issuing it
// before the result is needed to improve performance on
// modern out-of-order CPU as multiplication here is slower
// than the other instructions, we can get the end result faster
// doing multiplication first and let the CPU spends other cycles
// doing other computation and get multiplication result later.
let mul = result.checked_mul(radix as $int_ty);
let x = unwrap_or_PIE!((*c as char).to_digit(radix), InvalidDigit) as $int_ty;
result = unwrap_or_PIE!(mul, $overflow_err);
result = unwrap_or_PIE!(<$int_ty>::$checked_additive_op(result, x), $overflow_err);
digits = rest;
}
}};
}
run_swar!($swar);
if is_positive {
run_checked_loop!(checked_add, PosOverflow)
} else {
Expand All @@ -1877,9 +2088,22 @@ macro_rules! from_str_int_impl {
}
Ok(result)
}

maybe_define_swar!($int_ty, $swar);
}
)*}
};
}

macro_rules! from_str_int_impl {
($signedness:ident $($int_ty:ty)+) => {
$( from_str_int_impl_inner! { $signedness $int_ty, true } )+
};
($signedness:ident @no_swar $($int_ty:ty)+) => {
$( from_str_int_impl_inner! { $signedness $int_ty, false } )+
};
}

from_str_int_impl! { signed isize i8 i16 i32 i64 i128 }
from_str_int_impl! { unsigned usize u8 u16 u32 u64 u128 }
from_str_int_impl! { signed i64 i128 }
from_str_int_impl! { signed @no_swar isize i8 i16 i32 }
from_str_int_impl! { unsigned u64 u128 }
from_str_int_impl! { unsigned @no_swar usize u8 u16 u32 }
36 changes: 36 additions & 0 deletions library/coretests/benches/num/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,21 @@ const ASCII_NUMBERS: [&str; 19] = [
"c0ffee",
];

/// Long decimal strings (16-20 digits) that trigger the SWAR fast path
/// multiple times for 64-bit integer parsing.
const LONG_ASCII_NUMBERS: [&str; 10] = [
"1234567890123456", // 16 digits, exactly 2 SWAR chunks
"12345678901234567", // 17 digits
"123456789012345678", // 18 digits
"1234567890123456789", // 19 digits
"18446744073709551615", // 20 digits, u64::MAX
"9223372036854775807", // 19 digits, i64::MAX
"9999999999999999", // 16 digits
"10000000000000000", // 17 digits
"-9223372036854775808", // 19 digits + sign, i64::MIN
"0000123456789012", // 16 digits with leading zeros
];

macro_rules! from_str_bench {
($mac:ident, $t:ty) => {
#[bench]
Expand Down Expand Up @@ -63,6 +78,22 @@ macro_rules! from_str_radix_bench {
};
}

macro_rules! from_str_radix_long_bench {
($mac:ident, $t:ty, $radix:expr) => {
#[bench]
fn $mac(b: &mut Bencher) {
b.iter(|| {
LONG_ASCII_NUMBERS
.iter()
.cycle()
.take(5_000)
.filter_map(|s| <$t>::from_str_radix(black_box(s), $radix).ok())
.max()
})
}
};
}

from_str_bench!(bench_u8_from_str, u8);
from_str_radix_bench!(bench_u8_from_str_radix_2, u8, 2);
from_str_radix_bench!(bench_u8_from_str_radix_10, u8, 10);
Expand Down Expand Up @@ -110,3 +141,8 @@ from_str_radix_bench!(bench_i64_from_str_radix_2, i64, 2);
from_str_radix_bench!(bench_i64_from_str_radix_10, i64, 10);
from_str_radix_bench!(bench_i64_from_str_radix_16, i64, 16);
from_str_radix_bench!(bench_i64_from_str_radix_36, i64, 36);

// Long-string benchmarks: 16-20 digit decimal numbers that exercise
// the SWAR fast path (2+ iterations of 8-digit-at-a-time parsing).
from_str_radix_long_bench!(bench_u64_from_str_radix_10_long, u64, 10);
from_str_radix_long_bench!(bench_i64_from_str_radix_10_long, i64, 10);
Loading
Loading