diff --git a/.github/workflows/tests.yml b/.github/workflows/tests.yml index a40ff963..1abdb09f 100644 --- a/.github/workflows/tests.yml +++ b/.github/workflows/tests.yml @@ -162,4 +162,9 @@ jobs: with: toolchain: stable components: clippy - - run: cargo clippy --all-features --all-targets --workspace --exclude dashu-python -- -D warnings \ No newline at end of file + - name: Clippy (default / 64-bit Word) + run: cargo clippy --all-features --all-targets --workspace --exclude dashu-python -- -D warnings + - name: Clippy (32-bit Word) + env: + RUSTFLAGS: --cfg force_bits="32" + run: cargo clippy --all-features --all-targets --workspace --exclude dashu-python -- -D warnings \ No newline at end of file diff --git a/Cargo.toml b/Cargo.toml index 2f7761ca..0a6d0ba5 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -36,6 +36,7 @@ std = ["dashu-base/std", "dashu-int/std", "dashu-float/std", "dashu-ratio/std"] # stable features serde = ["dashu-int/serde", "dashu-float/serde", "dashu-ratio/serde"] num-order = ["dashu-int/num-order", "dashu-float/num-order", "dashu-ratio/num-order"] +tuning = ["dashu-int/tuning"] zeroize = ["dashu-int/zeroize", "dashu-float/zeroize", "dashu-ratio/zeroize"] # unstable features diff --git a/TODO.md b/TODO.md deleted file mode 100644 index 9d459cf1..00000000 --- a/TODO.md +++ /dev/null @@ -1,69 +0,0 @@ -## dashu-int Improvements - -### High impact - -- **`submul_1` fused primitive** — Multiply-and-subtract in one pass for the division inner loop correction step. Currently done as separate mul + sub, doubling memory passes. Reference: `ramp/src/ll/mul.rs:134-182`. - -- **`div_preinv` / 3-by-2 division with pre-inverted divisor** — Ramp computes `invert_pi(d1, d0)` for fast approximate quotient without the x86 `div` instruction, with separate `divrem_1` (single-limb) and `divrem_2` (two-limb) fast paths. Could replace the `num-modular` dependency. Affects many hot paths: modular arithmetic, formatting, base conversion. Reference: `ramp/src/ll/limb.rs:768-783`, `ramp/src/ll/div.rs:208-253`. - -- **Dedicated Toom-2 squaring** — `sqr_toom2` exploits symmetry: uses `z1 = x0*x1` instead of `(x0-x0)*(y1-y0)`, eliminating subtraction operations — only 3 sub-products instead of 4, and the cross term is `2*z1` without signed arithmetic. dashu has Karatsuba and Toom-3 general multiplication but no squaring-specific variant that takes advantage of `x == y`. dashu only uses specialized squaring up to 30 words. Reference: `ramp/src/ll/mul.rs:473-512`. - -- **Toom-22 as intermediate multiplication** — Ramp uses Toom-22 above 20 limbs before falling back to unbalanced mul. dashu goes Karatsuba → Toom-3 at 192 words. Toom-22 could fill the 24–192 word gap. Reference: `ramp/src/ll/mul.rs:243-390`. - -### Medium impact - -- **Trailing zero stripping in GCD loop** — Strip trailing zeros after each subtraction, not just at initialization. Helps for random inputs where intermediate results often gain trailing zeros. Reference: `ramp/src/ll/gcd.rs:20-86`. - -- **Trailing zero stripping in pow** — Factor out `(m * 2^k)^exp = m^exp * 2^(k*exp)` to reduce operand size. Reference: `ramp/src/ll/pow.rs:41-118`. - -### Low impact / ergonomics - -- **Build-time BASES table** — Pre-compute `digits_per_limb` and `big_base` per base via `build.rs` so base-10 conversion avoids repeated division. dashu's `integer/src/fmt/non_power_two.rs` uses simpler chunking (`CHUNK_LEN = 16`) without precomputed powers. Reference: `ramp/src/ll/base.rs:31-40`, `ramp/build.rs`. - -- **Scratch allocator improvements** — Ramp's `TmpAllocator` uses a linked list of dynamic allocations freed on drop, vs. dashu's pre-computed layout approach. Might be simpler for algorithms with hard-to-predict memory needs. Reference: `ramp/src/mem.rs`. - -## dashu-ratio Improvements - -- GCD: An idea of fast gcd check for rational number: don't do gcd reduction after every operation. - For small numerators or denominators, we can directly do a gcd, otherwise, we first do gcd with a primorial that - fits in a word (min is u16), and only remove these small divisors. - Further improvement: store a const divisor for the prime factors in the primorial, thus supports a fast factorial of - the gcd result, and then divide with these const divisor. - -## dashu-float Improvements - -- **Trig (`sin`/`cos`) — current baseline** — `float/src/math/trig.rs` uses dynamic guard-digit - work precision, simple `x mod (π/2)` range reduction, and Taylor series on the reduced - argument `r`. Adequate for moderate precision and moderate `|x|`; items below target large - arguments and very high precision. Reference: MPFR `mpfr_sin` / `mpfr_sin_cos`. - -### High impact - -- **Payne–Hanek range reduction** — For large `|x|`, replace `k = round(x/(π/2)); r = x - k·(π/2)` with - multiplication by precomputed blocks of `2/π`, extracting the integer part without a full high-precision - division. Avoids catastrophic cancellation that currently forces `work_precision ≈ precision + log|x| + guards` - (`compute_work_context`). This is the main gap vs. MPFR for huge arguments. Reference: `float/src/math/trig.rs`. - -- **Binary splitting for Taylor core** — At high precision, evaluate the `sin`/`cos` series via binary splitting - (same technique as Chudnovsky π in `float/src/math/consts.rs`) instead of naive term-by-term accumulation. - Reduces cost from O(p²) to roughly O(M(p) log p) for p-bit results. - -### Medium impact - -- **Remez minimax polynomial + Clenshaw (low/medium p)** — For `p ≲ 512`, use a fixed-degree minimax polynomial - on `[-π/4, π/4]` evaluated with Clenshaw recurrence instead of Taylor. MPFR uses this for its fast path; switch - to series/binary splitting only when p is large. - -- **Cody–Waite π/2 split** — Represent `π/2 = hi + lo` and compute `r = ((x - k·hi) - k·lo)` to reduce guard - digit pressure for moderate `|x|` before Payne–Hanek is needed. Complements the existing `reduce_to_quadrant`. - -- **Argument shrinking for `|r| > π/4`** — Use `sin(r) = cos(π/2 - r)` (and the cosine analogue) so the Taylor - series runs on a smaller interval, needing fewer terms when `r` is near ±π/2. - -### Low impact / ergonomics - -- **Cache π at common precisions** — Avoid recomputing Chudnovsky π on every trig call when work precision repeats. - TODO already noted in `float/src/math/consts.rs`. - -- **Precomputed `2/π` block table** — Storage for Payne–Hanek: blocks of `2/π` bits (e.g. 32/64 bits per entry), - generated once or lazily on first use at a given precision. diff --git a/float/src/parse.rs b/float/src/parse.rs index 1b164cee..3bf106c5 100644 --- a/float/src/parse.rs +++ b/float/src/parse.rs @@ -27,6 +27,12 @@ impl Repr { pub fn from_str_native(mut src: &str) -> Result<(Self, usize), ParseError> { assert!(MIN_RADIX as Word <= B && B <= MAX_RADIX as Word); + // B is guaranteed to be in 2..=36 by the assert above; the cast to u32 + // is needed because `from_str_radix` takes a u32 radix. On 32-bit Word + // targets the cast is a no-op. + #[allow(clippy::unnecessary_cast)] + let radix: u32 = B as u32; + // parse and remove the sign let sign = match src.strip_prefix('-') { Some(s) => { @@ -100,14 +106,14 @@ impl Repr { return Err(ParseError::UnsupportedRadix); } else { let digits = int_str.len() - int_str.matches('_').count(); - (UBig::from_str_radix(&src[..dot], B as u32)?, digits, B as u32) + (UBig::from_str_radix(&src[..dot], radix)?, digits, radix) } } else { if pmarker { // prefix is required for using `p` as scale marker return Err(ParseError::UnsupportedRadix); } - (UBig::ZERO, 0, B as u32) + (UBig::ZERO, 0, radix) }; // parse fractional part @@ -139,7 +145,7 @@ impl Repr { return Err(ParseError::UnsupportedRadix); } else { ndigits = src.len() - src.matches('_').count(); - UBig::from_str_radix(src, B as u32)? + UBig::from_str_radix(src, radix)? } }; diff --git a/float/src/third_party/num_traits.rs b/float/src/third_party/num_traits.rs index 1d5e3dbe..0c93e7ef 100644 --- a/float/src/third_party/num_traits.rs +++ b/float/src/third_party/num_traits.rs @@ -133,7 +133,7 @@ impl num_traits::Num for FBig { #[inline] fn from_str_radix(s: &str, radix: u32) -> Result { // the conversion might a fail with 16-bit words. - #[allow(clippy::unnecessary_fallible_conversions)] + #[allow(clippy::unnecessary_fallible_conversions, clippy::useless_conversion)] let r: Word = radix.try_into().map_err(|_| ParseError::UnsupportedRadix)?; if r == B { #[allow(deprecated)] // TODO(v0.5): remove after from_str_native is made private. diff --git a/integer/CHANGELOG.md b/integer/CHANGELOG.md index 71fb36ab..17e4b422 100644 --- a/integer/CHANGELOG.md +++ b/integer/CHANGELOG.md @@ -3,15 +3,29 @@ ## Unreleased ### Add +- NTT-based multiplication using Proth primes (`K·2^N + 1`), combined via Garner CRT. Supports 64-bit and 32-bit Word targets. Threshold at 4 000 words (~256 kbits). +- Asymmetric NTT chunking: when one operand is much larger than the other, the shorter operand is forward-transformed once and reused across chunks. - `UBig::from_u64` and `IBig::from_i64`, const on 32-bit and 64-bit targets. ### Improve - Basecase (schoolbook) multiplication now uses an dword mult inner kernel (two multiplier words per sweep over the accumulator, mirroring GMP's `mpn_addmul_2` and `mpn_submul_2`), roughly halving accumulator memory traffic. - Addition and subtraction carry/borrow propagation now uses `Word` (u64/u32) instead of `bool` throughout the architecture-specific `add_with_carry` and `sub_with_borrow` functions, eliminating `bool`↔Word conversions in the inner loops. +- Lowered the Karatsuba→Toom-3 multiplication threshold from 192 to 96 words, giving Toom-Cook-3 at ~6000 bits instead of ~12000 bits — closes the gap with malachite at ~10000-bit sizes. +- NTT coefficient width increased from 16 to 64 bits (K_eff=3 for 64-bit, K_eff=2 otherwise), roughly halving the transform length at each step. +- NTT multiplication auto-selects `K_eff = 2` primes when headroom allows, skipping the third prime. +- Multiplication thresholds can be overridden at runtime via `DASHU_THRESHOLD_SIMPLE`, `DASHU_THRESHOLD_KARATSUBA`, and `DASHU_THRESHOLD_NTT` environment variables (requires `tuning` feature). -### Improve -- Logarithm for very large values uses power-sequence decomposition, replacing iterative single-step multiplication. -- Improve power-of-two base formatting ([#3](https://github.com/cmpute/dashu/pull/3)) +### Change +- NTT multiplication now uses Proth primes (`K·2^N + 1`) instead of Solinas primes, improving modular reduction speed. +- NTT threshold lowered from 40 000 to 4 000 words. +- NTT enabled for 32-bit Word targets. +- Arch-specific NTT prime definitions under `arch/generic_{32,64}_bit/ntt.rs`. + +### Fix +- `pack.rs` test used 64-bit literals that overflowed `Word` (`u32`) on 32-bit targets, breaking the test build. +- `pack.rs` now uses native `Word`/`Lane` types throughout instead of `u64`/`u32`, fixing clippy `unnecessary_cast` warnings on 64-bit. +- `test_unpack_carry_propagation` had a hardcoded 64-bit shift assumption; now derived from `Word::BITS` so it works on 32-bit. +- Various clippy warnings (`let_and_return`, `too_many_arguments`, `needless_range_loop`, `type_complexity`) resolved across the NTT module. ## 0.4.2 diff --git a/integer/Cargo.toml b/integer/Cargo.toml index 047e701e..9f8fb6ff 100644 --- a/integer/Cargo.toml +++ b/integer/Cargo.toml @@ -19,6 +19,7 @@ all-features = true [features] default = ["std", "num-order"] std = ["dashu-base/std"] +tuning = ["std"] # unstable dependencies rand = ["rand_v08"] @@ -30,7 +31,7 @@ dashu-base = { version = "0.4.1", default-features = false, path = "../base" } cfg-if = { version = "1.0.0" } static_assertions = { version = "1.1" } rustversion = { version = "1.0.0" } -num-modular = { version = "0.6.1" } +num-modular = { version = "0.6.4" } # stable dependencies num-order = { optional = true, version = "1.2.0", default-features = false } diff --git a/integer/benches/primitive.rs b/integer/benches/primitive.rs index 2958fc54..6bcee837 100644 --- a/integer/benches/primitive.rs +++ b/integer/benches/primitive.rs @@ -150,6 +150,31 @@ fn ubig_ilog_large(criterion: &mut Criterion) { group.finish(); } +fn ubig_mul_asymmetric(criterion: &mut Criterion) { + let mut rng = StdRng::seed_from_u64(SEED); + let mut group = criterion.benchmark_group("ubig_mul_asymmetric"); + group.plot_config(PlotConfiguration::default().summary_scale(AxisScale::Logarithmic)); + + // b just above the NTT threshold (4 000 words = 256 kbits → use 500 kbits). + let b_bits = 500_000; + let b = random_ubig(b_bits, &mut rng); + + // a ranges from 1 kbit (below Karatsuba threshold) to heavily + // asymmetric (10×), exercising all chunked-mul code paths. + for &a_bits in &[ + 1_000, 10_000, 100_000, 500_000, 1_000_000, 2_000_000, 5_000_000, + ] { + let a = random_ubig(a_bits, &mut rng); + group.bench_with_input( + BenchmarkId::from_parameter(format!("{a_bits}/{b_bits}")), + &(a, &b), + |bencher, (ta, tb)| bencher.iter(|| ta * *tb), + ); + } + + group.finish(); +} + criterion_group!( benches, ubig_add, @@ -163,6 +188,7 @@ criterion_group!( ubig_modulo_pow, ubig_pow_large_base, ubig_ilog_large, + ubig_mul_asymmetric, ); criterion_main!(benches); diff --git a/integer/src/arch/generic_32_bit/mod.rs b/integer/src/arch/generic_32_bit/mod.rs index eaab456f..fd98d73d 100644 --- a/integer/src/arch/generic_32_bit/mod.rs +++ b/integer/src/arch/generic_32_bit/mod.rs @@ -4,4 +4,5 @@ pub(crate) mod add; #[path = "../generic/digits.rs"] pub(crate) mod digits; +pub(crate) mod ntt; pub(crate) mod word; diff --git a/integer/src/arch/generic_32_bit/ntt.rs b/integer/src/arch/generic_32_bit/ntt.rs new file mode 100644 index 00000000..746bc188 --- /dev/null +++ b/integer/src/arch/generic_32_bit/ntt.rs @@ -0,0 +1,84 @@ +//! NTT primes and constants for 32-bit Word targets. +//! +//! Uses Proth primes of the form `K * 2^N + 1`. +//! All constants computed by `integer/src/mul/ntt/compute_constants.py`. + +use num_modular::FixedProth32; + +// Proth reducer instances — each with a different (N, K) pair. +pub type Rp0 = FixedProth32<26, 7>; +pub type Rp1 = FixedProth32<27, 15>; +pub type Rp2 = FixedProth32<27, 17>; + +pub const P0: Rp0 = FixedProth32::<26, 7>; +pub const P1: Rp1 = FixedProth32::<27, 15>; +pub const P2: Rp2 = FixedProth32::<27, 17>; + +pub const K: usize = 3; +pub const MAX_LOG_N: u32 = 26; +pub const B_PACK_MIN: u32 = 8; +pub const B_PACK_CANDIDATES: &[u32] = &[32, 16, 8]; + +pub type Lane = u32; + +/// Primitive `MAX_LOG_N`-th roots of unity for each prime. +pub const OMEGA_MAX: [Lane; K] = [ + 0x0000088b, // P0 + 0x3a26eef8, // P1 + 0x1aa0ab5e, // P2 +]; + +pub const CRT_INV_IJ: [[Lane; K]; K] = [[0, 0x4e42c85b, 0x5fb425ef], [0, 0, 0x44000009], [0, 0, 0]]; + +/// Prime moduli indexed by PI. +pub const MODULI: [Lane; K] = [Rp0::MODULUS, Rp1::MODULUS, Rp2::MODULUS]; + +#[cfg(test)] +mod tests { + use super::*; + use num_modular::Reducer; + + type ReducerFns = (fn(Lane) -> Lane, fn(Lane) -> Lane, fn(Lane) -> Lane); + + #[test] + fn test_primes_proth_form() { + assert_eq!(MODULI[0], 7u32 * (1u32 << 26) + 1); + assert_eq!(MODULI[1], 15u32 * (1u32 << 27) + 1); + assert_eq!(MODULI[2], 17u32 * (1u32 << 27) + 1); + } + + #[test] + fn test_primes_v2() { + for &p in &MODULI { + let v2 = (p - 1).trailing_zeros(); + assert!(v2 >= MAX_LOG_N, "v2(p-1) = {v2} < MAX_LOG_N"); + } + } + + #[test] + fn test_omega_order() { + for (pi, &omega_max) in OMEGA_MAX.iter().enumerate() { + let p = MODULI[pi]; + let (sqr, to_m, from_m): ReducerFns = match pi { + 0 => { + (|w| P0.reduce((w as u64) * (w as u64)), |v| P0.transform(v), |v| P0.residue(v)) + } + 1 => { + (|w| P1.reduce((w as u64) * (w as u64)), |v| P1.transform(v), |v| P1.residue(v)) + } + 2 => { + (|w| P2.reduce((w as u64) * (w as u64)), |v| P2.transform(v), |v| P2.residue(v)) + } + _ => unreachable!(), + }; + + let mut w = to_m(omega_max); + for _ in 0..MAX_LOG_N - 1 { + w = sqr(w); + } + assert_eq!(from_m(w), p - 1, "omega^(2^(MAX_LOG_N-1)) != -1 mod p for prime {pi}"); + w = sqr(w); + assert_eq!(from_m(w), 1, "omega^(2^MAX_LOG_N) != 1 mod p for prime {pi}"); + } + } +} diff --git a/integer/src/arch/generic_32_bit/word.rs b/integer/src/arch/generic_32_bit/word.rs index 4b3769d3..2cb44b61 100644 --- a/integer/src/arch/generic_32_bit/word.rs +++ b/integer/src/arch/generic_32_bit/word.rs @@ -9,3 +9,7 @@ pub type DoubleWord = u64; /// Signed double machine word. pub type SignedDoubleWord = i64; + +/// Accumulator for the product of three primes (3 × 2^32 ≈ 2^96). +#[derive(Clone, Copy, Debug, Default)] +pub struct TripleWord(pub [u32; 3]); diff --git a/integer/src/arch/generic_64_bit/mod.rs b/integer/src/arch/generic_64_bit/mod.rs index eaab456f..fd98d73d 100644 --- a/integer/src/arch/generic_64_bit/mod.rs +++ b/integer/src/arch/generic_64_bit/mod.rs @@ -4,4 +4,5 @@ pub(crate) mod add; #[path = "../generic/digits.rs"] pub(crate) mod digits; +pub(crate) mod ntt; pub(crate) mod word; diff --git a/integer/src/arch/generic_64_bit/ntt.rs b/integer/src/arch/generic_64_bit/ntt.rs new file mode 100644 index 00000000..03896b5b --- /dev/null +++ b/integer/src/arch/generic_64_bit/ntt.rs @@ -0,0 +1,95 @@ +//! NTT primes and constants for 64-bit Word targets. +//! +//! Uses Proth primes of the form `K * 2^N + 1`. +//! All constants computed by `integer/src/mul/ntt/compute_constants.py`. + +use num_modular::FixedProth64; + +// Proth reducer instances — each with a different (N, K) pair. +pub type Rp0 = FixedProth64<57, 29>; +pub type Rp1 = FixedProth64<57, 71>; +pub type Rp2 = FixedProth64<57, 75>; + +pub const P0: Rp0 = FixedProth64::<57, 29>; +pub const P1: Rp1 = FixedProth64::<57, 71>; +pub const P2: Rp2 = FixedProth64::<57, 75>; + +pub const K: usize = 3; +pub const MAX_LOG_N: u32 = 57; +pub const B_PACK_MIN: u32 = 16; +pub const B_PACK_CANDIDATES: &[u32] = &[64, 32, 16]; + +pub type Lane = u64; + +/// Primitive `MAX_LOG_N`-th roots of unity for each prime: +/// `omega_max[i]` = `g^{(p_i-1) / 2^MAX_LOG_N} mod p_i`. +pub const OMEGA_MAX: [Lane; K] = [ + 0x00003e6b41437d93, // P0 + 0x2f754195e85edc63, // P1 + 0x75544cac36cebb29, // P2 +]; + +pub const CRT_INV_IJ: [[Lane; K]; K] = [ + [0, 0x3979e79e79e79e7c, 0x8c37a6f4de9bd37d], + [0, 0, 0x2580000000000013], + [0, 0, 0], +]; + +/// Prime moduli indexed by PI. +pub const MODULI: [Lane; K] = [Rp0::MODULUS, Rp1::MODULUS, Rp2::MODULUS]; + +#[cfg(test)] +mod tests { + use super::*; + use num_modular::Reducer; + + type ReducerFns = (fn(Lane) -> Lane, fn(Lane) -> Lane, fn(Lane) -> Lane); + + #[test] + fn test_primes_proth_form() { + assert_eq!(MODULI[0], 29u64 * (1u64 << 57) + 1); + assert_eq!(MODULI[1], 71u64 * (1u64 << 57) + 1); + assert_eq!(MODULI[2], 75u64 * (1u64 << 57) + 1); + } + + #[test] + fn test_primes_v2() { + for &p in &MODULI { + let v2 = (p - 1).trailing_zeros(); + assert!(v2 >= MAX_LOG_N, "v2(p-1) = {v2} < MAX_LOG_N"); + } + } + + #[test] + fn test_omega_order() { + for (pi, &omega_max) in OMEGA_MAX.iter().enumerate() { + let p = MODULI[pi]; + let (sqr, to_m, from_m): ReducerFns = match pi { + 0 => ( + |w| P0.reduce((w as u128) * (w as u128)), + |v| P0.transform(v), + |v| P0.residue(v), + ), + 1 => ( + |w| P1.reduce((w as u128) * (w as u128)), + |v| P1.transform(v), + |v| P1.residue(v), + ), + 2 => ( + |w| P2.reduce((w as u128) * (w as u128)), + |v| P2.transform(v), + |v| P2.residue(v), + ), + _ => unreachable!(), + }; + + let mut w = to_m(omega_max); + for _ in 0..MAX_LOG_N - 1 { + w = sqr(w); + } + assert_eq!(from_m(w), p - 1, "omega^(2^(MAX_LOG_N-1)) != -1 mod p for prime {pi}"); + w = sqr(w); + assert_eq!(from_m(w), 1, "omega^(2^MAX_LOG_N) != 1 mod p for prime {pi}"); + } + } +} diff --git a/integer/src/arch/generic_64_bit/word.rs b/integer/src/arch/generic_64_bit/word.rs index fefd7e7f..a5f15ed9 100644 --- a/integer/src/arch/generic_64_bit/word.rs +++ b/integer/src/arch/generic_64_bit/word.rs @@ -9,3 +9,7 @@ pub type DoubleWord = u128; /// Signed double machine word. pub type SignedDoubleWord = i128; + +/// Accumulator for the product of three primes (3 × 2^64 = 2^192). +#[derive(Clone, Copy, Debug, Default)] +pub struct TripleWord(pub [u64; 3]); diff --git a/integer/src/arch/mod.rs b/integer/src/arch/mod.rs index 3bcd88fa..dd11b686 100644 --- a/integer/src/arch/mod.rs +++ b/integer/src/arch/mod.rs @@ -6,6 +6,9 @@ pub(crate) use arch_impl::add; pub(crate) use arch_impl::digits; pub(crate) use arch_impl::word; +#[cfg(not(any(force_bits = "16", target_pointer_width = "16")))] +pub(crate) use arch_impl::ntt; + // Architecture choice. The logic works like this: // 1. If the configuration option force_bits is set to 16, 32 or 64, use generic__bit. // 2. Otherwise if target_arch is known, select that architecture. diff --git a/integer/src/arch/x86/mod.rs b/integer/src/arch/x86/mod.rs index bf7fccf5..c370f391 100644 --- a/integer/src/arch/x86/mod.rs +++ b/integer/src/arch/x86/mod.rs @@ -3,5 +3,8 @@ pub(crate) mod add; #[path = "../generic/digits.rs"] pub(crate) mod digits; +#[path = "../generic_32_bit/ntt.rs"] +pub(crate) mod ntt; + #[path = "../generic_32_bit/word.rs"] pub(crate) mod word; diff --git a/integer/src/arch/x86_64/mod.rs b/integer/src/arch/x86_64/mod.rs index 57d34fbe..fef5d136 100644 --- a/integer/src/arch/x86_64/mod.rs +++ b/integer/src/arch/x86_64/mod.rs @@ -3,5 +3,8 @@ pub(crate) mod add; #[path = "../generic/digits.rs"] pub(crate) mod digits; +#[path = "../generic_64_bit/ntt.rs"] +pub(crate) mod ntt; + #[path = "../generic_64_bit/word.rs"] pub(crate) mod word; diff --git a/integer/src/mul/mod.rs b/integer/src/mul/mod.rs index 5b2b1eaf..6f438957 100644 --- a/integer/src/mul/mod.rs +++ b/integer/src/mul/mod.rs @@ -14,18 +14,76 @@ use core::mem; use static_assertions::const_assert; /// If smaller operand length <= this, simple multiplication will be used. -const THRESHOLD_SIMPLE: usize = 24; -const_assert!(THRESHOLD_SIMPLE <= simple::MAX_SMALLER_LEN); -const_assert!(THRESHOLD_SIMPLE + 1 >= karatsuba::MIN_LEN); +const THRESHOLD_SIMPLE_DEFAULT: usize = 24; +const_assert!(THRESHOLD_SIMPLE_DEFAULT <= simple::MAX_SMALLER_LEN); +const_assert!(THRESHOLD_SIMPLE_DEFAULT + 1 >= karatsuba::MIN_LEN); /// If smaller operand length <= this, Karatsuba multiplication will be used. -const THRESHOLD_KARATSUBA: usize = 192; -const_assert!(THRESHOLD_KARATSUBA + 1 >= toom_3::MIN_LEN); +/// Tuned so that Toom-3 kicks in earlier (~96 words vs the old 192), +/// closing the gap with malachite/rug at ~10000-bit sizes. +const THRESHOLD_KARATSUBA_DEFAULT: usize = 96; +const_assert!(THRESHOLD_KARATSUBA_DEFAULT + 1 >= toom_3::MIN_LEN); + +/// If smaller operand length > this, NTT multiplication will be used. +#[cfg(not(any(force_bits = "16", target_pointer_width = "16")))] +const THRESHOLD_NTT_DEFAULT: usize = ntt::THRESHOLD_NTT; +/// NTT unavailable on 16/32-bit word targets — use `usize::MAX` so dispatch never +/// routes to the NTT path. +#[cfg(any(force_bits = "16", target_pointer_width = "16"))] +const THRESHOLD_NTT_DEFAULT: usize = usize::MAX; +#[cfg(not(any(force_bits = "16", target_pointer_width = "16")))] +const_assert!(THRESHOLD_NTT_DEFAULT + 1 >= toom_3::MIN_LEN); + +/// Environment-variable overrides for multiplication thresholds. +/// +/// When the `tuning` feature is active the user may set `DASHU_THRESHOLD_SIMPLE`, +/// `DASHU_THRESHOLD_KARATSUBA` or `DASHU_THRESHOLD_NTT` to override the +/// compile-time defaults. +mod threshold { + #[inline] + pub fn simple() -> usize { + #[cfg(feature = "tuning")] + { + if let Ok(s) = std::env::var("DASHU_THRESHOLD_SIMPLE") { + if let Ok(v) = s.parse() { + return v; + } + } + } + super::THRESHOLD_SIMPLE_DEFAULT + } + #[inline] + pub fn karatsuba() -> usize { + #[cfg(feature = "tuning")] + { + if let Ok(s) = std::env::var("DASHU_THRESHOLD_KARATSUBA") { + if let Ok(v) = s.parse() { + return v; + } + } + } + super::THRESHOLD_KARATSUBA_DEFAULT + } + #[inline] + pub fn ntt() -> usize { + #[cfg(feature = "tuning")] + { + if let Ok(s) = std::env::var("DASHU_THRESHOLD_NTT") { + if let Ok(v) = s.parse() { + return v; + } + } + } + super::THRESHOLD_NTT_DEFAULT + } +} mod helpers; mod karatsuba; +#[cfg(not(any(force_bits = "16", target_pointer_width = "16")))] +pub(crate) mod ntt; mod simple; -mod toom_3; +pub(crate) mod toom_3; /// Multiply a word sequence by a `Word` in place. /// @@ -156,17 +214,29 @@ pub fn sub_mul_word_same_len_in_place(words: &mut [Word], mult: Word, rhs: &[Wor } /// Temporary scratch space required for multiplication. -pub fn memory_requirement_up_to(_total_len: usize, smaller_len: usize) -> Layout { - if smaller_len <= THRESHOLD_SIMPLE { +pub fn memory_requirement_up_to(total_len: usize, smaller_len: usize) -> Layout { + if smaller_len <= threshold::simple() { memory::zero_layout() - } else if smaller_len <= THRESHOLD_KARATSUBA { + } else if smaller_len <= threshold::karatsuba() { karatsuba::memory_requirement_up_to(smaller_len) - } else { + } else if smaller_len <= threshold::ntt() { toom_3::memory_requirement_up_to(smaller_len) + } else { + // NTT path — only available on 64-bit word targets. + #[cfg(not(any(force_bits = "16", target_pointer_width = "16")))] + { + ntt::memory_requirement_up_to(total_len, smaller_len) + } + #[cfg(any(force_bits = "16", target_pointer_width = "16"))] + { + let _ = (total_len, smaller_len); + unreachable!("NTT unavailable on 16-bit targets"); + } } } /// Temporary scratch space required for multiplication. +#[inline] pub fn memory_requirement_exact(total_len: usize, smaller_len: usize) -> Layout { memory_requirement_up_to(total_len, smaller_len) } @@ -195,12 +265,22 @@ pub fn add_signed_mul<'a>( mem::swap(&mut a, &mut b); } - if b.len() <= THRESHOLD_SIMPLE { + if b.len() <= threshold::simple() { simple::add_signed_mul(c, sign, a, b, memory) - } else if b.len() <= THRESHOLD_KARATSUBA { + } else if b.len() <= threshold::karatsuba() { karatsuba::add_signed_mul(c, sign, a, b, memory) - } else { + } else if b.len() <= threshold::ntt() { toom_3::add_signed_mul(c, sign, a, b, memory) + } else { + #[cfg(not(any(force_bits = "16", target_pointer_width = "16")))] + { + ntt::add_signed_mul(c, sign, a, b, memory) + } + #[cfg(any(force_bits = "16", target_pointer_width = "16"))] + { + let _ = (c, sign, a, b, memory); + unreachable!("NTT unavailable on 16-bit targets"); + } } } @@ -218,11 +298,233 @@ pub fn add_signed_mul_same_len( let n = a.len(); debug_assert!(b.len() == n && c.len() == 2 * n); - if n <= THRESHOLD_SIMPLE { + if n <= threshold::simple() { simple::add_signed_mul_same_len(c, sign, a, b, memory) - } else if n <= THRESHOLD_KARATSUBA { + } else if n <= threshold::karatsuba() { karatsuba::add_signed_mul_same_len(c, sign, a, b, memory) - } else { + } else if n <= threshold::ntt() { toom_3::add_signed_mul_same_len(c, sign, a, b, memory) + } else { + #[cfg(not(any(force_bits = "16", target_pointer_width = "16")))] + { + ntt::add_signed_mul_same_len(c, sign, a, b, memory) + } + #[cfg(any(force_bits = "16", target_pointer_width = "16"))] + { + let _ = (c, sign, a, b, memory); + unreachable!("NTT unavailable on 16-bit targets"); + } + } +} + +#[cfg(all(test, feature = "std"))] +mod threshold_tests { + use super::*; + use crate::arch::word::Word; + use crate::Sign::Positive; + + /// Compare karatsuba vs toom-3 at various word counts to find [`THRESHOLD_KARATSUBA`]. + /// Run with: + /// cargo test -p dashu-int --release -- mul::threshold_tests::crossover_karatsuba --nocapture --ignored + #[test] + #[ignore] + #[cfg(feature = "std")] + fn crossover_karatsuba() { + use std::time::Instant; + + let sizes: &[usize] = &[80, 100, 120, 140, 160, 180, 200, 240, 280, 320, 360, 400]; + + println!("{:>8} {:>14} {:>14} {:>10}", "words", "karatsuba(µs)", "toom-3(µs)", "ratio"); + println!("{}", "-".repeat(50)); + + for &n in sizes { + let a: Vec = (0..n) + .map(|i| (i as Word + 1).wrapping_mul(0x9E3779B97F4A7C15u64 as Word)) + .collect(); + let b: Vec = (0..n) + .map(|i| (i as Word + 1).wrapping_mul(0xC6A4A7935BD1E995u64 as Word)) + .collect(); + let mut c_kara = vec![0 as Word; 2 * n]; + let mut c_toom = vec![0 as Word; 2 * n]; + let layout_kara = karatsuba::memory_requirement_up_to(n); + let layout_toom = toom_3::memory_requirement_up_to(n); + // Use the larger layout so both algorithms get enough memory. + let layout = if layout_kara.size() > layout_toom.size() { + layout_kara + } else { + layout_toom + }; + let warmup = 5; + let iters = 20; + + // Time karatsuba + let t_kara = { + let mut best = f64::MAX; + for _ in 0..warmup { + let mut alloc = crate::memory::MemoryAllocation::new(layout); + let mut mem = alloc.memory(); + c_kara.fill(0); + let _c = + karatsuba::add_signed_mul_same_len(&mut c_kara, Positive, &a, &b, &mut mem); + } + for _ in 0..iters { + let mut alloc = crate::memory::MemoryAllocation::new(layout); + let mut mem = alloc.memory(); + c_kara.fill(0); + let start = Instant::now(); + let _c = + karatsuba::add_signed_mul_same_len(&mut c_kara, Positive, &a, &b, &mut mem); + let elapsed = start.elapsed().as_secs_f64() * 1_000_000.0; + if elapsed < best { + best = elapsed; + } + } + best + }; + + // Time toom-3 + let t_toom = { + let mut best = f64::MAX; + for _ in 0..warmup { + let mut alloc = crate::memory::MemoryAllocation::new(layout); + let mut mem = alloc.memory(); + c_toom.fill(0); + let _c = + toom_3::add_signed_mul_same_len(&mut c_toom, Positive, &a, &b, &mut mem); + } + for _ in 0..iters { + let mut alloc = crate::memory::MemoryAllocation::new(layout); + let mut mem = alloc.memory(); + c_toom.fill(0); + let start = Instant::now(); + let _c = + toom_3::add_signed_mul_same_len(&mut c_toom, Positive, &a, &b, &mut mem); + let elapsed = start.elapsed().as_secs_f64() * 1_000_000.0; + if elapsed < best { + best = elapsed; + } + } + best + }; + + assert_eq!(&c_kara[..], &c_toom[..], "mismatch at n={n}"); + println!("{:>8} {:>14.1} {:>14.1} {:>9.2}x", n, t_kara, t_toom, t_toom / t_kara); + } + } + + /// Compare NTT against toom-3 at various word counts to find [`THRESHOLD_NTT`]. + /// + /// Run with (set a huge NTT threshold to keep toom-3 pure): + /// ```sh + /// DASHU_THRESHOLD_NTT=99999999 cargo test -p dashu-int --features tuning --release \ + /// -- mul::threshold_tests::crossover_ntt --ignored --nocapture + /// ``` + /// + /// The output is a table: words, b_pack, N, toom-3 time, NTT time, ratio. + #[test] + #[ignore] + #[allow(clippy::let_underscore_must_use)] + #[cfg(all( + feature = "std", + not(any( + force_bits = "16", + force_bits = "32", + target_pointer_width = "16", + target_pointer_width = "32" + )) + ))] + fn crossover_ntt() { + use std::time::Instant; + + let sizes: &[usize] = &[ + 1_000, 2_000, 3_000, 4_000, 5_000, 6_000, 7_000, 8_000, 9_000, 10_000, 20_000, 40_000, + 80_000, + ]; + + println!( + "{:>10} {:>4} {:>8} {:>12} {:>12} {:>10}", + "words", "bp", "N", "toom-3(ms)", "ntt(ms)", "ratio" + ); + println!("{}", "-".repeat(68)); + + for &n in sizes { + let a: Vec = (0..n) + .map(|i| (i as u64 + 1).wrapping_mul(0x9E3779B97F4A7C15)) + .collect(); + let b: Vec = (0..n) + .map(|i| (i as u64 + 1).wrapping_mul(0xC6A4A7935BD1E995)) + .collect(); + let mut c_toom = vec![0u64; 2 * n]; + let mut c_ntt = vec![0u64; 2 * n]; + + let layout_ntt = super::ntt::memory_requirement_up_to(2 * n, n); + let layout_toom = super::toom_3::memory_requirement_up_to(n); + let layout = if layout_ntt.size() > layout_toom.size() { + layout_ntt + } else { + layout_toom + }; + let warmup = 2; + let iters = 5; + + // toom-3 (may use NTT internally depending on DASHU_THRESHOLD_NTT) + let t_toom = { + let mut best = f64::MAX; + for _ in 0..warmup { + let mut alloc = crate::memory::MemoryAllocation::new(layout); + let mut mem = alloc.memory(); + c_toom.fill(0); + let _ = super::toom_3::add_signed_mul(&mut c_toom, Positive, &a, &b, &mut mem); + } + for _ in 0..iters { + let mut alloc = crate::memory::MemoryAllocation::new(layout); + let mut mem = alloc.memory(); + c_toom.fill(0); + let start = Instant::now(); + let _ = super::toom_3::add_signed_mul(&mut c_toom, Positive, &a, &b, &mut mem); + let elapsed = start.elapsed().as_secs_f64() * 1000.0; + if elapsed < best { + best = elapsed; + } + } + best + }; + + // NTT (via public entry, bypasses dispatch) + let t_ntt = { + let mut best = f64::MAX; + for _ in 0..warmup { + let mut alloc = crate::memory::MemoryAllocation::new(layout); + let mut mem = alloc.memory(); + c_ntt.fill(0); + let _ = super::ntt::add_signed_mul(&mut c_ntt, Positive, &a, &b, &mut mem); + } + for _ in 0..iters { + let mut alloc = crate::memory::MemoryAllocation::new(layout); + let mut mem = alloc.memory(); + c_ntt.fill(0); + let start = Instant::now(); + let _ = super::ntt::add_signed_mul(&mut c_ntt, Positive, &a, &b, &mut mem); + let elapsed = start.elapsed().as_secs_f64() * 1000.0; + if elapsed < best { + best = elapsed; + } + } + best + }; + + assert_eq!(&c_ntt[..], &c_toom[..], "mismatch at n={n}"); + + let (b_pack, nn, _k_eff) = super::ntt::select_params(n, n); + println!( + "{:>10} {:>4} {:>8} {:>12.3} {:>12.3} {:>9.2}x", + n, + b_pack, + nn, + t_toom, + t_ntt, + t_ntt / t_toom + ); + } } } diff --git a/integer/src/mul/ntt/compute_constants.py b/integer/src/mul/ntt/compute_constants.py new file mode 100644 index 00000000..dfc3e864 --- /dev/null +++ b/integer/src/mul/ntt/compute_constants.py @@ -0,0 +1,149 @@ +#!/usr/bin/env python3 +"""Compute omega_max and CRT constants for Proth NTT primes. + +Prints the computed constants in a readable format — does NOT generate Rust code. +Copy the values into the arch ntt.rs files by hand. +""" + +# --- 64-bit Proth primes --- +PRIMES_64 = [ + (0x3a00000000000001, 57, 29), # Proth(57, 29) + (0x8e00000000000001, 57, 71), # Proth(57, 71) + (0x9600000000000001, 57, 75), # Proth(57, 75) +] +MAX_LOG_N_64 = 57 + +# --- 32-bit Proth primes --- +PRIMES_32 = [ + (0x1c000001, 26, 7), # Proth(26, 7) + (0x78000001, 27, 15), # Proth(27, 15) + (0x88000001, 27, 17), # Proth(27, 17) +] +MAX_LOG_N_32 = 26 + + +def mod_pow(base, exp, mod): + """base**exp mod mod.""" + result = 1 + while exp > 0: + if exp & 1: + result = (result * base) % mod + base = (base * base) % mod + exp >>= 1 + return result + + +def mod_inv(a, mod): + """Inverse of a mod mod (mod is prime).""" + return mod_pow(a, mod - 2, mod) + + +def factorize(n): + """Return list of distinct prime factors of n.""" + factors = [] + d = 2 + m = n + while d * d <= m: + if m % d == 0: + factors.append(d) + while m % d == 0: + m //= d + d += 1 if d == 2 else 2 # skip even after 2 + if m > 1: + factors.append(m) + return factors + + +def is_primitive_root(g, p, factors_of_pm1): + """Check if g is a primitive root mod p.""" + for q in factors_of_pm1: + if mod_pow(g, (p - 1) // q, p) == 1: + return False + return True + + +def find_primitive_root(p): + """Find a primitive root mod p by brute force.""" + factors = factorize(p - 1) + for g in range(2, min(p, 2000)): + if is_primitive_root(g, p, factors): + return g + raise ValueError(f"No primitive root found for p = {p} (tried up to 2000)") + + +def compute_omega(p, g, max_log_n): + """omega = g^((p-1) / 2^max_log_n) mod p.""" + assert (p - 1) % (1 << max_log_n) == 0, \ + f"max_log_n={max_log_n} does not divide p-1 for p={p:#x}" + exp = (p - 1) >> max_log_n + return mod_pow(g, exp, p) + + +def verify_omega(omega, p, max_log_n): + """Verify omega^(2^(max_log_n-1)) == -1 and omega^(2^max_log_n) == 1.""" + w = omega + for _ in range(max_log_n - 1): + w = (w * w) % p + assert w == p - 1, f"omega^(2^{max_log_n-1}) != -1 mod p, got {w:#x}" + w = (w * w) % p + assert w == 1, f"omega^(2^{max_log_n}) != 1 mod p, got {w:#x}" + + +def compute_crt_constants(primes): + """Compute Garner CRT: inv(p_i mod p_j) mod p_j for i < j.""" + k = len(primes) + crt = [[0] * k for _ in range(k)] + for i in range(k): + for j in range(i + 1, k): + pi = primes[i] + pj = primes[j] + crt[i][j] = mod_inv(pi % pj, pj) + return crt + + +def print_results(name, primes_data, max_log_n): + """Pretty-print computed constants for one architecture.""" + primes = [p for p, _, _ in primes_data] + + print(f"===== {name} =====") + print(f" MAX_LOG_N = {max_log_n}") + print() + for i, (p, n, k) in enumerate(primes_data): + print(f" PI={i}: p = {p:#018x} ({k} * 2^{n} + 1)") + v2 = ((p - 1) & -(p - 1)).bit_length() - 1 # trailing zeros + print(f" v2(p-1) = {v2}") + + print() + + # Primitive roots & omega + for i, (p, n, k) in enumerate(primes_data): + print(f" PI={i}: finding primitive root...") + g = find_primitive_root(p) + omega = compute_omega(p, g, max_log_n) + verify_omega(omega, p, max_log_n) + print(f" g = {g}") + bit_width = 64 if max_log_n == 57 else 32 + print(f" omega_max = {omega:#0{bit_width//4 + 2}x}") + + print() + + # CRT constants + crt = compute_crt_constants(primes) + bit_width = 64 if max_log_n == 57 else 32 + print(f" CRT_INV_IJ:") + for i in range(len(primes)): + for j in range(len(primes)): + if crt[i][j] != 0: + print(f" inv(p{i} mod p{j}) = {crt[i][j]:#0{bit_width//4 + 2}x}") + + # Two-prime product for headroom checks + prod_01 = primes[0] * primes[1] + print(f"\n p0 * p1 = {prod_01:#x}") + + print() + + +# --- Main --- +if __name__ == "__main__": + print_results("64-bit", PRIMES_64, MAX_LOG_N_64) + print_results("32-bit", PRIMES_32, MAX_LOG_N_32) diff --git a/integer/src/mul/ntt/crt.rs b/integer/src/mul/ntt/crt.rs new file mode 100644 index 00000000..af6ba991 --- /dev/null +++ b/integer/src/mul/ntt/crt.rs @@ -0,0 +1,221 @@ +//! Garner CRT: combine `K` residues modulo `K` primes into a small integer. + +use crate::arch::ntt::K; +use num_modular::ModularCoreOps; + +/// Accumulator for Garner CRT. +/// +/// Implemented by [`crate::arch::word::TripleWord`] (192 bits on 64-bit +/// targets, 96 bits on 32-bit targets). +pub trait CrtAccum: Default + Copy { + type Lane: Copy + + Default + + Into + + for<'a> ModularCoreOps; + fn from_lane(v: Self::Lane) -> Self; + /// `self += t * factor` + fn add_product(&mut self, t: Self::Lane, factor: u128); + /// `self mod m` + fn rem_lane(&self, m: Self::Lane) -> Self::Lane; + /// Write the value into `out` as little-endian `Word` values, + /// returning the number of non-zero words written. + fn write_words(&self, out: &mut [crate::arch::word::Word; 6]) -> u32; +} + +// ── Garner combine ───────────────────────────────────────────────────── + +/// Combine `K` residues into a [`CrtAccum`] via Garner's algorithm. +/// +/// All arithmetic is standard-form. `crt_inv_ij[i][j]` +/// holds `inv(p_i mod p_j) mod p_j` for `i < j`. +/// `primes` contains the prime values (only `primes[0..k]` are used). +pub fn garner_combine( + residues: &[A::Lane], + crt_inv_ij: &[[A::Lane; K]; K], + primes: &[A::Lane; K], +) -> A { + let k = residues.len(); + assert!(k <= K, "CRT supports up to {K} primes"); + + let p0 = primes[0]; + let p1 = primes[1]; + let p2 = primes[2]; + let mut x = A::from_lane(residues[0]); + + if k == 1 { + return x; + } + + // t_1 = (r_1 - x mod p1) * inv(p0 mod p1) mod p1 + let x_mod_p1 = x.rem_lane(p1); + let diff1 = residues[1].subm(x_mod_p1, &p1); + let t1 = diff1.mulm(crt_inv_ij[0][1], &p1); + x.add_product(t1, p0.into()); + + if k == 2 { + return x; + } + + // t_2 = (r_2 - x mod p2) * inv(p0*p1 mod p2) mod p2 + let x_mod_p2 = x.rem_lane(p2); + let diff2 = residues[2].subm(x_mod_p2, &p2); + let inv_prod = crt_inv_ij[0][2].mulm(crt_inv_ij[1][2], &p2); + let t2 = diff2.mulm(inv_prod, &p2); + x.add_product(t2, p0.into() * p1.into()); + + x +} + +// ── TripleWord impls (cfg-gated per arch) ───────────────────────────── + +/// 64-bit: 3 × u64 = 192 bits. +#[cfg(not(any(force_bits = "32", target_pointer_width = "32")))] +mod triple_impl { + use super::CrtAccum; + use crate::arch::word::TripleWord; + + impl CrtAccum for TripleWord { + type Lane = u64; + + #[inline] + fn from_lane(v: u64) -> Self { + TripleWord([v, 0, 0]) + } + + #[inline] + fn add_product(&mut self, t: u64, factor: u128) { + let fac_lo = factor as u64; + let fac_hi = (factor >> 64) as u64; + let m_lo_full = (t as u128) * (fac_lo as u128); + let lo = m_lo_full as u64; + let m_lo = (m_lo_full >> 64) as u64; + let m_hi_full = (t as u128) * (fac_hi as u128); + let m_hi = m_hi_full as u64; + let hi = (m_hi_full >> 64) as u64; + let (mid, c) = m_lo.overflowing_add(m_hi); + let hi_word = hi.wrapping_add(c as u64); + let (r0, c0) = self.0[0].overflowing_add(lo); + self.0[0] = r0; + let (r1, c1) = self.0[1].overflowing_add(mid.wrapping_add(c0 as u64)); + self.0[1] = r1; + self.0[2] = self.0[2].wrapping_add(hi_word.wrapping_add(c1 as u64)); + } + + #[inline] + fn rem_lane(&self, m: u64) -> u64 { + let m128 = m as u128; + let mut r: u128 = 0; + for &word in self.0.iter().rev() { + r = (r << 64) | (word as u128); + r %= m128; + } + r as u64 + } + + #[inline] + fn write_words(&self, out: &mut [crate::arch::word::Word; 6]) -> u32 { + out[0] = self.0[0]; + out[1] = self.0[1]; + out[2] = self.0[2]; + if self.0[2] != 0 { + 3 + } else if self.0[1] != 0 { + 2 + } else { + 1 + } + } + } +} + +/// 32-bit: 3 × u32 = 96 bits. +#[cfg(any(force_bits = "32", target_pointer_width = "32"))] +mod triple_impl { + use super::CrtAccum; + use crate::arch::word::TripleWord; + + impl CrtAccum for TripleWord { + type Lane = u32; + + #[inline] + fn from_lane(v: u32) -> Self { + TripleWord([v, 0, 0]) + } + + #[inline] + fn add_product(&mut self, t: u32, factor: u128) { + let factor_lo = factor as u32; + let factor_hi = (factor >> 32) as u32; + let m_lo = (t as u64) * (factor_lo as u64); + let lo = m_lo as u32; + let m_mid = (m_lo >> 32) as u32; + let m_hi = (t as u64) * (factor_hi as u64) + m_mid as u64; + let mid = m_hi as u32; + let hi = (m_hi >> 32) as u32; + let (r0, c0) = self.0[0].overflowing_add(lo); + self.0[0] = r0; + let (r1, c1) = self.0[1].overflowing_add(mid.wrapping_add(c0 as u32)); + self.0[1] = r1; + self.0[2] = self.0[2].wrapping_add(hi.wrapping_add(c1 as u32)); + } + + #[inline] + fn rem_lane(&self, m: u32) -> u32 { + let m64 = m as u64; + let mut r: u64 = 0; + for &word in self.0.iter().rev() { + r = (r << 32) | (word as u64); + r %= m64; + } + r as u32 + } + + #[inline] + fn write_words(&self, out: &mut [crate::arch::word::Word; 6]) -> u32 { + out[0] = self.0[0]; + out[1] = self.0[1]; + out[2] = self.0[2]; + if self.0[2] != 0 { + 3 + } else if self.0[1] != 0 { + 2 + } else { + 1 + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::arch::word::TripleWord; + #[cfg(not(feature = "std"))] + use alloc::vec; + + #[test] + fn test_garner_roundtrip() { + use crate::arch::ntt::{CRT_INV_IJ, MODULI}; + + type Lane = ::Lane; + let p0 = MODULI[0]; + let p1 = MODULI[1]; + let p2 = MODULI[2]; + let primes = [p0, p1, p2]; + + let residues = vec![12345u64 as Lane, 67890u64 as Lane, 11111u64 as Lane]; + let x = garner_combine::(&residues, &CRT_INV_IJ, &primes); + assert_eq!(x.rem_lane(p0), residues[0]); + assert_eq!(x.rem_lane(p1), residues[1]); + assert_eq!(x.rem_lane(p2), residues[2]); + + let x = garner_combine::(&residues[..2], &CRT_INV_IJ, &primes); + assert_eq!(x.rem_lane(p0), residues[0]); + assert_eq!(x.rem_lane(p1), residues[1]); + + let x = garner_combine::(&residues[..1], &CRT_INV_IJ, &primes); + let mut buf = [crate::arch::word::Word::default(); 6]; + x.write_words(&mut buf); + assert_eq!(buf[0], residues[0]); + } +} diff --git a/integer/src/mul/ntt/mod.rs b/integer/src/mul/ntt/mod.rs new file mode 100644 index 00000000..a0f21200 --- /dev/null +++ b/integer/src/mul/ntt/mod.rs @@ -0,0 +1,690 @@ +//! NTT-based multiplication for very large integers. +//! +//! Uses Number Theoretic Transforms over Proth primes of the form +//! `K * 2^N + 1` combined with the Chinese Remainder Theorem (CRT). + +use crate::{ + add, + arch::word::{SignedWord, Word}, + memory::{self, Memory}, + Sign::{self, *}, +}; +use alloc::alloc::Layout; +use core::mem; + +pub(crate) mod crt; +mod pack; +mod transform; + +use crate::arch::ntt::{ + B_PACK_CANDIDATES, B_PACK_MIN, CRT_INV_IJ, K, MAX_LOG_N, MODULI, OMEGA_MAX, P0, P1, P2, +}; +use crate::mul::ntt::crt::{garner_combine, CrtAccum}; +use num_modular::Reducer; + +/// Minimum smaller-operand length (in words) for the NTT path. +/// +/// Crossover with Toom-3 lies at ~3 200 words; chosen at 4 000 where +/// NTT is a clear 30%+ faster. +pub const THRESHOLD_NTT: usize = 4_000; + +/// Select NTT parameters for operands with the given word lengths. +/// +/// Returns `(b_pack, N, K_eff)`. +pub fn select_params(la_words: usize, lb_words: usize) -> (u32, usize, usize) { + let word_bits = Word::BITS; + let la_bits = la_words as u64 * word_bits as u64; + let lb_bits = lb_words as u64 * word_bits as u64; + let prod_2 = (MODULI[0] as u128) * (MODULI[1] as u128); + + for &b_pack in B_PACK_CANDIDATES { + let coeffs_a = (la_bits + b_pack as u64 - 1) / b_pack as u64; + let coeffs_b = (lb_bits + b_pack as u64 - 1) / b_pack as u64; + let total_coeffs = (coeffs_a + coeffs_b - 1) as usize; + let n = total_coeffs.next_power_of_two().max(2); + + if (n.trailing_zeros()) > MAX_LOG_N { + continue; + } + + // Compute max coefficient value, guarding against u128 overflow for + // b_pack = 64 where (2^64−1)² ≈ 2^128 and n/2 can push it past 2^128. + let coeff_max = (1u128 << b_pack) - 1; + let max_coeff = coeff_max + .checked_mul(coeff_max) + .and_then(|sq| (n as u128 / 2).checked_mul(sq)); + + let k_eff = match max_coeff { + Some(mc) if mc < prod_2 => 2, + _ => K, + }; + return (b_pack, n, k_eff); + } + + unreachable!( + "b_pack = {} always passes the headroom check", + B_PACK_CANDIDATES.last().unwrap() + ) +} + +/// Estimate bit length from a word slice (excludes leading zeros). +fn bit_len(words: &[Word]) -> u64 { + let leading_zeros = words.iter().rev().take_while(|&&w| w == 0).count(); + let used = words.len() - leading_zeros; + if used == 0 { + return 0; + } + let hi_word = words[used - 1]; + let hi_bits = Word::BITS - hi_word.leading_zeros(); + (used as u64 - 1) * Word::BITS as u64 + hi_bits as u64 +} + +/// Count number of coefficients needed for a given bit length. +fn coeff_count(bit_len: u64, b_pack: u32) -> usize { + ((bit_len + b_pack as u64 - 1) / b_pack as u64) as usize +} + +/// Worst-case scratch memory bound. +pub fn memory_requirement_up_to(total_len: usize, _smaller_len: usize) -> Layout { + use crate::arch::ntt::Lane; + + let word_bits = Word::BITS; + let max_coeffs = + (total_len as u64 * word_bits as u64 + B_PACK_MIN as u64 - 1) / B_PACK_MIN as u64; + let n_max = ((max_coeffs + 1) as usize).next_power_of_two().max(2); + + let lanes = 2 * n_max; + let residues = K * n_max; + let twiddles = n_max; + let product = total_len; + + let lane_bytes = mem::size_of::(); + let word_bytes = mem::size_of::(); + + let lanes_words = lanes * lane_bytes / word_bytes; + let residues_words = residues * lane_bytes / word_bytes; + let twiddles_words = twiddles * lane_bytes / word_bytes; + let total_words = product + lanes_words + residues_words + twiddles_words; + + memory::array_layout::(total_words) +} + +/// `c += sign * a * b` with equal-length operands. +/// +/// Returns carry. +#[must_use] +#[inline] +pub fn add_signed_mul_same_len( + c: &mut [Word], + sign: Sign, + a: &[Word], + b: &[Word], + memory: &mut Memory, +) -> SignedWord { + let n = a.len(); + debug_assert!(b.len() == n && c.len() == 2 * n); + add_signed_mul_conv(c, sign, a, b, memory) +} + +/// `c += sign * a * b` (general, a may be longer than b). +/// +/// When `a ≫ b` the implementation forks: +/// - If `b` is below [`THRESHOLD_NTT`], dispatch already routes to +/// `toom_3::add_signed_mul` (which uses +/// `add_signed_mul_split_into_chunks` from +/// [`helpers`](crate::mul::helpers)). +/// - If `b` is above [`THRESHOLD_NTT`], this function pre-transforms +/// `b` once per prime and reuses `b̂` across chunks of `a` via +/// [`add_signed_mul_chunked`]. +/// +/// Returns carry. +#[must_use] +#[inline] +pub fn add_signed_mul( + c: &mut [Word], + sign: Sign, + a: &[Word], + b: &[Word], + memory: &mut Memory, +) -> SignedWord { + debug_assert!(a.len() >= b.len() && c.len() == a.len() + b.len()); + if a.len() > 2 * b.len() { + return add_signed_mul_chunked(c, sign, a, b, memory); + } + add_signed_mul_conv(c, sign, a, b, memory) +} + +/// NTT multiplication with asymmetric chunking. +/// +/// When `la > 2 * lb`, transform `b` once and reuse `b̂` across chunks +/// of `a`, reducing total transform work from O((la+lb)·log(la+lb)) +/// to O(la + lb·log(lb)). +fn add_signed_mul_chunked( + c: &mut [Word], + sign: Sign, + a: &[Word], + b: &[Word], + memory: &mut Memory, +) -> SignedWord { + use crate::arch::ntt::Lane; + use crate::mul::helpers::add_signed_mul_split_into_chunks; + + let lb = b.len(); + let chunk_len = lb * 2; + + // Parameters for chunk-sized transforms. + let (b_pack, nn_chunk, k_eff) = select_params(chunk_len, lb); // a_chunk ≈ 2*lb + + // ---- Allocate long-lived buffers ---- + + // Per-prime forward-transformed b̂ and cached twiddles (fwd + inv). + // Twiddles depend only on (pi, nn_chunk, omega_max) — precompute + // once so the per-chunk closure can copy instead of recomputing. + let b_hat_len = k_eff * nn_chunk; + let twiddle_len = k_eff * (nn_chunk / 2); + let (b_hat, mut mem) = memory.allocate_slice_fill::(b_hat_len, 0); + let (fwd_tw_cache, mut mem) = mem.allocate_slice_fill::(twiddle_len, 0); + let (inv_tw_cache, mut mem) = mem.allocate_slice_fill::(twiddle_len, 0); + + // ---- Transform b once per prime; also precompute twiddles ---- + let geom = NttGeometry { + nn: nn_chunk, + b_pack, + k_eff, + output_coeffs: 0, // unused by prepare_b_hat_and_twiddles + }; + prepare_b_hat_and_twiddles(b_hat, fwd_tw_cache, inv_tw_cache, b, &geom, &mut mem); + + // ---- Setup for the closure ---- + let lb_bits = bit_len(b); + let coeffs_b = coeff_count(lb_bits, b_pack); + + // ---- Chunked multiply ---- + add_signed_mul_split_into_chunks( + c, + sign, + a, + b, + chunk_len, + &mut mem, + |c_slice, sign, a_chunk, b, mem| { + let a_bits = bit_len(a_chunk); + if a_bits == 0 { + return 0; + } + let coeffs_a = coeff_count(a_bits, b_pack); + let output_coeffs = coeffs_a + coeffs_b - 1; + let out_words = a_chunk.len() + b.len(); + + let geom = NttGeometry { + nn: nn_chunk, + b_pack, + k_eff, + output_coeffs, + }; + run_ntt_pipeline( + a_chunk, + b_hat, + fwd_tw_cache, + inv_tw_cache, + &geom, + out_words, + c_slice, + sign, + mem, + ) + }, + ) +} + +/// Run the full NTT pipeline: allocate → per-prime transform → CRT → fold into `c_out`. +/// +/// `b_hat`, `fwd_tw_cache`, and `inv_tw_cache` must have been precomputed by the +/// caller (pack + Montgomery convert + bit-reverse + forward-transform for `b_hat`; +/// forward/inverse twiddle tables for the caches). See `transform_b_forward` and +/// `transform::precompute_twiddles`. +/// +/// Shared body of `add_signed_mul_conv` and the per-chunk callback in +/// `add_signed_mul_chunked`. +#[allow(clippy::too_many_arguments)] +fn run_ntt_pipeline( + a: &[Word], + b_hat: &[crate::arch::ntt::Lane], + fwd_tw_cache: &[crate::arch::ntt::Lane], + inv_tw_cache: &[crate::arch::ntt::Lane], + geom: &NttGeometry, + out_words: usize, + c_out: &mut [Word], + sign: Sign, + mem: &mut Memory, +) -> SignedWord { + use crate::arch::ntt::Lane; + + let nn = geom.nn; + let k_eff = geom.k_eff; + + let (prod, mut m) = mem.allocate_slice_fill::(out_words, 0); + let (residues, mut m) = m.allocate_slice_fill::(k_eff * nn, 0); + let (a_lane, mut m) = m.allocate_slice_fill::(nn, 0); + let (b_lane, mut m) = m.allocate_slice_fill::(nn, 0); + let (fwd_twiddles, mut m) = m.allocate_slice_fill::(nn / 2, 0); + let (inv_twiddles, _) = m.allocate_slice_fill::(nn / 2, 0); + + let mut ctx = TransformCtx { + a_lane, + b_lane, + fwd_twiddles, + inv_twiddles, + geom: NttGeometry { ..*geom }, + }; + + for pi in 0..k_eff { + let tw_off = pi * (nn / 2); + ctx.fwd_twiddles + .copy_from_slice(&fwd_tw_cache[tw_off..tw_off + nn / 2]); + ctx.inv_twiddles + .copy_from_slice(&inv_tw_cache[tw_off..tw_off + nn / 2]); + + let b_hat_slice = &b_hat[pi * nn..(pi + 1) * nn]; + + match pi { + 0 => process_prime(a, b_hat_slice, &mut ctx, residues, pi, &P0), + 1 => process_prime(a, b_hat_slice, &mut ctx, residues, pi, &P1), + 2 => process_prime(a, b_hat_slice, &mut ctx, residues, pi, &P2), + _ => unreachable!(), + } + } + + do_crt::(prod, residues, &ctx, &MODULI, &CRT_INV_IJ); + + match sign { + Positive => add::add_signed_in_place(&mut c_out[..out_words], Positive, &prod[..out_words]), + Negative => add::add_signed_in_place(&mut c_out[..out_words], Negative, &prod[..out_words]), + } +} + +/// Core implementation: c += sign * a * b. +/// +/// Does a single NTT convolution of the full operands (no chunking). +fn add_signed_mul_conv( + c: &mut [Word], + sign: Sign, + a: &[Word], + b: &[Word], + memory: &mut Memory, +) -> SignedWord { + use crate::arch::ntt::Lane; + + let la = a.len(); + let lb = b.len(); + + debug_assert!(la > 0 && lb > 0); + let (b_pack, nn, k_eff) = select_params(la, lb); + let la_bits = bit_len(a); + let lb_bits = bit_len(b); + debug_assert!(la_bits > 0 && lb_bits > 0); + + let coeffs_a = coeff_count(la_bits, b_pack); + let coeffs_b = coeff_count(lb_bits, b_pack); + let output_coeffs = coeffs_a + coeffs_b - 1; + + // Pre-transform b and precompute twiddles. + let b_hat_len = k_eff * nn; + let twiddle_len = k_eff * (nn / 2); + let (b_hat, mut mem) = memory.allocate_slice_fill::(b_hat_len, 0); + let (fwd_tw_cache, mut mem) = mem.allocate_slice_fill::(twiddle_len, 0); + let (inv_tw_cache, mut mem) = mem.allocate_slice_fill::(twiddle_len, 0); + + let geom = NttGeometry { + nn, + b_pack, + k_eff, + output_coeffs, + }; + prepare_b_hat_and_twiddles(b_hat, fwd_tw_cache, inv_tw_cache, b, &geom, &mut mem); + run_ntt_pipeline(a, b_hat, fwd_tw_cache, inv_tw_cache, &geom, la + lb, c, sign, &mut mem) +} + +/// CRT + accumulate, generic over the accumulator type. +fn do_crt( + prod: &mut [Word], + residues: &[A::Lane], + ctx: &TransformCtx<'_>, + primes: &[A::Lane; K], + crt_inv: &[[A::Lane; K]; K], +) { + let g = &ctx.geom; + for k in 0..g.output_coeffs { + let mut coeff_residues = [A::Lane::default(); 3]; + #[allow(clippy::needless_range_loop)] + for pi in 0..g.k_eff { + coeff_residues[pi] = residues[pi * g.nn + k]; + } + let crt_val = garner_combine::(&coeff_residues[..g.k_eff], crt_inv, primes); + let mut crt_buf = [Word::default(); 6]; + let crt_n = crt_val.write_words(&mut crt_buf); + add_shifted_to_prod(prod, &crt_buf[..crt_n as usize], crt_n, k, g.b_pack); + } +} + +/// Geometry constants for an NTT pipeline invocation. +struct NttGeometry { + nn: usize, + b_pack: u32, + k_eff: usize, + output_coeffs: usize, +} + +/// Scratch buffers and geometry for the per-prime NTT pipeline. +struct TransformCtx<'a> { + a_lane: &'a mut [crate::arch::ntt::Lane], + b_lane: &'a mut [crate::arch::ntt::Lane], + fwd_twiddles: &'a mut [crate::arch::ntt::Lane], + inv_twiddles: &'a mut [crate::arch::ntt::Lane], + geom: NttGeometry, +} + +/// Transform `b` and leave the result in `b_lane` (forward-transformed, +/// Montgomery form). `fwd_twiddles` must already be precomputed. +fn transform_b_forward>( + b_lane: &mut [crate::arch::ntt::Lane], + b: &[Word], + nn: usize, + b_pack: u32, + fwd_twiddles: &[crate::arch::ntt::Lane], + r: &R, +) { + pack::pack(b_lane, b, b_pack, nn); + for c in b_lane[..nn].iter_mut() { + *c = r.transform(*c); + } + transform::bit_reverse(&mut b_lane[..nn]); + transform::forward(&mut b_lane[..nn], fwd_twiddles, r); +} + +/// Pre-transform `b` and precompute twiddles, storing results into the +/// pre-allocated cache slices. +/// +/// `b_hat` must have length `geom.k_eff * geom.nn`, `fwd_tw_cache` and +/// `inv_tw_cache` each `geom.k_eff * (geom.nn / 2)`. +fn prepare_b_hat_and_twiddles( + b_hat: &mut [crate::arch::ntt::Lane], + fwd_tw_cache: &mut [crate::arch::ntt::Lane], + inv_tw_cache: &mut [crate::arch::ntt::Lane], + b: &[Word], + geom: &NttGeometry, + mem: &mut Memory, +) { + use crate::arch::ntt::Lane; + + let nn = geom.nn; + let b_pack = geom.b_pack; + let k_eff = geom.k_eff; + + for (pi, &omega) in OMEGA_MAX.iter().enumerate().take(k_eff) { + let (b_lane, mut rest) = mem.allocate_slice_fill::(nn, 0); + let (fwd_tw, mut rest) = rest.allocate_slice_fill::(nn / 2, 0); + let (inv_tw, _) = rest.allocate_slice_fill::(nn / 2, 0); + + match pi { + 0 => { + transform::precompute_twiddles(fwd_tw, nn, omega, false, &P0); + transform::precompute_twiddles(inv_tw, nn, omega, true, &P0); + transform_b_forward(b_lane, b, nn, b_pack, fwd_tw, &P0); + } + 1 => { + transform::precompute_twiddles(fwd_tw, nn, omega, false, &P1); + transform::precompute_twiddles(inv_tw, nn, omega, true, &P1); + transform_b_forward(b_lane, b, nn, b_pack, fwd_tw, &P1); + } + 2 => { + transform::precompute_twiddles(fwd_tw, nn, omega, false, &P2); + transform::precompute_twiddles(inv_tw, nn, omega, true, &P2); + transform_b_forward(b_lane, b, nn, b_pack, fwd_tw, &P2); + } + _ => unreachable!(), + } + + let b_off = pi * nn; + let tw_off = pi * (nn / 2); + b_hat[b_off..b_off + nn].copy_from_slice(b_lane); + fwd_tw_cache[tw_off..tw_off + nn / 2].copy_from_slice(fwd_tw); + inv_tw_cache[tw_off..tw_off + nn / 2].copy_from_slice(inv_tw); + } +} + +/// Per-prime NTT pipeline. +/// +/// `b_hat_slice` must already be forward-transformed (packed, Montgomery +/// form, bit-reversed). `ctx.fwd_twiddles` and `ctx.inv_twiddles` must +/// already be precomputed. +fn process_prime>( + a: &[Word], + b_hat_slice: &[crate::arch::ntt::Lane], + ctx: &mut TransformCtx<'_>, + residues: &mut [crate::arch::ntt::Lane], + pi: usize, + r: &R, +) { + let nn = ctx.geom.nn; + let b_pack = ctx.geom.b_pack; + + // Transform a + pack::pack(ctx.a_lane, a, b_pack, nn); + for c in ctx.a_lane[..nn].iter_mut() { + *c = r.transform(*c); + } + transform::bit_reverse(&mut ctx.a_lane[..nn]); + transform::forward(&mut ctx.a_lane[..nn], ctx.fwd_twiddles, r); + + // Copy pre-transformed b + ctx.b_lane[..nn].copy_from_slice(b_hat_slice); + + transform::pointwise_mul(&mut ctx.a_lane[..nn], &ctx.b_lane[..nn], r); + transform::inverse(&mut ctx.a_lane[..nn], ctx.inv_twiddles, r); + for c in ctx.a_lane[..nn].iter_mut() { + *c = r.residue(*c); + } + + let offset = pi * nn; + residues[offset..offset + nn].copy_from_slice(&ctx.a_lane[..nn]); +} + +/// Add a CRT value (as `Word`-sized limbs) to `prod`, shifted left by +/// `k * b_pack` bits. +fn add_shifted_to_prod(prod: &mut [Word], words: &[Word], count: u32, k: usize, b_pack: u32) { + let shift_bits = (k as u32).wrapping_mul(b_pack); + let word_bits = Word::BITS; + let start_idx = (shift_bits / word_bits) as usize; + let bit_shift = shift_bits % word_bits; + + let mut carry: Word = 0; + + for (vi, &word) in words.iter().enumerate().take(count as usize) { + let limb = word.wrapping_add(carry); + let idx = start_idx + vi; + if idx >= prod.len() { + return; + } + + if bit_shift == 0 { + let (r, c) = prod[idx].overflowing_add(limb); + prod[idx] = r; + carry = Word::from(c); + } else { + let val = (limb as u128) << bit_shift; + let lo = val as Word; + let hi = (val >> word_bits) as Word; + + let (r, c1) = prod[idx].overflowing_add(lo); + prod[idx] = r; + carry = Word::from(c1).wrapping_add(hi); + + if idx + 1 < prod.len() && carry != 0 { + let (r2, c2) = prod[idx + 1].overflowing_add(carry); + prod[idx + 1] = r2; + carry = Word::from(c2); + } + } + } + + let mut idx = start_idx + count as usize; + while carry != 0 && idx < prod.len() { + let (r, c) = prod[idx].overflowing_add(carry); + prod[idx] = r; + carry = Word::from(c); + idx += 1; + } +} + +#[cfg(test)] +mod tests { + use super::*; + #[cfg(not(feature = "std"))] + use alloc::vec; + #[cfg(not(feature = "std"))] + use alloc::vec::Vec; + + #[test] + fn test_select_params_small() { + let (b_pack, n, k_eff) = select_params(10, 10); + // On 64-bit: B_PACK_CANDIDATES[0] = 64, needs K_eff = 3. + // On 32-bit: B_PACK_CANDIDATES[0] = 32, likely K_eff = 2. + assert!(b_pack >= 32); + assert!(n >= 2 && n.is_power_of_two()); + assert!((2..=K).contains(&k_eff)); + } + + #[test] + fn test_select_params_large() { + let (b_pack, n, _k_eff) = select_params(THRESHOLD_NTT, THRESHOLD_NTT); + assert!(b_pack >= 32); + assert!(n.is_power_of_two()); + let coeffs_a = + (THRESHOLD_NTT * Word::BITS as usize + b_pack as usize - 1) / b_pack as usize; + let coeffs_b = coeffs_a; + let min_n = (coeffs_a + coeffs_b).next_power_of_two().max(2); + assert!(n >= min_n, "n={n} < min_n={min_n}"); + } + + #[test] + fn test_bit_len() { + assert_eq!(bit_len(&[]), 0); + assert_eq!(bit_len(&[0]), 0); + assert_eq!(bit_len(&[1]), 1); + } + + #[test] + fn test_ntt_sign_negative() { + let a: Vec = (0..30).map(|i| (i as Word + 1) * 100).collect(); + let b: Vec = (0..30).map(|i| (i as Word + 1) * 200).collect(); + + let mut c = vec![0u64 as Word; a.len() + b.len()]; + let layout = memory_requirement_up_to(c.len(), b.len()); + let mut alloc = crate::memory::MemoryAllocation::new(layout); + let mut memory = alloc.memory(); + let _ = add_signed_mul_conv(&mut c, Positive, &a, &b, &mut memory); + + let layout2 = memory_requirement_up_to(c.len(), b.len()); + let mut alloc2 = crate::memory::MemoryAllocation::new(layout2); + let mut memory2 = alloc2.memory(); + let _ = add_signed_mul_conv(&mut c, Negative, &a, &b, &mut memory2); + + assert!(c.iter().all(|&w| w == 0)); + } + + /// Naive schoolbook multiplication for comparison. + fn schoolbook_mul(a: &[Word], b: &[Word]) -> Vec { + let mut c = vec![0u64 as Word; a.len() + b.len()]; + for (i, &ai) in a.iter().enumerate() { + let mut carry: u128 = 0; + for (j, &bj) in b.iter().enumerate() { + let idx = i + j; + let prod = (ai as u128) * (bj as u128) + (c[idx] as u128) + carry; + c[idx] = prod as Word; + carry = prod >> Word::BITS; + } + let mut k = i + b.len(); + while carry != 0 { + let sum = (c[k] as u128) + carry; + c[k] = sum as Word; + carry = sum >> Word::BITS; + k += 1; + } + } + c + } + + fn run_ntt_vs_schoolbook(la: usize, lb: usize) { + let a: Vec = (0..la) + .map(|i| (i as Word + 1).wrapping_mul(0x9E3779B97F4A7C15u64 as Word)) + .collect(); + let b: Vec = (0..lb) + .map(|i| (i as Word + 1).wrapping_mul(0xC6A4A7935BD1E995u64 as Word)) + .collect(); + let expected = schoolbook_mul(&a, &b); + + let mut c = vec![0u64 as Word; a.len() + b.len()]; + let layout = memory_requirement_up_to(c.len(), b.len()); + let mut alloc = crate::memory::MemoryAllocation::new(layout); + let mut memory = alloc.memory(); + let carry = add_signed_mul_conv(&mut c, Positive, &a, &b, &mut memory); + assert_eq!(carry, 0, "carry should be 0"); + assert_eq!(&c[..], &expected[..], "NTT mismatch: la={la}, lb={lb}"); + } + + #[test] + fn test_ntt_vs_schoolbook_equal() { + for &len in &[20, 30, 50, 64, 100, 128] { + run_ntt_vs_schoolbook(len, len); + } + } + + #[test] + fn test_ntt_vs_schoolbook_unequal() { + for &(la, lb) in &[(30, 20), (50, 30), (100, 50), (128, 64), (100, 20)] { + run_ntt_vs_schoolbook(la, lb); + } + } + + #[test] + fn test_ntt_vs_schoolbook_asymmetric() { + for &(la, lb) in &[(200, 30), (150, 20)] { + run_ntt_vs_schoolbook(la, lb); + } + } + + #[test] + fn test_ntt_all_ones() { + for &len in &[20, 50] { + let a = vec![Word::MAX; len]; + let b = vec![Word::MAX; len]; + let expected = schoolbook_mul(&a, &b); + + let mut c = vec![0u64 as Word; a.len() + b.len()]; + let layout = memory_requirement_up_to(c.len(), b.len()); + let mut alloc = crate::memory::MemoryAllocation::new(layout); + let mut memory = alloc.memory(); + add_signed_mul_conv(&mut c, Positive, &a, &b, &mut memory); + assert_eq!(&c[..], &expected[..], "all-ones mismatch len={len}"); + } + } + + #[test] + fn test_ntt_high_low_zero_limbs() { + let mut a = vec![0u64 as Word; 80]; + let mut b = vec![0u64 as Word; 80]; + for i in 20..60 { + a[i] = (i as Word + 1).wrapping_mul(0xDEADBEEF); + b[i] = (i as Word + 1).wrapping_mul(0xCAFEBABE); + } + let expected = schoolbook_mul(&a, &b); + + let mut c = vec![0u64 as Word; a.len() + b.len()]; + let layout = memory_requirement_up_to(c.len(), b.len()); + let mut alloc = crate::memory::MemoryAllocation::new(layout); + let mut memory = alloc.memory(); + add_signed_mul_conv(&mut c, Positive, &a, &b, &mut memory); + assert_eq!(&c[..], &expected[..], "sparse operand mismatch"); + } +} diff --git a/integer/src/mul/ntt/pack.rs b/integer/src/mul/ntt/pack.rs new file mode 100644 index 00000000..05111224 --- /dev/null +++ b/integer/src/mul/ntt/pack.rs @@ -0,0 +1,190 @@ +//! Bit-level packing / unpacking of `b`-bit coefficients. + +use crate::arch::ntt::Lane; +use crate::arch::word::Word; + +/// Pack a big integer (given as `&[Word]`, little-endian) into `out`, +/// producing `n` coefficients of `b_pack` bits each, zero-padded. +/// +/// Each coefficient `c_i` satisfies `0 ≤ c_i < 2^{b_pack}`. +/// Panics if `out.len() < n`. +pub fn pack(out: &mut [Lane], words: &[Word], b_pack: u32, n: usize) { + assert!(out.len() >= n); + + // Fast path: one coefficient per word, no bit shifting needed. + if b_pack == Word::BITS { + let len = words.len().min(n); + // SAFETY: NTT path requires Word and Lane have the same size. + #[allow(clippy::unnecessary_cast)] + let words_lane = unsafe { &*(words as *const [Word] as *const [Lane]) }; + out[..len].copy_from_slice(&words_lane[..len]); + out[len..n].fill(0); + return; + } + + let word_bits = Word::BITS; + let mask: Word = if b_pack < word_bits { + (1 << b_pack) - 1 + } else { + Word::MAX + }; + let mut word_idx = 0usize; + let mut bit_offset = 0u32; + + for coeff in out.iter_mut().take(n) { + if word_idx >= words.len() { + *coeff = 0; + continue; + } + + if bit_offset + b_pack <= word_bits { + *coeff = (words[word_idx] >> bit_offset) & mask; + bit_offset += b_pack; + if bit_offset == word_bits { + bit_offset = 0; + word_idx += 1; + } + } else { + let bits_first = word_bits - bit_offset; + let bits_second = b_pack - bits_first; + let mut val = (words[word_idx] >> bit_offset) & ((1 << bits_first) - 1); + word_idx += 1; + if word_idx < words.len() { + val |= (words[word_idx] & ((1 << bits_second) - 1)) << bits_first; + } + *coeff = val; + bit_offset = bits_second; + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + #[cfg(not(feature = "std"))] + use alloc::vec; + #[cfg(not(feature = "std"))] + use alloc::vec::Vec; + + /// Accumulate CRT-recovered convolution coefficients into the output limb + /// array with carry propagation. + /// + /// Each coefficient `c_k` contributes `c_k << (k * b_pack)` bits to the + /// output. `output` must have capacity for `c.len()` coefficients plus any + /// carry overflow. + fn unpack_accumulate(output: &mut [Word], coeffs: &[Lane], b_pack: u32, output_len: usize) { + let word_bits = Word::BITS; + + for (k, &coeff) in coeffs.iter().enumerate().take(output_len) { + if coeff == 0 { + continue; + } + let shift_bits = (k as u32).wrapping_mul(b_pack); + let word_idx = (shift_bits / word_bits) as usize; + let bit_shift = shift_bits % word_bits; + + let lo = coeff; + let mut carry: Word; + let mut idx = word_idx; + + if bit_shift == 0 { + let (sum, c) = output.get(idx).copied().unwrap_or(0).overflowing_add(lo); + carry = Word::from(c); + if idx < output.len() { + output[idx] = sum; + } + idx += 1; + } else { + let val = (lo as u128) << bit_shift; + let lo_part = val as Word; + let hi_part = (val >> word_bits) as Word; + + let (sum, c1) = output + .get(idx) + .copied() + .unwrap_or(0) + .overflowing_add(lo_part); + carry = Word::from(c1); + if idx < output.len() { + output[idx] = sum; + } + idx += 1; + + let (sum2, c2) = output + .get(idx) + .copied() + .unwrap_or(0) + .overflowing_add(hi_part + carry); + carry = Word::from(c2); + if idx < output.len() { + output[idx] = sum2; + } + idx += 1; + } + + // Propagate remaining carry + while carry != 0 && idx < output.len() { + let (sum, c) = output[idx].overflowing_add(carry); + output[idx] = sum; + carry = Word::from(c); + idx += 1; + } + } + } + + #[test] + fn test_pack_unpack_roundtrip() { + let b_pack = 16u32; + let test_words: Vec = vec![0xDEADBEEF, 0x12345678]; + let coeffs_per_word = (Word::BITS / b_pack) as usize; + let n = test_words.len() * coeffs_per_word; + + let mut packed: Vec = vec![0; n]; + pack(&mut packed, &test_words, b_pack, n); + + let output_len = test_words.len() + 1; + let mut output: Vec = vec![0; output_len]; + unpack_accumulate(&mut output, &packed, b_pack, n); + assert_eq!(&output[..test_words.len()], &test_words[..]); + } + + #[test] + fn test_pack_zero_pads() { + let words: Vec = vec![0xFFFF]; + let n = 32; + let mut packed: Vec = vec![0; n]; + pack(&mut packed, &words, 16, n); + assert_eq!(packed[0], 0xFFFF); + for &c in packed.iter().skip(1) { + assert_eq!(c, 0); + } + } + + #[test] + fn test_pack_empty_input() { + let mut packed: Vec = vec![0; 8]; + pack(&mut packed, &[], 16, 8); + assert_eq!(packed, vec![0; 8]); + } + + #[test] + fn test_unpack_single_coeff() { + let mut output: Vec = vec![0; 2]; + unpack_accumulate(&mut output, &[0xABCD], 16, 1); + assert_eq!(output[0], 0xABCD); + assert_eq!(output[1], 0); + } + + #[test] + fn test_unpack_carry_propagation() { + // Coefficient at k = Word::BITS/16 shifts by exactly Word::BITS bits = 1 word. + let k = (Word::BITS / 16) as usize; + let mut coeffs: Vec = vec![0; k + 1]; + coeffs[k] = 1; + let mut output: Vec = vec![0; 3]; + unpack_accumulate(&mut output, &coeffs, 16, k + 1); + assert_eq!(output[0], 0); + assert_eq!(output[1], 1); + assert_eq!(output[2], 0); + } +} diff --git a/integer/src/mul/ntt/transform.rs b/integer/src/mul/ntt/transform.rs new file mode 100644 index 00000000..4b8866d6 --- /dev/null +++ b/integer/src/mul/ntt/transform.rs @@ -0,0 +1,226 @@ +//! Iterative in-place radix-2 NTT over Proth primes `K * 2^N + 1`. +//! +//! All functions are generic over `R: Reducer` so each prime's +//! reducer is monomorphized at the call site. + +use crate::arch::ntt::{Lane, MAX_LOG_N}; +use num_modular::Reducer; + +// ---- public API ---- + +/// Fill `out[0..n/2]` with twiddle factors `omega_n^k` in Montgomery form. +/// +/// Panics if `out.len() < n / 2`. +pub fn precompute_twiddles>( + out: &mut [Lane], + n: usize, + omega_max: Lane, + inverse: bool, + r: &R, +) { + assert!(out.len() >= n / 2); + let shift = MAX_LOG_N - n.trailing_zeros(); + let omega_max_mont = r.transform(omega_max); + let omega_n_mont = r.pow(omega_max_mont, &((1u64 as Lane) << shift)); + + let base_mont = if inverse { + r.inv(omega_n_mont).expect("omega_n not invertible") + } else { + omega_n_mont + }; + + out[0] = r.transform(1); + for k in 1..(n / 2) { + out[k] = r.mul(&out[k - 1], &base_mont); + } +} + +/// Bit-reverse `a` in place. Length must be a power of two. +pub fn bit_reverse(a: &mut [Lane]) { + let n = a.len(); + assert!(n.is_power_of_two()); + let log_n = n.trailing_zeros(); + for i in 0..n { + let j = i.reverse_bits() >> (usize::BITS - log_n); + if i < j { + a.swap(i, j); + } + } +} + +/// Forward NTT in place (decimation-in-time, radix-2). +pub fn forward>(a: &mut [Lane], twiddles: &[Lane], r: &R) { + ntt_core(a, twiddles, r); +} + +/// Inverse NTT in place. +/// +/// Computed as `bit_reverse → forward(ω⁻¹) → scale`, producing output +/// in **natural order**. +/// +/// `twiddles` must have been precomputed with `inverse = true`. +pub fn inverse>(a: &mut [Lane], twiddles: &[Lane], r: &R) { + let n = a.len(); + bit_reverse(a); + ntt_core(a, twiddles, r); + let n_mont = r.transform(n as Lane); + let n_inv_mont = r.inv(n_mont).expect("n not invertible mod p"); + for x in a.iter_mut() { + *x = r.mul(x, &n_inv_mont); + } +} + +/// In-place radix-2 DIT NTT (Cooley–Tukey). +fn ntt_core>(a: &mut [Lane], twiddles: &[Lane], r: &R) { + let n = a.len(); + debug_assert!(n.is_power_of_two() && twiddles.len() == n / 2); + + let mut sub_len = 2usize; + while sub_len <= n { + let half = sub_len / 2; + let step = n / sub_len; + + for i in (0..n).step_by(sub_len) { + for j in 0..half { + let u = a[i + j]; + let v = r.mul(&a[i + j + half], &twiddles[j * step]); + a[i + j] = r.add(&u, &v); + a[i + j + half] = r.sub(&u, &v); + } + } + + sub_len *= 2; + } +} + +/// Pointwise multiply of two transformed vectors in place. +pub fn pointwise_mul>(a_hat: &mut [Lane], b_hat: &[Lane], r: &R) { + assert_eq!(a_hat.len(), b_hat.len()); + for (a, &b_val) in a_hat.iter_mut().zip(b_hat.iter()) { + *a = r.mul(a, &b_val); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::arch::ntt::{K, MODULI, OMEGA_MAX, P0, P1, P2}; + #[cfg(not(feature = "std"))] + use alloc::vec; + #[cfg(not(feature = "std"))] + use alloc::vec::Vec; + + fn assert_all_eq(a: &[Lane], b_val: &[Lane], context: &str) { + assert_eq!(a.len(), b_val.len(), "{context}: length mismatch"); + for (i, (x, y)) in a.iter().zip(b_val.iter()).enumerate() { + assert_eq!(x, y, "{context}: mismatch at index {i}: {x} != {y}"); + } + } + + macro_rules! for_each_prime { + ($r:ident, $p:ident, $omega:ident, $body:block) => { + for idx in 0..K { + let $p = MODULI[idx]; + let $omega = OMEGA_MAX[idx]; + match idx { + 0 => { + fn go>($r: &R, $p: Lane, $omega: Lane) $body + go::(&P0, $p, $omega); + } + 1 => { + fn go>($r: &R, $p: Lane, $omega: Lane) $body + go::(&P1, $p, $omega); + } + 2 => { + fn go>($r: &R, $p: Lane, $omega: Lane) $body + go::(&P2, $p, $omega); + } + _ => unreachable!(), + } + } + }; + } + + #[test] + fn test_forward_inverse_roundtrip() { + for_each_prime!(r, p, omega, { + for &n in &[2, 4, 8, 16, 32, 64, 128, 256, 512] { + let mut fwd_twiddles = alloc::vec![0u64 as Lane; n / 2]; + let mut inv_twiddles = alloc::vec![0u64 as Lane; n / 2]; + precompute_twiddles(&mut fwd_twiddles, n, omega, false, r); + precompute_twiddles(&mut inv_twiddles, n, omega, true, r); + + let mut a: Vec = (0..n) + .map(|i| ((i as Lane + 1).wrapping_mul(123456789)) % p) + .collect(); + for val in a.iter_mut() { + *val = r.transform(*val); + } + let orig = a.clone(); + + bit_reverse(&mut a); + forward(&mut a, &fwd_twiddles, r); + inverse(&mut a, &inv_twiddles, r); + + assert_all_eq(&a, &orig, "roundtrip failed for n={n}"); + } + }); + } + + #[test] + fn test_convolution_via_ntt() { + for_each_prime!(r, p, omega, { + for len_a in [1, 2, 3, 5] { + for len_b in [1, 2, 3, 5] { + let conv_len: usize = len_a + len_b - 1; + let n = conv_len.next_power_of_two().max(2); + + let a: Vec = (0..len_a).map(|i| ((i + 1) as Lane * 12345) % p).collect(); + let b_vec: Vec = + (0..len_b).map(|i| ((i + 1) as Lane * 67890) % p).collect(); + + let mut expected = vec![0u64 as Lane; conv_len]; + for (i, &ai) in a.iter().enumerate() { + for (j, &bj) in b_vec.iter().enumerate() { + let prod = (ai as u128 * bj as u128 % p as u128) as Lane; + expected[i + j] = r.add(&expected[i + j], &prod); + } + } + + let mut fwd_twiddles = alloc::vec![0u64 as Lane; n / 2]; + let mut inv_twiddles = alloc::vec![0u64 as Lane; n / 2]; + precompute_twiddles(&mut fwd_twiddles, n, omega, false, r); + precompute_twiddles(&mut inv_twiddles, n, omega, true, r); + + let mut a_pad = vec![0u64 as Lane; n]; + let mut b_pad = vec![0u64 as Lane; n]; + for i in 0..len_a { + a_pad[i] = r.transform(a[i]); + } + for i in 0..len_b { + b_pad[i] = r.transform(b_vec[i]); + } + + bit_reverse(&mut a_pad); + bit_reverse(&mut b_pad); + forward(&mut a_pad, &fwd_twiddles, r); + forward(&mut b_pad, &fwd_twiddles, r); + pointwise_mul(&mut a_pad, &b_pad, r); + inverse(&mut a_pad, &inv_twiddles, r); + for val in a_pad[..conv_len].iter_mut() { + *val = r.residue(*val); + } + + assert_all_eq(&a_pad[..conv_len], &expected, "convolution mismatch"); + } + } + }); + } + + #[test] + fn test_bit_reverse() { + let mut a: Vec = (0..8).map(|i| i as Lane).collect(); + bit_reverse(&mut a); + assert_eq!(a, vec![0, 4, 2, 6, 1, 5, 3, 7]); + } +}