From ff6aa3b39370b5e8827b6ea7b1884b5c2ecf077d Mon Sep 17 00:00:00 2001 From: Dev Ojha Date: Wed, 2 Sep 2026 18:50:03 +0200 Subject: [PATCH 1/2] Cache per-level FFT twiddle tables --- crates/halo2_proofs/src/poly/domain.rs | 178 ++++++++++++++++++++++--- 1 file changed, 163 insertions(+), 15 deletions(-) diff --git a/crates/halo2_proofs/src/poly/domain.rs b/crates/halo2_proofs/src/poly/domain.rs index adf27195..a3d247f6 100644 --- a/crates/halo2_proofs/src/poly/domain.rs +++ b/crates/halo2_proofs/src/poly/domain.rs @@ -35,7 +35,7 @@ type PolynomialTransformBatch = ( /// The fields are private so caches can only be built by /// [`EvaluationDomain::proving_key_twiddles`] or cloned from one it built. /// -/// Clones share both allocations. This matters because [`ProvingKey`] is +/// Clones share all retained allocations. This matters because [`ProvingKey`] is /// cloneable and these tables are intended to be retained, not rebuilt. /// Exact equivalence with the uncached transforms across the supported domain /// shapes is covered by @@ -46,6 +46,8 @@ type PolynomialTransformBatch = ( pub(crate) struct ProvingKeyTwiddles { base_inverse: Arc<[F]>, extended_forward: Arc<[F]>, + base_inverse_tables: Arc<[Vec]>, + extended_forward_tables: Arc<[Vec]>, } impl fmt::Debug for ProvingKeyTwiddles { @@ -314,9 +316,16 @@ impl> EvaluationDomain { /// Builds the FFT twiddles retained by a proving key. pub(crate) fn proving_key_twiddles(&self) -> ProvingKeyTwiddles { + let base_inverse = twiddle_table(self.omega_inv, 1 << self.k); + let extended_forward = twiddle_table(self.extended_omega, self.extended_len()); + let base_inverse_tables = butterfly_twiddle_tables(&base_inverse, 1 << self.k); + let extended_forward_tables = + butterfly_twiddle_tables(&extended_forward, self.extended_len()); ProvingKeyTwiddles { - base_inverse: Arc::from(twiddle_table(self.omega_inv, 1 << self.k)), - extended_forward: Arc::from(twiddle_table(self.extended_omega, self.extended_len())), + base_inverse: Arc::from(base_inverse), + extended_forward: Arc::from(extended_forward), + base_inverse_tables: Arc::from(base_inverse_tables), + extended_forward_tables: Arc::from(extended_forward_tables), } } @@ -336,6 +345,8 @@ impl> EvaluationDomain { 1, 1, &twiddles.base_inverse, + &twiddles.base_inverse_tables, + 0, parallel_depth(), ); normalize_inverse_fft(&mut polynomial.values, self.k, self.ifft_divisor, true); @@ -362,6 +373,7 @@ impl> EvaluationDomain { self.k, self.extended_k, &twiddles.extended_forward, + &twiddles.extended_forward_tables, parallel_depth(), ); @@ -415,6 +427,8 @@ impl> EvaluationDomain { 1, 1, &twiddles.base_inverse, + &twiddles.base_inverse_tables, + 0, INNER_PARALLEL_DEPTH, ); normalize_inverse_fft(&mut values, self.k, self.ifft_divisor, false); @@ -430,6 +444,7 @@ impl> EvaluationDomain { self.k, self.extended_k, &twiddles.extended_forward, + &twiddles.extended_forward_tables, INNER_PARALLEL_DEPTH, ); let extended = Polynomial { @@ -532,6 +547,8 @@ impl> EvaluationDomain { 1, 1, &twiddles.extended_forward, + &twiddles.extended_forward_tables, + 0, parallel_depth(), ); @@ -591,6 +608,8 @@ impl> EvaluationDomain { 1, 1, &twiddles.extended_forward, + &twiddles.extended_forward_tables, + 0, parallel_depth(), ); polynomial.values[1..].reverse(); @@ -701,11 +720,13 @@ impl> EvaluationDomain { } let twiddles = twiddle_table(omega, extended_n); + let tables = butterfly_twiddle_tables(&twiddles, extended_n); Self::fft_zero_padded_with_twiddles( coefficients, log_n, extended_log_n, &twiddles, + &tables, parallel_depth(), ); } @@ -715,6 +736,7 @@ impl> EvaluationDomain { log_n: u32, extended_log_n: u32, twiddles: &[F], + tables: &[Vec], parallel_depth: u32, ) { assert!(log_n <= extended_log_n); @@ -773,7 +795,15 @@ impl> EvaluationDomain { // Each 2 * extension chunk contains the first retained stage. Continue // with the other log_n - 1 butterfly stages. A zero parallel depth // still traverses recursively to retain cache locality. - recursive_butterfly_after_prefix(&mut values, 2 * extension, 1, twiddles, parallel_depth); + recursive_butterfly_after_prefix( + &mut values, + 2 * extension, + 1, + twiddles, + tables, + 0, + parallel_depth, + ); *coefficients = values; } @@ -983,11 +1013,30 @@ fn bitreverse(value: usize, bits: u32) -> usize { } } +/// Node half-lengths at or below this size use the strided scalar combine +/// loop; larger nodes use a contiguous per-level twiddle table. +const CACHED_TWIDDLE_MIN_HALF: usize = 32; + +/// Builds contiguous per-level twiddle tables for the butterfly combine step. +fn butterfly_twiddle_tables(twiddles: &[F], len: usize) -> Vec> { + let mut tables = Vec::new(); + let mut half = len / 2; + let mut chunk = 1; + while half > CACHED_TWIDDLE_MIN_HALF { + tables.push((1..half).map(|index| twiddles[index * chunk]).collect()); + half /= 2; + chunk *= 2; + } + tables +} + fn recursive_butterfly_after_prefix( values: &mut [F], completed_chunk_len: usize, twiddle_chunk: usize, twiddles: &[F], + tables: &[Vec], + level: usize, parallel_depth: u32, ) { let len = values.len(); @@ -1005,6 +1054,8 @@ fn recursive_butterfly_after_prefix( completed_chunk_len, twiddle_chunk * 2, twiddles, + tables, + level + 1, parallel_depth - 1, ) }, @@ -1014,6 +1065,8 @@ fn recursive_butterfly_after_prefix( completed_chunk_len, twiddle_chunk * 2, twiddles, + tables, + level + 1, parallel_depth - 1, ) }, @@ -1025,11 +1078,13 @@ fn recursive_butterfly_after_prefix( completed_chunk_len, twiddle_chunk * 2, twiddles, + tables, + level + 1, ); } } - butterfly_chunk(left, right, twiddle_chunk, twiddles); + butterfly_chunk(left, right, twiddle_chunk, twiddles, tables, level); } /// Recursively processes two equal-sized field FFT chunks together. The FFT @@ -1042,6 +1097,8 @@ fn recursive_butterfly_pair_after_prefix( completed_chunk_len: usize, twiddle_chunk: usize, twiddles: &[F], + tables: &[Vec], + level: usize, ) { debug_assert_eq!(first.len(), second.len()); let len = first.len(); @@ -1061,6 +1118,8 @@ fn recursive_butterfly_pair_after_prefix( completed_chunk_len, twiddle_chunk * 2, twiddles, + tables, + level + 1, ); recursive_butterfly_pair_after_prefix( second_left, @@ -1068,6 +1127,8 @@ fn recursive_butterfly_pair_after_prefix( completed_chunk_len, twiddle_chunk * 2, twiddles, + tables, + level + 1, ); } @@ -1078,6 +1139,8 @@ fn recursive_butterfly_pair_after_prefix( second_right, twiddle_chunk, twiddles, + tables, + level, ); } @@ -1087,6 +1150,8 @@ fn butterfly_chunk( right: &mut [F], twiddle_chunk: usize, twiddles: &[F], + tables: &[Vec], + level: usize, ) { // Handle the unity twiddle without a field multiplication. let (first_left, left) = left.split_at_mut(1); @@ -1096,16 +1161,30 @@ fn butterfly_chunk( first_left[0] += &t; first_right[0] -= &t; - for (index, (left, right)) in left.iter_mut().zip(right.iter_mut()).enumerate() { - let mut t = *right; - t *= &twiddles[(index + 1) * twiddle_chunk]; - *right = *left; - *left += &t; - *right -= &t; + if let Some(table) = tables.get(level) { + debug_assert_eq!(table.len(), right.len()); + for (right, twiddle) in right.iter_mut().zip(table) { + *right *= twiddle; + } + for (left, right) in left.iter_mut().zip(right.iter_mut()) { + let t = *right; + *right = *left; + *left += &t; + *right -= &t; + } + } else { + for (index, (left, right)) in left.iter_mut().zip(right.iter_mut()).enumerate() { + let mut t = *right; + t *= &twiddles[(index + 1) * twiddle_chunk]; + *right = *left; + *left += &t; + *right -= &t; + } } } #[inline(always)] +#[allow(clippy::too_many_arguments)] fn butterfly_chunk_pair( first_left: &mut [F], first_right: &mut [F], @@ -1113,6 +1192,8 @@ fn butterfly_chunk_pair( second_right: &mut [F], twiddle_chunk: usize, twiddles: &[F], + tables: &[Vec], + level: usize, ) { debug_assert_eq!(first_left.len(), first_right.len()); debug_assert_eq!(first_left.len(), second_left.len()); @@ -1132,6 +1213,31 @@ fn butterfly_chunk_pair( first_b[0] -= &first_t; second_b[0] -= &second_t; + if let Some(table) = tables.get(level) { + debug_assert_eq!(table.len(), first_right.len()); + for (right, twiddle) in first_right.iter_mut().zip(table) { + *right *= twiddle; + } + for (right, twiddle) in second_right.iter_mut().zip(table) { + *right *= twiddle; + } + for ((first_left, first_right), (second_left, second_right)) in first_left + .iter_mut() + .zip(first_right.iter_mut()) + .zip(second_left.iter_mut().zip(second_right.iter_mut())) + { + let first_t = *first_right; + let second_t = *second_right; + *first_right = *first_left; + *second_right = *second_left; + *first_left += &first_t; + *second_left += &second_t; + *first_right -= &first_t; + *second_right -= &second_t; + } + return; + } + for (index, ((first_left, first_right), (second_left, second_right))) in first_left .iter_mut() .zip(first_right.iter_mut()) @@ -1581,15 +1687,49 @@ fn test_sparse_quotient_division_matches_pointwise_on_basis_vectors() { fn test_orchard_proving_key_twiddle_cache_size_and_sharing() { use crate::pasta::pallas::Scalar; - let domain = EvaluationDomain::::new(9, 11); + const ORCHARD_DEGREE: u32 = 9; + const ORCHARD_K: u32 = 11; + const EXPECTED_FLAT_PAYLOAD_BYTES: usize = 288 * 1024; + const EXPECTED_LEVEL_PAYLOAD_BYTES: usize = 585_312; + #[cfg(target_pointer_width = "64")] + const EXPECTED_LEVEL_RETAINED_BYTES: usize = 585_656; + + let domain = EvaluationDomain::::new(ORCHARD_DEGREE, ORCHARD_K); let twiddles = domain.proving_key_twiddles(); assert_eq!(std::mem::size_of::(), 32); - assert_eq!(twiddles.base_inverse.len(), 1 << 10); - assert_eq!(twiddles.extended_forward.len(), 1 << 13); + assert_eq!(twiddles.base_inverse.len(), 1 << (ORCHARD_K - 1)); + assert_eq!( + twiddles.extended_forward.len(), + 1 << (domain.extended_k - 1) + ); assert_eq!( (twiddles.base_inverse.len() + twiddles.extended_forward.len()) * std::mem::size_of::(), - 288 * 1024 + EXPECTED_FLAT_PAYLOAD_BYTES + ); + + let level_tables = twiddles + .base_inverse_tables + .iter() + .chain(twiddles.extended_forward_tables.iter()) + .collect::>(); + assert!( + level_tables + .iter() + .all(|table| table.len() == table.capacity()) + ); + let level_payload_bytes = + level_tables.iter().map(|table| table.len()).sum::() * std::mem::size_of::(); + assert_eq!(level_payload_bytes, EXPECTED_LEVEL_PAYLOAD_BYTES); + + // This stable accounting excludes allocator metadata and the two `Arc` + // reference-count headers, whose sizes are implementation details. + #[cfg(target_pointer_width = "64")] + assert_eq!( + level_payload_bytes + + level_tables.len() * std::mem::size_of::>() + + 2 * std::mem::size_of::]>>(), + EXPECTED_LEVEL_RETAINED_BYTES ); let cloned = twiddles.clone(); @@ -1598,6 +1738,14 @@ fn test_orchard_proving_key_twiddle_cache_size_and_sharing() { &twiddles.extended_forward, &cloned.extended_forward )); + assert!(Arc::ptr_eq( + &twiddles.base_inverse_tables, + &cloned.base_inverse_tables + )); + assert!(Arc::ptr_eq( + &twiddles.extended_forward_tables, + &cloned.extended_forward_tables + )); } #[test] From 4b72e484433da528d4672a3e240a1fe8afe0c290 Mon Sep 17 00:00:00 2001 From: Dev Ojha Date: Wed, 2 Sep 2026 23:37:20 +0200 Subject: [PATCH 2/2] Add changelog for FFT twiddle cache Co-authored-by: burbach7 --- docs/changelog/unreleased/324.md | 12 ++++++++++++ 1 file changed, 12 insertions(+) create mode 100644 docs/changelog/unreleased/324.md diff --git a/docs/changelog/unreleased/324.md b/docs/changelog/unreleased/324.md new file mode 100644 index 00000000..46b0dc02 --- /dev/null +++ b/docs/changelog/unreleased/324.md @@ -0,0 +1,12 @@ +## zakura-halo2-proofs + +### Changed + +- Cached contiguous per-level FFT twiddles in proving keys. A + component-isolated benchmark reduced first-proof latency for a prepared + Orchard k=11 four-action proof by 0.565% on 10-worker Apple arm64; one action + and the six-worker x86_64 Linux gates were neutral. The cache retains an + additional 585,656 bytes per independently generated Orchard proving key on + 64-bit targets, excluding allocator and reference-count metadata, and is + shared by clones. Proof format and verification are unchanged + ([#324](https://github.com/zakura-core/common/pull/324)).