Skip to main content

binius_prover/protocols/bitand/
ntt_lookup.rs

1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4//! # NTT Lookup Table Module
5//!
6//! This module provides a precomputed lookup table implementation for fast Number Theoretic
7//! Transform (NTT) operations on 64-bit binary field elements. The implementation is specifically
8//! optimized for the Binius64 protocol's constraint system.
9//!
10//! ## Overview
11//!
12//! The NTT lookup table achieves significant performance improvements by precomputing all possible
13//! NTT evaluations for 8-bit input chunks. This allows the full 64-bit NTT to be computed by:
14//!
15//! 1. Splitting the 64 input bits into eight 8-bit chunks
16//! 2. Looking up precomputed NTT values for each chunk
17//! 3. Adding the results together (exploiting the linearity of the NTT)
18//!
19//! ## Algorithm
20//!
21//! The transformation is really a *low-degree extension* (LDE) over a binary subspace, not a plain
22//! NTT. It maps the 64 input bits — viewed as evaluations of a polynomial over the input domain
23//! (the lower half of the subspace) — to 64 evaluations of that same polynomial over the output
24//! domain (the upper half, a coset shift of the input domain). The LDE is the composition of an
25//! inverse NTT over the input domain (recovering the polynomial's coefficients) with a forward NTT
26//! over the coset-shifted output domain; see `LowDegreeExtension::transform`.
27//!
28//! Because the LDE is linear, the LDE of a 64-bit input is the sum of the LDEs of its eight bytes:
29//!
30//! - **Input**: 64 1-bit coefficients (evaluations over the input domain)
31//! - **Output**: 64 field elements (evaluations over the output domain)
32//! - **Optimization**: precompute all 256 LDE images for each 8-bit position, so the full LDE is 8
33//!   table lookups plus 7 packed additions.
34//!
35//! ## Compressed storage
36//!
37//! Storing an independent 256-entry table for each of the eight byte positions would take
38//! `8 * 256 * 64` field elements. We instead store tables for only two byte positions and
39//! reconstruct the rest by permuting the packed evaluations.
40//!
41//! This exploits a *translation-invariance* property of the LDE matrix observed in the Flock paper
42//! (<https://github.com/succinctlabs/flock/blob/main/paper/flock-paper.pdf>, Section 4.2): because
43//! the output domain is a coset shift of the input domain over a binary subspace, the LDE images of
44//! the unit inputs at different byte positions are fixed permutations of one another. The LDE image
45//! for byte position `b` is therefore recovered from the stored table for parity `b % 2` by
46//! permuting its packed evaluations according to `b / 2` (see [`NTTLookup::ntt`]). This shrinks the
47//! table footprint by 4x — to `2 * 256 * 64` field elements — adding only a cheap in-register
48//! permute to each lookup.
49
50use std::{array, marker::PhantomData};
51
52use binius_core::Word;
53use binius_field::{
54	BinaryField, BinaryField1b as B1, Divisible, PackedField, PackedRijndael64x8b,
55	Rijndael8b as B8, UnderlierView, arch::M128, util::expand_subset_sums_array,
56};
57use binius_math::{
58	BinarySubspace, FieldBuffer,
59	ntt::{AdditiveNTT, NeighborsLastReference, domain_context::GenericOnTheFly},
60};
61use binius_verifier::protocols::bitand::{ROWS_PER_HYPERCUBE_VERTEX, SKIPPED_VARS};
62
63/// A precomputed lookup table for fast LDE operations on 64-bit binary field elements.
64///
65/// This structure stores precomputed LDE evaluations for all possible 8-bit input combinations,
66/// enabling fast computation of the full 64-bit LDE through table lookups and additions. See the
67/// module-level documentation for the LDE definition and the Flock compression it uses.
68///
69/// ## Structure
70///
71/// The internal data structure is a boxed array `Box<[[PackedRijndael64x8b; 256]; 2]>` where:
72/// - **First dimension**: the byte-position parity `b % 2`. Only these two tables are stored; the
73///   remaining six byte positions are recovered by permutation (see the module-level "Compressed
74///   storage" notes).
75/// - **Second dimension**: the 8-bit value (0-255) for that byte.
76///
77/// Each entry holds the `ROWS_PER_HYPERCUBE_VERTEX` LDE evaluations of that byte's coefficients,
78/// packed into a single [`PackedRijndael64x8b`].
79#[derive(Debug, Clone)]
80pub struct NTTLookup(Box<[[PackedRijndael64x8b; 256]; 2]>);
81
82impl NTTLookup {
83	/// Creates a new NTT lookup table by precomputing all possible NTT evaluations
84	/// for 8-bit input chunks across all byte positions in a 64-bit word.
85	///
86	/// ## Parameters
87	///
88	/// - `subspace`: Binary subspace of dimension `SKIPPED_VARS + 1`. Its lower half defines the
89	///   NTT input domain and its upper half the output domain at which evaluations are
90	///   precomputed.
91	///
92	/// ## Constraints
93	///
94	/// - Subspace dimension must equal `SKIPPED_VARS + 1`
95	///
96	/// # Panics
97	///
98	/// Panics if `subspace.dim() != SKIPPED_VARS + 1`.
99	pub fn new(subspace: &BinarySubspace<B8>) -> Self {
100		assert_eq!(subspace.dim(), SKIPPED_VARS + 1);
101
102		let lde = LowDegreeExtension::<PackedRijndael64x8b>::new(subspace);
103		let lde_mat = array::from_fn::<_, 2, _>(|b| {
104			array::from_fn::<_, 8, _>(|i| {
105				let output = lde.transform(1 << (8 * b + i));
106				assert_eq!(output.log_len(), SKIPPED_VARS + 1);
107				// Pull out the second element, corresponding to the output domain
108				output.as_ref()[1]
109			})
110		});
111
112		let lookup = lde_mat.map(expand_subset_sums_array::<_, 8, 256>);
113		NTTLookup(Box::new(lookup))
114	}
115
116	/// Computes the LDE of 64 1-bit coefficients using the precomputed lookup tables.
117	///
118	/// The 64-bit `input` is split into eight bytes B₀, B₁, ..., B₇. By linearity the LDE of the
119	/// input is the sum of the per-byte LDEs: `LDE(input) = LDE(B₀) + LDE(B₁) + ... + LDE(B₇)`.
120	///
121	/// Only two byte-position tables are stored, so the LDE of byte position `b` is reconstructed
122	/// from the table for parity `b % 2` by permuting its packed evaluations according to `b / 2`.
123	/// This permutation is the translation of the output domain exploited by the Flock compression
124	/// (see the module-level "Compressed storage" notes); in the packed representation it is a
125	/// permutation of the four 128-bit lanes, `lane[i] <- lane[i ^ (b / 2)]`. The loop is unrolled
126	/// over `b`, so the parity select and lane index are compile-time constants.
127	///
128	/// Used directly only in tests; `univariate_round_message_extension_domain` accesses the tables
129	/// inline to compute three LDE evaluations at once, which is more efficient.
130	///
131	/// ## Returns
132	///
133	/// A [`PackedRijndael64x8b`] holding the `ROWS_PER_HYPERCUBE_VERTEX` LDE evaluations over the
134	/// output domain.
135	#[inline]
136	#[must_use]
137	pub fn ntt(&self, input: Word) -> PackedRijndael64x8b {
138		let input_bytes = input.as_u64().to_le_bytes();
139
140		let mut out = PackedRijndael64x8b::default();
141		// This will get unrolled, so indexing arithmetic washes away.
142		for b in 0..8 {
143			let packed = &self.0[b % 2][input_bytes[b] as usize];
144			let bitvec = packed.to_underlier_ref();
145			let dst_bitvec = Divisible::<M128>::from_iter(
146				(0..4).map(|i| Divisible::<M128>::get(bitvec, i ^ (b / 2))),
147			);
148			out += PackedRijndael64x8b::from_underlier(dst_bitvec);
149		}
150		out
151	}
152}
153
154struct LowDegreeExtension<P: PackedField> {
155	interpolation: NeighborsLastReference<GenericOnTheFly<P::Scalar>>,
156	extrapolation: NeighborsLastReference<GenericOnTheFly<P::Scalar>>,
157	_marker: PhantomData<P>,
158}
159
160impl<F, P> LowDegreeExtension<P>
161where
162	F: BinaryField,
163	P: PackedField<Scalar = F>,
164{
165	fn new(subspace: &BinarySubspace<F>) -> Self {
166		assert_eq!(subspace.dim(), SKIPPED_VARS + 1);
167
168		let input_subspace = subspace.reduce_dim(SKIPPED_VARS);
169		let input_domain_context = GenericOnTheFly::generate_from_subspace(&input_subspace);
170		let output_domain_context = GenericOnTheFly::generate_from_subspace(subspace);
171
172		Self {
173			interpolation: NeighborsLastReference {
174				domain_context: input_domain_context,
175			},
176			extrapolation: NeighborsLastReference {
177				domain_context: output_domain_context,
178			},
179			_marker: PhantomData,
180		}
181	}
182
183	/// Computes the low-degree extension of a 64-bit input.
184	///
185	/// The LDE is the composition of two additive NTTs over the binary subspace:
186	///
187	/// 1. an **inverse NTT** over the input domain (the lower half of the subspace), which
188	///    interprets the 64 input bits as evaluations over that domain and recovers the
189	///    polynomial's coefficients;
190	/// 2. a **forward NTT** over the full subspace, which re-evaluates that polynomial over the
191	///    output domain (the upper half) — a coset shift of the input domain.
192	///
193	/// The returned buffer holds both halves; the output-domain evaluations are the upper half.
194	fn transform(&self, input: u64) -> FieldBuffer<P> {
195		let mut values = FieldBuffer::<P>::zeros(SKIPPED_VARS + 1);
196
197		// Inverse NTT the inputs in the first half of the buffer.
198		{
199			let mut values_split = values.split_half_mut();
200			let (mut input_elems, _) = values_split.halves();
201
202			for i in 0..ROWS_PER_HYPERCUBE_VERTEX {
203				input_elems.set(i, F::from(B1::from((input >> i) & 1 == 1)));
204			}
205			self.interpolation.inverse_transform(input_elems, 0, 0);
206		}
207
208		// Forward NTT the zero-padded coefficients.
209		self.extrapolation
210			.forward_transform(values.as_mut_view(), 0, 0);
211
212		values
213	}
214}
215
216#[cfg(test)]
217mod test {
218	use binius_field::Divisible;
219	use binius_math::BinarySubspace;
220	use rand::prelude::*;
221
222	use super::*;
223
224	#[test]
225	fn test_against_ntt() {
226		let subspace = BinarySubspace::with_dim(SKIPPED_VARS + 1);
227		let lde = LowDegreeExtension::<B8>::new(&subspace);
228		let ntt_lookup = NTTLookup::new(&subspace);
229
230		// Repeat for 10 random values
231		let mut rng = StdRng::seed_from_u64(0);
232		for _ in 0..10 {
233			let input = rng.random::<u64>();
234
235			let lde_result = lde.transform(input);
236			let ntt_lookup_result = ntt_lookup.ntt(Word(input));
237			for i in 0..ROWS_PER_HYPERCUBE_VERTEX {
238				let lookup_result = ntt_lookup_result.get(i);
239				assert_eq!(lookup_result, lde_result.get(i + ROWS_PER_HYPERCUBE_VERTEX));
240			}
241		}
242	}
243}