Skip to main content

binius_math/
binary_subspace.rs

1// Copyright 2024-2025 Irreducible Inc.
2
3//! Binary subspaces: $\mathbb{F}_2$-linear spans of a binary field, enumerated in order.
4//!
5//! A subspace element is a subset-XOR of ordered basis elements, chosen by the bits of an index.
6//!
7//! Walking the elements in order is a binary-counter increment.
8//! Each step XORs in or out only the basis elements whose bit changed.
9
10use std::ops::Deref;
11
12use binius_field::{BinaryField, BinaryField1b};
13
14/// An $\mathbb{F}_2$-linear subspace of a binary field.
15///
16/// The subspace is the span of an ordered basis under XOR.
17/// The basis order fixes an order on the subspace's own elements too.
18#[derive(Debug, Clone, PartialEq, Eq)]
19pub struct BinarySubspace<F, Data: Deref<Target = [F]> = Vec<F>> {
20	/// Ordered basis elements the subspace is spanned by.
21	basis: Data,
22}
23
24impl<F: BinaryField, Data: Deref<Target = [F]>> BinarySubspace<F, Data> {
25	/// Creates a new subspace from a vector of ordered basis elements.
26	///
27	/// This constructor does not check that the basis elements are linearly independent.
28	pub const fn new_unchecked(basis: Data) -> Self {
29		Self { basis }
30	}
31
32	/// Creates a new subspace isomorphic to this one, over a different field type.
33	///
34	/// Maps each basis element into `FIso` via `From`, keeping the same order.
35	pub fn isomorphic<FIso>(&self) -> BinarySubspace<FIso>
36	where
37		FIso: BinaryField + From<F>,
38	{
39		BinarySubspace {
40			// Convert every basis element into the target field, in the same order.
41			basis: self.basis.iter().copied().map(Into::into).collect(),
42		}
43	}
44
45	/// Returns the dimension of the subspace.
46	pub fn dim(&self) -> usize {
47		self.basis.len()
48	}
49
50	/// Returns the slice of ordered basis elements.
51	pub fn basis(&self) -> &[F] {
52		&self.basis
53	}
54
55	/// Returns the subspace element selected by `index`.
56	///
57	/// Bit `i` of `index` selects whether basis element `i` is included in the sum.
58	/// Basis elements combine over $\mathbb{F}_2$, so "included" means XORed in.
59	///
60	/// # Arguments
61	/// - `index`: which subspace element to return, in `0..2^dim`.
62	///
63	/// # Panics
64	/// Panics if `index` is at least `2^dim`. Once `dim` reaches `usize::BITS`, every `usize`
65	/// is below `2^dim`, so no index can panic.
66	pub fn get(&self, index: usize) -> F {
67		// Once the dimension reaches usize::BITS, 2^dim exceeds every usize, so the bound
68		// holds for any index and computing it would overflow.
69		assert!(
70			self.dim() >= usize::BITS as usize || index < 1 << self.dim(),
71			"precondition: index must be less than 2^dim"
72		);
73
74		element_at(&self.basis, index)
75	}
76
77	/// Returns an iterator over every element of the subspace, in index order.
78	///
79	/// # Panics
80	/// Panics if the subspace's dimension is at least `usize::BITS`.
81	/// An index that large would not fit in a `usize`.
82	pub fn iter(&self) -> BinarySubspaceIterator<'_, F> {
83		BinarySubspaceIterator::new(&self.basis)
84	}
85}
86
87impl<F: BinaryField> BinarySubspace<F> {
88	/// Creates a subspace spanned by the field's first `dim` default basis elements.
89	///
90	/// Uses a prefix of the field's own canonical $\mathbb{F}_2$ basis.
91	/// So a smaller dimension is always a prefix of a larger one's basis.
92	///
93	/// # Panics
94	/// Panics if `dim` is greater than `F::DEGREE`.
95	pub fn with_dim(dim: usize) -> Self {
96		assert!(dim <= F::DEGREE, "precondition: dim must be at most F::DEGREE");
97
98		// Take the field's own first `dim` basis elements, in order.
99		let basis = (0..dim).map(|i| F::basis(i)).collect();
100		Self { basis }
101	}
102
103	/// Creates a smaller subspace using a prefix of this subspace's basis.
104	///
105	/// # Panics
106	/// Panics if `dim` is greater than this subspace's own dimension.
107	pub fn reduce_dim(&self, dim: usize) -> Self {
108		assert!(dim <= self.dim(), "precondition: dim must be at most this subspace's dimension");
109
110		Self {
111			basis: self.basis[..dim].to_vec(),
112		}
113	}
114}
115
116/// Computes the subset-XOR of `basis` selected by the bits of `index`.
117fn element_at<F: BinaryField>(basis: &[F], index: usize) -> F {
118	basis
119		.iter()
120		// A basis element past the top of the index has no bit to select it, so the sum stops
121		// there. This also keeps the shift below in range without a per-element branch.
122		.take(usize::BITS as usize)
123		.enumerate()
124		// Keep basis_i when bit i of index is set.
125		// Drop it (multiply by 0) otherwise.
126		.map(|(i, &basis_i)| basis_i * BinaryField1b::from((index >> i) & 1 == 1))
127		.sum()
128}
129
130/// Iterator over every element of a binary subspace, in index order.
131///
132/// Each element is a subset-XOR of the basis elements.
133/// Stepping forward reuses the previous value instead of recomputing a full subset sum.
134/// Skipping ahead computes the landing value directly instead.
135#[derive(Debug, Clone)]
136pub struct BinarySubspaceIterator<'a, F> {
137	/// The subspace's ordered basis elements.
138	basis: &'a [F],
139	/// Index of the next element this iterator will yield.
140	index: usize,
141	/// The next element's value, precomputed so stepping never redoes a full sum.
142	next: Option<F>,
143}
144
145impl<'a, F: BinaryField> BinarySubspaceIterator<'a, F> {
146	/// Creates an iterator starting at index 0.
147	///
148	/// # Panics
149	/// Panics if the basis has `usize::BITS` or more elements.
150	/// An index that large would not fit in a `usize`.
151	pub fn new(basis: &'a [F]) -> Self {
152		assert!(basis.len() < usize::BITS as usize);
153
154		// Index 0 selects no basis elements, so its value is always zero.
155		Self {
156			basis,
157			index: 0,
158			next: Some(F::ZERO),
159		}
160	}
161}
162
163impl<'a, F: BinaryField> Iterator for BinarySubspaceIterator<'a, F> {
164	type Item = F;
165
166	/// Advances to the next index.
167	/// Reuses the previous value instead of recomputing a full subset sum.
168	///
169	/// # Algorithm
170	/// Moving from `index` to `index + 1` is a binary-counter increment.
171	/// A run of trailing 1-bits flips to 0, then the next 0-bit flips to 1.
172	///
173	/// A bit flip here means XOR-ing a basis element in or out.
174	/// So the update only touches elements whose bit actually changed.
175	#[inline]
176	fn next(&mut self) -> Option<Self::Item> {
177		let ret = self.next?;
178
179		// Length of the trailing run of 1-bits.
180		// Found with one hardware instruction, not a bit-by-bit scan.
181		let ones = self.index.trailing_ones() as usize;
182
183		// Undo every bit in that run: each one flips to 0, so XOR its basis element back out.
184		let mut next = ret;
185		for &basis_i in &self.basis[..ones] {
186			next -= basis_i;
187		}
188
189		// The bit right after that run flips from 0 to 1: XOR its basis element in.
190		// No such bit left in the basis means this was the last element.
191		self.next = self.basis.get(ones).map(|&basis_i| next + basis_i);
192
193		self.index += 1;
194		Some(ret)
195	}
196
197	fn size_hint(&self) -> (usize, Option<usize>) {
198		// Total element count is 2^dim.
199		// Subtract how many indices are already past.
200		let last = 1 << self.basis.len();
201		let remaining = last - self.index;
202		(remaining, Some(remaining))
203	}
204
205	/// Skips ahead `n` elements, computing the landing value directly.
206	/// Never steps through everything in between.
207	///
208	/// Exhausts the iterator instead of panicking when `n` overflows the index.
209	/// Also exhausts it if the landing point is past the last element.
210	fn nth(&mut self, n: usize) -> Option<Self::Item> {
211		match self.index.checked_add(n) {
212			// Lands inside range: jump straight to that index's value.
213			Some(new_index) if new_index < 1 << self.basis.len() => {
214				self.index = new_index;
215				self.next = Some(element_at(self.basis, new_index));
216			}
217			// Overflowed, or landed past the end: exhaust the iterator.
218			_ => {
219				self.index = 1 << self.basis.len();
220				self.next = None;
221			}
222		}
223
224		self.next()
225	}
226}
227
228impl<'a, F: BinaryField> ExactSizeIterator for BinarySubspaceIterator<'a, F> {
229	fn len(&self) -> usize {
230		// Same total as the length hint, without the Option wrapper this trait doesn't need.
231		let last = 1 << self.basis.len();
232		last - self.index
233	}
234}
235
236impl<'a, F: BinaryField> std::iter::FusedIterator for BinarySubspaceIterator<'a, F> {}
237
238impl<F: BinaryField> Default for BinarySubspace<F> {
239	/// The default subspace spans the whole field, using its full canonical basis.
240	fn default() -> Self {
241		// Every basis element of the field, in canonical order.
242		let basis = (0..F::DEGREE).map(|i| F::basis(i)).collect();
243		Self { basis }
244	}
245}
246
247#[cfg(test)]
248mod tests {
249	use binius_field::{ExtensionField, Field, Ghash128b as B128, Rijndael8b as B8};
250
251	use super::*;
252
253	#[test]
254	fn test_default_binary_subspace_iterates_elements() {
255		// The default basis for an 8-bit field is the powers of two.
256		// So get(i) reconstructs i exactly, for every byte value.
257		let subspace = BinarySubspace::<B8>::default();
258		for i in 0..=255 {
259			assert_eq!(subspace.get(i), B8::new(i as u8));
260		}
261	}
262
263	#[test]
264	fn test_get_on_a_subspace_wider_than_a_usize() {
265		// The default basis of a 128-bit field has 128 elements, so 2^dim does not fit in a
266		// usize. Every index that fits in a usize is still in range and must be selectable.
267		let basis = <B128 as ExtensionField<BinaryField1b>>::basis;
268		let subspace = BinarySubspace::<B128>::default();
269		assert_eq!(subspace.dim(), 128);
270		assert_eq!(subspace.get(0), B128::ZERO);
271		assert_eq!(subspace.get(1), basis(0));
272		assert_eq!(subspace.get(5), basis(0) + basis(2));
273		// The basis elements above bit 63 have no index bit to select them.
274		let low_bits: B128 = (0..usize::BITS as usize).map(basis).sum();
275		assert_eq!(subspace.get(usize::MAX), low_bits);
276	}
277
278	#[test]
279	#[should_panic(expected = "precondition")]
280	fn test_binary_subspace_range_error() {
281		// dim = 8, so valid indices are 0..256; 256 is one past the last valid index.
282		let subspace = BinarySubspace::<B8>::default();
283		let _ = subspace.get(256);
284	}
285
286	#[test]
287	fn test_default_binary_subspace() {
288		let subspace = BinarySubspace::<B8>::default();
289		assert_eq!(subspace.dim(), 8);
290		assert_eq!(subspace.basis().len(), 8);
291
292		// The default basis is the field's own bits, in order: 1, 2, 4, ..., 128.
293		assert_eq!(
294			subspace.basis(),
295			[
296				B8::new(0b00000001),
297				B8::new(0b00000010),
298				B8::new(0b00000100),
299				B8::new(0b00001000),
300				B8::new(0b00010000),
301				B8::new(0b00100000),
302				B8::new(0b01000000),
303				B8::new(0b10000000)
304			]
305		);
306
307		// With that basis, index and value coincide: get(i) is just i itself.
308		let expected_elements: [u8; 256] = (0..=255).collect::<Vec<_>>().try_into().unwrap();
309
310		for (i, &expected) in expected_elements.iter().enumerate() {
311			assert_eq!(subspace.get(i), B8::new(expected));
312		}
313	}
314
315	#[test]
316	fn test_with_dim_valid() {
317		// A 3-dimensional subspace only uses the field's first 3 basis elements: 1, 2, 4.
318		let subspace = BinarySubspace::<B8>::with_dim(3);
319		assert_eq!(subspace.dim(), 3);
320		assert_eq!(subspace.basis().len(), 3);
321
322		assert_eq!(subspace.basis(), [B8::new(0b001), B8::new(0b010), B8::new(0b100)]);
323
324		// So it spans exactly the 8 values expressible in 3 bits: 0..8.
325		let expected_elements: [u8; 8] = [0b000, 0b001, 0b010, 0b011, 0b100, 0b101, 0b110, 0b111];
326
327		for (i, &expected) in expected_elements.iter().enumerate() {
328			assert_eq!(subspace.get(i), B8::new(expected));
329		}
330	}
331
332	#[test]
333	#[should_panic(expected = "precondition")]
334	fn test_with_dim_invalid() {
335		// B8 has degree 8, so dimension 10 is out of range.
336		let _ = BinarySubspace::<B8>::with_dim(10);
337	}
338
339	#[test]
340	fn test_reduce_dim_valid() {
341		// Start from a 6-dimensional subspace, then keep only its first 4 basis elements.
342		let subspace = BinarySubspace::<B8>::with_dim(6);
343		let reduced = subspace.reduce_dim(4);
344		assert_eq!(reduced.dim(), 4);
345		assert_eq!(reduced.basis().len(), 4);
346
347		// A prefix of the basis, so this matches with_dim(4) exactly.
348		assert_eq!(
349			reduced.basis(),
350			[
351				B8::new(0b0001),
352				B8::new(0b0010),
353				B8::new(0b0100),
354				B8::new(0b1000)
355			]
356		);
357
358		let expected_elements: [u8; 16] = (0..16).collect::<Vec<_>>().try_into().unwrap();
359
360		for (i, &expected) in expected_elements.iter().enumerate() {
361			assert_eq!(reduced.get(i), B8::new(expected));
362		}
363	}
364
365	#[test]
366	#[should_panic(expected = "precondition")]
367	fn test_reduce_dim_invalid() {
368		// Can't reduce to a larger dimension than the subspace already has.
369		let subspace = BinarySubspace::<B8>::with_dim(4);
370		let _ = subspace.reduce_dim(6);
371	}
372
373	#[test]
374	fn test_isomorphic_conversion() {
375		let subspace = BinarySubspace::<B8>::with_dim(3);
376		// Re-express the same 3 basis elements as values of a much larger field.
377		let iso_subspace: BinarySubspace<B128> = subspace.isomorphic();
378		assert_eq!(iso_subspace.dim(), 3);
379		assert_eq!(iso_subspace.basis().len(), 3);
380
381		// Same basis values, just converted into B128 via From, in the same order.
382		assert_eq!(
383			iso_subspace.basis(),
384			[
385				B128::from(B8::new(0b001)),
386				B128::from(B8::new(0b010)),
387				B128::from(B8::new(0b100)),
388			]
389		);
390	}
391
392	#[test]
393	fn test_iterate_subspace() {
394		let subspace = BinarySubspace::<B8>::with_dim(3);
395		// Collecting the iterator gives exactly the 8 elements of a 3-dim subspace.
396		let elements: Vec<_> = subspace.iter().collect();
397		assert_eq!(elements.len(), 8);
398
399		let expected_elements: [u8; 8] = [0b000, 0b001, 0b010, 0b011, 0b100, 0b101, 0b110, 0b111];
400
401		for (i, &expected) in expected_elements.iter().enumerate() {
402			assert_eq!(elements[i], B8::new(expected));
403		}
404	}
405
406	#[test]
407	fn test_iterator_matches_get() {
408		let subspace = BinarySubspace::<B8>::with_dim(5);
409
410		// The incremental iterator and the direct per-index formula must always agree.
411		for (i, elem) in subspace.iter().enumerate() {
412			assert_eq!(elem, subspace.get(i), "Mismatch at index {}", i);
413		}
414	}
415
416	#[test]
417	#[allow(clippy::iter_nth_zero)]
418	fn test_iterator_nth() {
419		let subspace = BinarySubspace::<B8>::with_dim(4);
420
421		// nth(0) behaves like next(): returns the very next element, advances by 1.
422		let mut iter = subspace.iter();
423		assert_eq!(iter.nth(0), Some(subspace.get(0)));
424		assert_eq!(iter.nth(0), Some(subspace.get(1)));
425		// Larger skips land on the element that many steps further along.
426		assert_eq!(iter.nth(2), Some(subspace.get(4)));
427		assert_eq!(iter.nth(5), Some(subspace.get(10)));
428
429		// Landing exactly on the last valid index still works.
430		let mut iter = subspace.iter();
431		assert_eq!(iter.nth(15), Some(subspace.get(15)));
432		// One more step past the end exhausts the iterator.
433		assert_eq!(iter.nth(0), None);
434	}
435
436	#[test]
437	fn test_iterator_nth_skips_efficiently() {
438		let subspace = BinarySubspace::<B8>::with_dim(6);
439
440		// Jump straight to index 30, without stepping through 0..30 first.
441		let mut iter = subspace.iter();
442		assert_eq!(iter.nth(30), Some(subspace.get(30)));
443		// A plain next() afterward continues from exactly where the jump landed.
444		assert_eq!(iter.next(), Some(subspace.get(31)));
445
446		// A larger single jump works the same way.
447		let mut iter = subspace.iter();
448		assert_eq!(iter.nth(50), Some(subspace.get(50)));
449	}
450
451	#[test]
452	fn test_iterator_size_hint() {
453		let subspace = BinarySubspace::<B8>::with_dim(3);
454		let mut iter = subspace.iter();
455
456		// 3 dimensions means 8 elements total, all still ahead at the start.
457		assert_eq!(iter.size_hint(), (8, Some(8)));
458		iter.next();
459		assert_eq!(iter.size_hint(), (7, Some(7)));
460		// Skipping 3 ahead accounts for all 4 consumed elements at once.
461		iter.nth(3);
462		assert_eq!(iter.size_hint(), (3, Some(3)));
463	}
464
465	#[test]
466	fn test_iterator_exact_size() {
467		let subspace = BinarySubspace::<B8>::with_dim(4);
468		let mut iter = subspace.iter();
469
470		assert_eq!(iter.len(), 16);
471		iter.next();
472		assert_eq!(iter.len(), 15);
473		iter.nth(5);
474		assert_eq!(iter.len(), 9);
475	}
476
477	#[test]
478	fn test_iterator_empty_subspace() {
479		// Dimension 0 has exactly one element: the empty XOR-sum, zero.
480		let subspace = BinarySubspace::<B8>::with_dim(0);
481		let mut iter = subspace.iter();
482
483		assert_eq!(iter.len(), 1);
484		assert_eq!(iter.next(), Some(B8::ZERO));
485		assert_eq!(iter.next(), None);
486	}
487
488	#[test]
489	fn test_iterator_full_iteration() {
490		// The full 8-bit field has 256 elements.
491		// The iterator must produce all of them, matching get() at every index.
492		let subspace = BinarySubspace::<B8>::default();
493		let collected: Vec<_> = subspace.iter().collect();
494
495		assert_eq!(collected.len(), 256);
496		for (i, elem) in collected.iter().enumerate() {
497			assert_eq!(*elem, subspace.get(i));
498		}
499	}
500
501	#[test]
502	fn test_iterator_partial_then_nth() {
503		let subspace = BinarySubspace::<B8>::with_dim(5);
504		let mut iter = subspace.iter();
505
506		// Step through the first 3 elements one at a time.
507		assert_eq!(iter.next(), Some(subspace.get(0)));
508		assert_eq!(iter.next(), Some(subspace.get(1)));
509		assert_eq!(iter.next(), Some(subspace.get(2)));
510
511		// Jump ahead 5 more (landing on index 8), then continue normally.
512		assert_eq!(iter.nth(5), Some(subspace.get(8)));
513		assert_eq!(iter.next(), Some(subspace.get(9)));
514	}
515
516	#[test]
517	fn test_iterator_clone() {
518		let subspace = BinarySubspace::<B8>::with_dim(3);
519		let mut iter1 = subspace.iter();
520
521		iter1.next();
522		iter1.next();
523
524		// Cloning mid-iteration copies the current position, not a fresh start.
525		let mut iter2 = iter1.clone();
526
527		assert_eq!(iter1.next(), iter2.next());
528		assert_eq!(iter1.collect::<Vec<_>>(), iter2.collect::<Vec<_>>());
529	}
530}