binius_math/ntt/
domain_context.rs1use binius_field::BinaryField;
6
7use super::DomainContext;
8use crate::{BinarySubspace, binary_subspace::BinarySubspaceIterator};
9
10fn generate_evals_from_subspace<F: BinaryField>(subspace: &BinarySubspace<F>) -> Vec<Vec<F>> {
20 let l = subspace.dim();
21 let mut evals = Vec::with_capacity(l);
22
23 evals.push(subspace.basis().to_vec());
25 for i in 1..l {
26 evals.push(Vec::with_capacity(l - i));
28 for k in 1..evals[i - 1].len() {
29 let val = evals[i - 1][k] * (evals[i - 1][k] + evals[i - 1][0]);
33 evals[i].push(val);
34 }
35 }
36
37 for evals_i in evals.iter_mut() {
39 let w_i_b_i_inverse = unsafe { evals_i[0].invert() };
42 for eval_i_j in evals_i.iter_mut() {
43 *eval_i_j *= w_i_b_i_inverse;
44 }
45 }
46
47 evals
48}
49
50#[derive(Debug, Clone)]
57pub struct GenericOnTheFly<F> {
58 evals: Vec<Vec<F>>,
66}
67
68impl<F: BinaryField> GenericOnTheFly<F> {
69 pub fn generate_from_subspace(subspace: &BinarySubspace<F>) -> Self {
73 Self {
74 evals: generate_evals_from_subspace(subspace),
75 }
76 }
77}
78
79impl<F: BinaryField> DomainContext for GenericOnTheFly<F> {
80 type Field = F;
81
82 fn log_domain_size(&self) -> usize {
83 self.evals.len()
84 }
85
86 fn subspace(&self, i: usize) -> BinarySubspace<F> {
87 if i == 0 {
88 return BinarySubspace::with_dim(0);
89 }
90 BinarySubspace::new_unchecked(self.evals[self.log_domain_size() - i].clone())
91 }
92
93 fn twiddle(&self, layer: usize, block: usize) -> F {
94 let v = &self.evals[self.log_domain_size() - layer - 1];
95 BinarySubspace::new_unchecked(&v[1..]).get(block)
96 }
97}
98
99#[derive(Debug)]
101pub struct GenericPreExpanded<F> {
102 evals: Vec<Vec<F>>,
104 expanded: Vec<Vec<F>>,
110}
111
112impl<F: BinaryField> GenericPreExpanded<F> {
113 pub fn generate_from_subspace(subspace: &BinarySubspace<F>) -> Self {
117 let evals = generate_evals_from_subspace(subspace);
118
119 let mut expanded = Vec::with_capacity(evals.len());
120 for basis in evals.iter().rev() {
121 let mut expanded_i = Vec::with_capacity(1 << (basis.len() - 1));
122 expanded_i.push(F::ZERO);
123 for i in 1..basis.len() {
124 for j in 0..expanded_i.len() {
125 expanded_i.push(expanded_i[j] + basis[i]);
126 }
127 }
128 assert_eq!(expanded_i.len(), 1 << (basis.len() - 1));
129 expanded.push(expanded_i);
130 }
131 assert_eq!(expanded.len(), evals.len());
132
133 Self { evals, expanded }
134 }
135}
136
137impl<F: BinaryField> DomainContext for GenericPreExpanded<F> {
138 type Field = F;
139
140 fn log_domain_size(&self) -> usize {
141 self.evals.len()
142 }
143
144 fn subspace(&self, i: usize) -> BinarySubspace<F> {
145 if i == 0 {
146 return BinarySubspace::with_dim(0);
147 }
148 BinarySubspace::new_unchecked(self.evals[self.log_domain_size() - i].clone())
149 }
150
151 fn twiddle(&self, layer: usize, block: usize) -> F {
152 self.expanded[layer][block]
153 }
154
155 fn iter_twiddles(&self, layer: usize, log_step_by: usize) -> impl Iterator<Item = F> {
156 self.expanded[layer]
157 .iter()
158 .step_by(1 << log_step_by)
159 .copied()
160 }
161}
162
163fn gao_mateer_basis<F: BinaryField>(num_basis_elements: usize) -> Vec<F> {
171 const {
172 assert!(F::N_BITS.is_power_of_two(), "the field degree over F_2 must be a power of two");
173 }
174
175 let mut beta = F::TRACE_ONE_ELEMENT;
178
179 for _i in 0..(F::N_BITS - num_basis_elements) {
181 beta = beta.square() + beta;
182 }
183
184 let mut basis = vec![F::ZERO; num_basis_elements];
186 basis[num_basis_elements - 1] = beta;
187 for i in (1..num_basis_elements).rev() {
188 basis[i - 1] = basis[i].square() + basis[i];
189 }
190
191 assert_eq!(basis[0], F::ONE);
194
195 basis
196}
197
198#[derive(Debug, Clone)]
226pub struct GaoMateerOnTheFly<F> {
227 basis: Vec<F>,
229}
230
231impl<F: BinaryField> GaoMateerOnTheFly<F> {
232 pub fn generate(log_domain_size: usize) -> Self {
242 Self {
243 basis: gao_mateer_basis(log_domain_size),
244 }
245 }
246}
247
248impl<F: BinaryField> DomainContext for GaoMateerOnTheFly<F> {
249 type Field = F;
250
251 fn log_domain_size(&self) -> usize {
252 self.basis.len()
253 }
254
255 fn subspace(&self, i: usize) -> BinarySubspace<F> {
256 BinarySubspace::new_unchecked(self.basis[..i].to_vec())
257 }
258
259 fn twiddle(&self, layer: usize, block: usize) -> F {
260 BinarySubspace::new_unchecked(&self.basis[1..=layer]).get(block)
261 }
262
263 fn iter_twiddles(&self, layer: usize, log_step_by: usize) -> impl Iterator<Item = F> + '_ {
264 BinarySubspaceIterator::new(&self.basis[1 + log_step_by..=layer])
265 }
266}
267
268#[derive(Debug)]
273pub struct GaoMateerPreExpanded<F> {
274 basis: Vec<F>,
276 expanded: Vec<F>,
281}
282
283impl<F: BinaryField> GaoMateerPreExpanded<F> {
284 pub fn generate(log_domain_size: usize) -> Self {
294 let basis: Vec<F> = gao_mateer_basis(log_domain_size);
295
296 let mut expanded = Vec::with_capacity(1 << log_domain_size);
297 expanded.push(F::ZERO);
298 for i in 1..log_domain_size {
299 for j in 0..expanded.len() {
300 expanded.push(expanded[j] + basis[i]);
301 }
302 }
303 assert_eq!(expanded.len(), 1usize << (log_domain_size - 1));
304
305 Self { basis, expanded }
306 }
307}
308
309impl<F: BinaryField> DomainContext for GaoMateerPreExpanded<F> {
310 type Field = F;
311
312 fn log_domain_size(&self) -> usize {
313 self.basis.len()
314 }
315
316 fn subspace(&self, i: usize) -> BinarySubspace<F> {
317 BinarySubspace::new_unchecked(self.basis[..i].to_vec())
318 }
319
320 fn twiddle(&self, _layer: usize, block: usize) -> F {
321 self.expanded[block]
322 }
323
324 fn iter_twiddles(&self, layer: usize, log_step_by: usize) -> impl Iterator<Item = F> {
325 self.expanded[..1 << layer]
326 .iter()
327 .step_by(1 << log_step_by)
328 .copied()
329 }
330}
331
332#[cfg(test)]
333mod tests {
334 use binius_field::{GhashSq256b, Rijndael8b};
335
336 use super::*;
337 use crate::test_utils::B128;
338
339 fn test_equivalence<F: BinaryField>(
340 dc_1: &impl DomainContext<Field = F>,
341 dc_2: &impl DomainContext<Field = F>,
342 log_domain_size: usize,
343 ) {
344 assert_eq!(dc_1.log_domain_size(), log_domain_size);
345 assert_eq!(dc_2.log_domain_size(), log_domain_size);
346
347 for i in 0..log_domain_size {
348 assert_eq!(dc_1.subspace(i), dc_2.subspace(i));
349
350 for block in 0..1 << i {
351 assert_eq!(dc_1.twiddle(i, block), dc_2.twiddle(i, block));
352 }
353 }
354 assert_eq!(dc_1.subspace(log_domain_size), dc_2.subspace(log_domain_size));
355 }
356
357 #[test]
358 fn test_generic() {
359 const LOG_SIZE: usize = 5;
360
361 let subspace = BinarySubspace::with_dim(LOG_SIZE);
362
363 let dc_otf = GenericOnTheFly::<B128>::generate_from_subspace(&subspace);
364 let dc_pre = GenericPreExpanded::<B128>::generate_from_subspace(&subspace);
365
366 test_equivalence(&dc_otf, &dc_pre, LOG_SIZE);
367 }
368
369 #[test]
370 fn test_gao_mateer() {
371 const LOG_SIZE: usize = 5;
372
373 let dc_gm_otf = GaoMateerOnTheFly::<B128>::generate(LOG_SIZE);
374 let dc_gm_pre = GaoMateerPreExpanded::<B128>::generate(LOG_SIZE);
375 let dc_generic_otf =
376 GenericOnTheFly::<B128>::generate_from_subspace(&dc_gm_otf.subspace(LOG_SIZE));
377
378 test_equivalence(&dc_gm_otf, &dc_gm_pre, LOG_SIZE);
379 test_equivalence(&dc_gm_otf, &dc_generic_otf, LOG_SIZE);
380 }
381
382 #[test]
390 fn test_gao_mateer_over_every_field() {
391 fn check<F: BinaryField>(log_size: usize) {
392 let dc = GaoMateerPreExpanded::<F>::generate(log_size);
393 let basis = dc.subspace(log_size);
394 assert_eq!(basis.dim(), log_size);
395 assert_eq!(basis.basis()[0], F::ONE);
396
397 for window in basis.basis().windows(2) {
399 assert_eq!(window[0], window[1].square() + window[1]);
400 }
401 }
402
403 check::<Rijndael8b>(8);
405 check::<B128>(5);
406 check::<GhashSq256b>(5);
407 }
408
409 #[test]
410 fn test_iter_layer() {
411 const LOG_SIZE: usize = 7;
412
413 let dc_gm_otf = GaoMateerOnTheFly::<B128>::generate(LOG_SIZE);
414 let dc_gm_pre = GaoMateerPreExpanded::<B128>::generate(LOG_SIZE);
415 let subspace = BinarySubspace::with_dim(LOG_SIZE);
416 let dc_generic_otf = GenericOnTheFly::<B128>::generate_from_subspace(&subspace);
417 let dc_generic_pre = GenericPreExpanded::<B128>::generate_from_subspace(&subspace);
418
419 for layer in 0..LOG_SIZE {
421 let expected: Vec<_> = (0..1 << layer)
422 .map(|block| dc_gm_pre.twiddle(layer, block))
423 .collect();
424
425 assert_eq!(
426 dc_gm_otf.iter_twiddles(layer, 0).collect::<Vec<_>>(),
427 expected,
428 "GaoMateerOnTheFly iter_layer mismatch at layer {}",
429 layer
430 );
431 assert_eq!(
432 dc_gm_pre.iter_twiddles(layer, 0).collect::<Vec<_>>(),
433 expected,
434 "GaoMateerPreExpanded iter_layer mismatch at layer {}",
435 layer
436 );
437 assert_eq!(
438 dc_generic_otf.iter_twiddles(layer, 0).collect::<Vec<_>>(),
439 dc_generic_pre.iter_twiddles(layer, 0).collect::<Vec<_>>(),
440 "Generic iter_layer mismatch at layer {}",
441 layer
442 );
443 }
444 }
445}