1use binius_field::BinaryField;
69
70use super::DomainContext;
71
72#[derive(Debug, Clone)]
80pub struct NormalizedSubspacePolys<F> {
81 step_inv: Vec<F>,
86 beta_0_inv: F,
91}
92
93impl<F: BinaryField> NormalizedSubspacePolys<F> {
94 pub fn new<DC: DomainContext<Field = F>>(domain_context: &DC) -> Self {
100 let l = domain_context.log_domain_size();
101 assert!(l >= 1, "precondition: log_domain_size must be at least 1");
102
103 let step_inv = (0..l - 1)
106 .map(|k| {
107 let d = domain_context.subspace(l - k).basis()[1];
108 let step = d * (d + F::ONE);
109 assert_ne!(step, F::ZERO, "W_hat_{{k+1}} normalizer must be non-zero");
111 step.invert_or_zero()
112 })
113 .collect();
114
115 let beta_0 = domain_context.subspace(l).basis()[0];
117 assert_ne!(beta_0, F::ZERO, "beta_0 is a basis element, so it must be non-zero");
118
119 Self {
120 step_inv,
121 beta_0_inv: beta_0.invert_or_zero(),
122 }
123 }
124
125 pub const fn log_domain_size(&self) -> usize {
127 self.step_inv.len() + 1
128 }
129
130 pub fn evals_at(&self, x: F) -> Vec<F> {
143 let mut evals = Vec::with_capacity(self.log_domain_size());
144 let mut w = x * self.beta_0_inv;
145 evals.push(w);
146 for &step_inv in &self.step_inv {
147 w *= (w + F::ONE) * step_inv;
148 evals.push(w);
149 }
150 evals
151 }
152}
153
154pub fn evals_at_domain_index<F, DC>(domain_context: &DC, index: usize) -> Vec<F>
170where
171 F: BinaryField,
172 DC: DomainContext<Field = F>,
173{
174 let l = domain_context.log_domain_size();
175 assert!(index < 1 << l, "precondition: index must be less than 2^log_domain_size");
176
177 (0..l)
179 .map(|k| domain_context.subspace(l - k).get(index >> k))
180 .collect()
181}
182
183#[cfg(test)]
184mod tests {
185 use binius_compute::GlobalAllocator;
186 use binius_field::{Field, Ghash128b, util::expand_subset_products};
187 use proptest::prelude::*;
188 use rand::{SeedableRng, rngs::StdRng};
189
190 use super::*;
191 use crate::{
192 BinarySubspace, FieldBuffer,
193 bit_reverse::reverse_bits,
194 inner_product::inner_product,
195 ntt::{
196 AdditiveNTT, NeighborsLastSingleThread,
197 domain_context::{GaoMateerOnTheFly, GaoMateerPreExpanded, GenericPreExpanded},
198 },
199 reed_solomon::ReedSolomonCode,
200 test_utils::random_field_buffer,
201 };
202
203 type F = Ghash128b;
204
205 fn scaled_subspace(log_d: usize) -> BinarySubspace<F> {
211 let scale = F::new(5);
212 let basis = BinarySubspace::<F>::with_dim(log_d)
213 .basis()
214 .iter()
215 .map(|&b| b * scale)
216 .collect::<Vec<_>>();
217 BinarySubspace::new_unchecked(basis)
218 }
219
220 fn for_each_context(log_d: usize, check: impl Fn(&dyn Fn(usize) -> BinarySubspace<F>)) {
224 check(&|i| GaoMateerPreExpanded::<F>::generate(log_d).subspace(i));
225 let standard = GenericPreExpanded::generate_from_subspace(&BinarySubspace::with_dim(log_d));
226 check(&|i| standard.subspace(i));
227 let scaled = GenericPreExpanded::generate_from_subspace(&scaled_subspace(log_d));
228 check(&|i| scaled.subspace(i));
229 }
230
231 #[test]
232 fn evals_at_basis_elements_match_the_domain_context() {
233 for log_d in 1..7 {
234 for_each_context(log_d, |subspace| {
237 let dc = GenericPreExpanded::generate_from_subspace(&subspace(log_d));
238 let polys = NormalizedSubspacePolys::new(&dc);
239 let domain = subspace(log_d);
240 for k in 0..log_d {
241 for (j, &expected) in subspace(log_d - k).basis().iter().enumerate() {
242 let got = polys.evals_at(domain.basis()[k + j])[k];
243 assert_eq!(got, expected, "log_d={log_d} k={k} j={j}");
244 }
245 }
246 });
247 }
248 }
249
250 #[test]
251 fn what_k_is_normalized_and_vanishes_below_its_subspace() {
252 for log_d in 1..7 {
253 for_each_context(log_d, |subspace| {
254 let dc = GenericPreExpanded::generate_from_subspace(&subspace(log_d));
255 let polys = NormalizedSubspacePolys::new(&dc);
256 let domain = subspace(log_d);
257 for k in 0..log_d {
258 for j in 0..k {
260 assert_eq!(polys.evals_at(domain.basis()[j])[k], F::ZERO);
261 }
262 assert_eq!(polys.evals_at(domain.basis()[k])[k], F::ONE);
264 }
265 });
266 }
267 }
268
269 #[test]
270 #[should_panic(expected = "normalizer must be non-zero")]
271 fn new_rejects_a_dependent_basis() {
272 let dependent = BinarySubspace::new_unchecked(vec![F::new(5), F::new(22), F::new(19)]);
274 NormalizedSubspacePolys::new(&GenericPreExpanded::generate_from_subspace(&dependent));
275 }
276
277 #[test]
278 fn what_k_vanishes_at_zero() {
279 let dc = GaoMateerPreExpanded::<F>::generate(5);
280 let polys = NormalizedSubspacePolys::new(&dc);
281 assert!(polys.evals_at(F::ZERO).iter().all(|&w| w == F::ZERO));
283 }
284
285 #[test]
286 fn domain_index_route_matches_arbitrary_point_route() {
287 for log_d in 1..7 {
288 for_each_context(log_d, |subspace| {
289 let dc = GenericPreExpanded::generate_from_subspace(&subspace(log_d));
290 let polys = NormalizedSubspacePolys::new(&dc);
291 let domain = subspace(log_d);
292 for index in 0..1 << log_d {
294 assert_eq!(
295 evals_at_domain_index(&dc, index),
296 polys.evals_at(domain.get(index)),
297 "log_d={log_d} index={index}"
298 );
299 }
300 });
301 }
302 }
303
304 #[test]
305 #[should_panic(expected = "index must be less than 2^log_domain_size")]
306 fn evals_at_domain_index_rejects_an_out_of_range_index() {
307 let dc = GaoMateerPreExpanded::<F>::generate(3);
308 evals_at_domain_index(&dc, 8);
309 }
310
311 fn assert_tensor_row_matches_ntt<DC>(log_d: usize, dc: DC, seed: u64)
313 where
314 DC: DomainContext<Field = F>,
315 {
316 let mut rng = StdRng::seed_from_u64(seed);
317 let coeffs = random_field_buffer::<F>(&mut rng, log_d);
318
319 let mut transformed = coeffs.clone();
320 let ntt = NeighborsLastSingleThread::new(dc);
321 ntt.forward_transform(transformed.as_mut_view(), 0, 0);
322
323 for index in 0..1 << log_d {
324 let row = expand_subset_products(&evals_at_domain_index(ntt.domain_context(), index));
326 let dot = inner_product(coeffs.as_ref().iter().copied(), row);
327 assert_eq!(dot, transformed.as_ref()[index], "log_d={log_d} index={index}");
328 }
329 }
330
331 #[test]
332 fn tensor_row_matches_ntt_over_every_domain_context() {
333 for log_d in 1..7 {
334 assert_tensor_row_matches_ntt(log_d, GaoMateerPreExpanded::<F>::generate(log_d), 0);
335 let standard = BinarySubspace::<F>::with_dim(log_d);
336 assert_tensor_row_matches_ntt(
337 log_d,
338 GenericPreExpanded::generate_from_subspace(&standard),
339 1,
340 );
341 assert_tensor_row_matches_ntt(
342 log_d,
343 GenericPreExpanded::generate_from_subspace(&scaled_subspace(log_d)),
344 2,
345 );
346 }
347 }
348
349 fn generator_rows(code: &ReedSolomonCode<F>) -> Vec<Vec<F>> {
355 let dc = GaoMateerPreExpanded::<F>::generate(code.log_len());
356 (0..1 << code.log_len())
357 .map(|index| {
358 let mut evals = evals_at_domain_index(&dc, index);
359 evals.truncate(code.log_dim());
360 evals.reverse();
361 expand_subset_products(&evals)
362 })
363 .collect()
364 }
365
366 #[test]
367 fn tensor_row_matches_reed_solomon_encoding() {
368 for log_dim in 1..6 {
369 for log_inv_rate in 1..4 {
370 let code = ReedSolomonCode::<F>::new(log_dim, log_inv_rate);
371 let ntt = NeighborsLastSingleThread::new(GaoMateerOnTheFly::<F>::generate(
372 code.log_len(),
373 ));
374 let mut rng = StdRng::seed_from_u64(7);
375 let msg = random_field_buffer::<F>(&mut rng, log_dim);
376 let codeword = code.encode_batch(&ntt, msg.as_view(), 0, &GlobalAllocator);
377
378 for (index, row) in generator_rows(&code).into_iter().enumerate() {
379 let dot = inner_product(msg.as_ref().iter().copied(), row);
380 assert_eq!(
381 dot,
382 codeword.as_ref()[index],
383 "log_dim={log_dim} log_inv_rate={log_inv_rate} index={index}"
384 );
385 }
386 }
387 }
388 }
389
390 #[test]
391 fn interleaved_lanes_are_independent_codewords() {
392 for log_dim in 1..5 {
393 for log_inv_rate in 1..3 {
394 for log_batch in 1..3 {
395 let code = ReedSolomonCode::<F>::new(log_dim, log_inv_rate);
396 let ntt = NeighborsLastSingleThread::new(GaoMateerOnTheFly::<F>::generate(
397 code.log_len(),
398 ));
399 let mut rng = StdRng::seed_from_u64(11);
400 let msg = random_field_buffer::<F>(&mut rng, log_dim + log_batch);
401 let codeword =
402 code.encode_batch(&ntt, msg.as_view(), log_batch, &GlobalAllocator);
403 let rows = generator_rows(&code);
404
405 for lane in 0..1 << log_batch {
406 let base = reverse_bits(lane, log_batch as u32) << log_dim;
408 let lane_msg = (0..1 << log_dim)
409 .map(|j| msg.as_ref()[base | j])
410 .collect::<Vec<_>>();
411 let lane_msg = FieldBuffer::<F, Vec<F>>::new(log_dim, lane_msg);
412
413 for (index, row) in rows.iter().enumerate() {
415 let dot = inner_product(
416 lane_msg.as_ref().iter().copied(),
417 row.iter().copied(),
418 );
419 assert_eq!(
420 dot,
421 codeword.as_ref()[(index << log_batch) | lane],
422 "log_dim={log_dim} r={log_inv_rate} b={log_batch} lane={lane}"
423 );
424 }
425 }
426 }
427 }
428 }
429 }
430
431 proptest! {
432 #[test]
433 fn tensor_row_matches_ntt_on_random_coefficients(seed: u64) {
434 const LOG_D: usize = 5;
435 assert_tensor_row_matches_ntt(LOG_D, GaoMateerPreExpanded::<F>::generate(LOG_D), seed);
437 }
438
439 #[test]
440 fn evals_at_is_f2_linear(a: u64, b: u64) {
441 const LOG_D: usize = 6;
442 let dc = GaoMateerPreExpanded::<F>::generate(LOG_D);
443 let polys = NormalizedSubspacePolys::new(&dc);
444 let (x, y) = (F::new(a as u128), F::new(b as u128));
445 let sum = polys.evals_at(x + y);
447 let parts = std::iter::zip(polys.evals_at(x), polys.evals_at(y));
448 for (got, (wx, wy)) in std::iter::zip(sum, parts) {
449 prop_assert_eq!(got, wx + wy);
450 }
451 }
452 }
453}