diff --git a/library/core/src/num/mod.rs b/library/core/src/num/mod.rs index 25326f4f066c6..7c8b1c121a666 100644 --- a/library/core/src/num/mod.rs +++ b/library/core/src/num/mod.rs @@ -1589,6 +1589,160 @@ pub const fn can_not_overflow(radix: u32, is_signed_ty: bool, digits: &[u8]) radix <= 16 && digits.len() <= size_of::() * 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] @@ -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; @@ -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; @@ -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 { @@ -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 { @@ -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 } diff --git a/library/coretests/benches/num/mod.rs b/library/coretests/benches/num/mod.rs index a131b3454f0cc..07701ccd3fca5 100644 --- a/library/coretests/benches/num/mod.rs +++ b/library/coretests/benches/num/mod.rs @@ -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] @@ -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); @@ -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); diff --git a/library/coretests/tests/num/mod.rs b/library/coretests/tests/num/mod.rs index b1c3001790f07..ae6a5ae19f280 100644 --- a/library/coretests/tests/num/mod.rs +++ b/library/coretests/tests/num/mod.rs @@ -152,6 +152,36 @@ fn test_int_from_str_overflow() { test_parse::("-9223372036854775809", Err(IntErrorKind::NegOverflow)); } +#[test] +fn test_from_str_radix_10_swar_boundaries() { + // SWAR batches up to 16 leading digits for 64-bit types. These + // cases sit on the batch boundaries and put invalid digits inside + // the batched window. + test_parse::("9999999999999999", Ok(9_999_999_999_999_999)); + test_parse::("10000000000000000", Ok(10_000_000_000_000_000)); + test_parse::("18446744073709551615", Ok(18_446_744_073_709_551_615)); + test_parse::("18446744073709551616", Err(IntErrorKind::PosOverflow)); + test_parse::("99999999999999999999", Err(IntErrorKind::PosOverflow)); + + test_parse::("1234567890123456x", Err(IntErrorKind::InvalidDigit)); + test_parse::("x2345678901234567", Err(IntErrorKind::InvalidDigit)); + test_parse::("1234567890x2345678", Err(IntErrorKind::InvalidDigit)); + test_parse::("12345678901234567x", Err(IntErrorKind::InvalidDigit)); + + test_parse::("9223372036854775807", Ok(9_223_372_036_854_775_807)); + test_parse::("-9223372036854775808", Ok(-9_223_372_036_854_775_808)); + test_parse::("9223372036854775808", Err(IntErrorKind::PosOverflow)); + test_parse::("-9223372036854775809", Err(IntErrorKind::NegOverflow)); + test_parse::("-123456789012345678x", Err(IntErrorKind::InvalidDigit)); + test_parse::("123456789012345x6", Err(IntErrorKind::InvalidDigit)); + + // Short inputs must not be affected by the SWAR path. + test_parse::("0", Ok(0)); + test_parse::("+42", Ok(42)); + test_parse::("-42", Ok(-42)); + test_parse::("12a34", Err(IntErrorKind::InvalidDigit)); +} + #[test] fn test_can_not_overflow() { fn can_overflow(radix: u32, input: &str) -> bool diff --git a/tests/ui/consts/const-eval/parse_ints.stderr b/tests/ui/consts/const-eval/parse_ints.stderr index b8027fd951d5a..4cb7f81f044ed 100644 --- a/tests/ui/consts/const-eval/parse_ints.stderr +++ b/tests/ui/consts/const-eval/parse_ints.stderr @@ -14,7 +14,7 @@ note: inside `core::num::::from_ascii_bytes_radix_impl` ::: $SRC_DIR/core/src/num/mod.rs:LL:COL | = note: in this macro invocation - = note: this error originates in the macro `from_str_int_impl` (in Nightly builds, run with -Z macro-backtrace for more info) + = note: this error originates in the macro `from_str_int_impl_inner` which comes from the expansion of the macro `from_str_int_impl` (in Nightly builds, run with -Z macro-backtrace for more info) error[E0080]: evaluation panicked: from_ascii_bytes_radix: radix must lie in the range `[2, 36]` --> $DIR/parse_ints.rs:8:25 @@ -32,7 +32,7 @@ note: inside `core::num::::from_ascii_bytes_radix_impl` ::: $SRC_DIR/core/src/num/mod.rs:LL:COL | = note: in this macro invocation - = note: this error originates in the macro `from_str_int_impl` (in Nightly builds, run with -Z macro-backtrace for more info) + = note: this error originates in the macro `from_str_int_impl_inner` which comes from the expansion of the macro `from_str_int_impl` (in Nightly builds, run with -Z macro-backtrace for more info) error: aborting due to 2 previous errors