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}