Skip to main content

binius_math/
bit_reverse.rs

1// Copyright 2025 Irreducible Inc.
2// Copyright 2026 The Binius Developers
3
4use std::ptr;
5
6use binius_field::{PackedField, transpose_square_blocks};
7use binius_utils::{
8	checked_arithmetics::checked_log_2,
9	rayon::{prelude::*, task_size::min_len_for_bytes},
10};
11use bytemuck::zeroed_vec;
12
13use crate::field_buffer::FieldSliceMut;
14
15/// Reverses the low `bits` bits of an unsigned integer.
16///
17/// # Arguments
18///
19/// * `x` - The value whose bits to reverse
20/// * `bits` - The number of low-order bits to reverse
21///
22/// # Returns
23///
24/// The value with its low `bits` bits reversed
25pub const fn reverse_bits(x: usize, bits: u32) -> usize {
26	x.reverse_bits().unbounded_shr(usize::BITS - bits)
27}
28
29/// Bytes one row of a tile should span.
30///
31/// Rows sit far apart in the buffer, so each one costs a memory stream of its own.
32/// Four cache lines is where a stream is wide enough to reach peak bandwidth.
33///
34/// Going wider costs more than it returns: two squares have to fit in the first cache level.
35const TILE_ROW_BYTES: usize = 256;
36
37/// Base-2 log of the widest tile a permutation instance moves, counted in packed elements.
38///
39/// # Why this value
40///
41/// - Every tile-sized loop takes its bound from this parameter.
42/// - Each step it may take therefore doubles the number of instances compiled.
43/// - Eight elements already reach the row width above once an element is 32 bytes or wider.
44/// - A narrower element would prefer one step more, and one packing here is narrower.
45/// - Measured, that step gains past the last cache level and loses below it.
46const MAX_LOG_TILE_PACKED: usize = 3;
47
48/// Permutes the elements of a buffer by reversing the bits of every index, in place.
49///
50/// An index is reversed over the bit width that the length of the buffer takes.
51/// Packed representations are handled, so one exchange may move lanes within a single element.
52///
53/// ```text
54///     8 elements:   index 001 -> index 100
55/// ```
56///
57/// # Arguments
58///
59/// * `buffer` - the mutable view of packed field elements to permute
60pub fn bit_reverse_packed<P: PackedField>(buffer: FieldSliceMut<'_, P>) {
61	// A buffer shorter than a square of packed elements leaves the tiled path no room.
62	let log_len = buffer.log_len();
63	if log_len < 2 * P::LOG_WIDTH {
64		return bit_reverse_packed_naive(buffer);
65	}
66
67	// Scalars filling one tile row, which is what every gather and scatter moves at a time.
68	let log_scalar_bytes = size_of::<P::Scalar>().next_power_of_two().ilog2() as usize;
69	let log_row = checked_log_2(TILE_ROW_BYTES).saturating_sub(log_scalar_bytes);
70
71	// A tile holds at least one packed element, since a narrower one cannot be addressed.
72	// It stops at the widest instance compiled, and at half the length of a short buffer.
73	//
74	// The check above leaves half the length at or above the packing width.
75	// So neither clamp can cut a tile below one element.
76	let log_tile = log_row
77		.max(P::LOG_WIDTH)
78		.min(P::LOG_WIDTH + MAX_LOG_TILE_PACKED)
79		.min(log_len / 2);
80
81	// Why: a constant tile bound is what keeps the gather and scatter unrolled.
82	// A run-time bound lowers the moves of a one-element tile to a call.
83	// That call is most of the work at that width.
84	match log_tile - P::LOG_WIDTH {
85		0 => bit_reverse_paired::<P, 0>(buffer),
86		1 => bit_reverse_paired::<P, 1>(buffer),
87		2 => bit_reverse_paired::<P, 2>(buffer),
88		// The clamp caps the tile at the widest instance, which answers every larger value.
89		_ => bit_reverse_paired::<P, MAX_LOG_TILE_PACKED>(buffer),
90	}
91}
92
93/// Applies a bit-reversal permutation in a single pass over memory.
94///
95/// # Overview
96///
97/// Split the scalar index into three fields, with equal width at both ends:
98///
99/// ```text
100///     i = (h, m, l)        |h| = |l| = log_tile
101/// ```
102///
103/// A tile is `2^log_tile` consecutive scalars.
104/// So `l` picks a scalar inside a tile, and `h` picks a tile.
105///
106/// Reversing every bit sends `(h, m, l)` to `(rev l, rev m, rev h)`.
107/// Reading that destination-first is what lets one pass do the whole job:
108///
109/// ```text
110///     new(H, M, L) = old(rev L, rev M, rev H)
111/// ```
112///
113/// # Algorithm
114///
115/// The middle field alone decides where a scalar goes.
116/// Everything under middle index `rev M` lands under `M`, and nothing else does.
117///
118/// The scalars sharing a middle index form a square of `2^log_tile` tiles.
119/// So squares pair up, and a pair is the whole unit of work:
120///
121/// ```text
122///     square(M)      <-  transpose of square(rev M)
123///     square(rev M)  <-  transpose of square(M)
124/// ```
125///
126/// Both are read into scratch, transposed there, and written back crossed.
127/// A middle index equal to its own reversal is a fixed point, and needs one square.
128///
129/// # Why one pass
130///
131/// Two passes would transpose each square where it sits, then permute the squares.
132/// Each pass reads and writes the whole buffer, so the permutation costs four traversals.
133///
134/// Pairing the squares fuses those passes into one.
135/// Every scalar is then read once and written once, which halves the memory traffic.
136///
137/// # Preconditions
138///
139/// * `buffer.log_len() >= 2 * (P::LOG_WIDTH + LOG_TILE_PACKED)`
140fn bit_reverse_paired<P: PackedField, const LOG_TILE_PACKED: usize>(
141	mut buffer: FieldSliceMut<'_, P>,
142) {
143	// Tile width in scalars, and the middle index field the two ends leave over.
144	let log_tile = P::LOG_WIDTH + LOG_TILE_PACKED;
145	let log_len = buffer.log_len();
146	debug_assert!(log_len >= 2 * log_tile);
147	let log_mid = log_len - 2 * log_tile;
148
149	// Words in one square: `2^log_tile` rows, each `2^LOG_TILE_PACKED` words wide.
150	let log_square = log_tile + LOG_TILE_PACKED;
151
152	let data = buffer.as_mut();
153	// Holding an address rather than the slice is what lets disjoint tasks write one buffer.
154	let data_ptr = data.as_mut_ptr() as usize;
155
156	// One iteration moves two squares, and half of them return at once.
157	// The byte budget counts single words, so divide it by the words a square holds.
158	let min_len = (min_len_for_bytes::<P>() >> log_square).max(1);
159
160	(0..1 << log_mid)
161		.into_par_iter()
162		.with_min_len(min_len)
163		.for_each_init(
164			|| {
165				(
166					zeroed_vec::<P>(1 << log_square),
167					zeroed_vec::<P>(1 << log_square),
168					zeroed_vec::<P>(P::WIDTH),
169				)
170			},
171			|(square, mirror, block), m| {
172				let m_rev = reverse_bits(m, log_mid as u32);
173
174				// A pair is claimed by its lower middle index, so the higher one returns.
175				// A fixed point claims itself, and takes the single-square path below.
176				if m_rev < m {
177					return;
178				}
179
180				// First word of the tile at high index `h` under middle index `mm`.
181				// The three index fields hold disjoint bit ranges, so shifting places them.
182				let row = |h: usize, mm: usize| {
183					(h << (log_mid + LOG_TILE_PACKED)) | (mm << LOG_TILE_PACKED)
184				};
185
186				// SAFETY:
187				// - Every word addressed here carries `m` or its reversal in the middle field.
188				// - A pair is claimed once, so no two iterations name the same middle index.
189				// - Tasks therefore write disjoint ranges.
190				// - The widest address reached is `2^(log_len - P::LOG_WIDTH) - 1`, the last one.
191				// - The address stays live because the buffer outlives this loop.
192				unsafe {
193					let data = data_ptr as *mut P;
194
195					// Rows are visited at their high index reversed.
196					// That is what turns a plain transpose into a reversal of both end fields.
197					let gather = |dst: &mut [P], mm: usize| {
198						for j in 0..1 << log_tile {
199							let src = data.add(row(reverse_bits(j, log_tile as u32), mm));
200							let dst = dst.as_mut_ptr().add(j << LOG_TILE_PACKED);
201							for k in 0..1 << LOG_TILE_PACKED {
202								*dst.add(k) = *src.add(k);
203							}
204						}
205					};
206					let scatter = |src: &[P], mm: usize| {
207						for j in 0..1 << log_tile {
208							let src = src.as_ptr().add(j << LOG_TILE_PACKED);
209							let dst = data.add(row(reverse_bits(j, log_tile as u32), mm));
210							for k in 0..1 << LOG_TILE_PACKED {
211								*dst.add(k) = *src.add(k);
212							}
213						}
214					};
215
216					// Transposing in scratch keeps every exchange inside the first cache level.
217					gather(square, m);
218					transpose_tile::<P, LOG_TILE_PACKED>(square, block);
219
220					// A middle index that reverses to itself keeps its own square.
221					if m_rev == m {
222						scatter(square, m);
223						return;
224					}
225
226					// Otherwise the two squares trade places, each transposed on the way.
227					// Otherwise the two squares trade places, each transposed on the way.
228					// The mirror is read before the scatter that overwrites where it sat.
229					gather(mirror, m_rev);
230					transpose_tile::<P, LOG_TILE_PACKED>(mirror, block);
231					scatter(square, m_rev);
232					scatter(mirror, m);
233				}
234			},
235		);
236}
237
238/// Transposes in place the square matrix of scalars a tile buffer holds.
239///
240/// # Overview
241///
242/// The matrix is `2^log_n` rows of `2^log_n` scalars in row-major order.
243/// One row spans `2^LOG_TILE_PACKED` packed elements.
244/// That count is also the number of `P::WIDTH x P::WIDTH` sub-blocks along one side.
245///
246/// # Algorithm
247///
248/// The transpose factors over those sub-blocks.
249/// A sub-block of the transpose is the transpose of the sub-block mirrored across the diagonal:
250///
251/// ```text
252///     M^T[sub-block (c, r)] = (M[sub-block (r, c)])^T
253/// ```
254///
255/// Two steps therefore cover the whole job:
256///
257/// - Swap every sub-block with its mirror, which moves whole packed elements.
258/// - Transpose each sub-block in place, which is the widest transpose a packing can express.
259///
260/// # Arguments
261///
262/// * `tile` - the matrix to transpose, in row-major order
263/// * `block` - scratch space for one sub-block
264///
265/// # Preconditions
266///
267/// * `tile.len() == 1 << (P::LOG_WIDTH + 2 * LOG_TILE_PACKED)`
268/// * `block.len() == P::WIDTH`
269fn transpose_tile<P: PackedField, const LOG_TILE_PACKED: usize>(tile: &mut [P], block: &mut [P]) {
270	debug_assert_eq!(tile.len(), 1 << (P::LOG_WIDTH + 2 * LOG_TILE_PACKED));
271	debug_assert_eq!(block.len(), P::WIDTH);
272
273	// A matrix one sub-block wide is already a single square of lanes.
274	// So it needs neither the mirror swaps nor the gather below.
275	if LOG_TILE_PACKED == 0 {
276		return transpose_square_blocks(P::LOG_WIDTH, tile);
277	}
278
279	// Element holding lane row `a` of sub-block `(r, c)`.
280	// A sub-block spans `P::WIDTH` matrix rows and takes one element from each, all at column `c`.
281	let block_elem =
282		|r: usize, a: usize, c: usize| (((r << P::LOG_WIDTH) | a) << LOG_TILE_PACKED) | c;
283
284	// Step 1: swap each sub-block with its mirror across the diagonal.
285	for r in 0..1 << LOG_TILE_PACKED {
286		for c in r + 1..1 << LOG_TILE_PACKED {
287			for a in 0..P::WIDTH {
288				tile.swap(block_elem(r, a, c), block_elem(c, a, r));
289			}
290		}
291	}
292
293	// Step 2: transpose each sub-block internally.
294	// A packing of one scalar has no lanes to exchange, and the bound folds away per instance.
295	if P::LOG_WIDTH > 0 {
296		for r in 0..1 << LOG_TILE_PACKED {
297			for c in 0..1 << LOG_TILE_PACKED {
298				// The elements of one sub-block sit a whole matrix row apart.
299				// Gathering them makes the square contiguous, which is what a lane transpose takes.
300				for (a, block_i) in block.iter_mut().enumerate() {
301					*block_i = tile[block_elem(r, a, c)];
302				}
303				transpose_square_blocks(P::LOG_WIDTH, block);
304				for (a, &block_i) in block.iter().enumerate() {
305					tile[block_elem(r, a, c)] = block_i;
306				}
307			}
308		}
309	}
310}
311
312/// Applies a bit-reversal permutation to packed field elements using a simple algorithm.
313///
314/// This is a straightforward reference implementation that directly swaps field elements
315/// according to the bit-reversal permutation. It serves as a baseline for correctness
316/// testing of optimized implementations.
317///
318/// # Arguments
319///
320/// * `buffer` - Mutable slice of packed field elements to permute
321fn bit_reverse_packed_naive<P: PackedField>(mut buffer: FieldSliceMut<'_, P>) {
322	let bits = buffer.log_len() as u32;
323	for i in 0..buffer.len() {
324		let i_rev = reverse_bits(i, bits);
325		if i < i_rev {
326			let tmp = buffer.get(i);
327			buffer.set(i, buffer.get(i_rev));
328			buffer.set(i_rev, tmp);
329		}
330	}
331}
332
333/// Applies a bit-reversal permutation to elements in a slice using parallel iteration.
334///
335/// This function permutes the elements such that element at index `i` is moved to
336/// index `reverse_bits(i, log2(length))`. The permutation is performed in-place
337/// by swapping elements in parallel.
338///
339/// # Arguments
340///
341/// * `buffer` - Mutable slice of elements to permute
342///
343/// # Panics
344///
345/// Panics if the buffer length is not a power of two.
346pub fn bit_reverse_indices<T>(buffer: &mut [T]) {
347	bit_reverse_groups::<T, 0>(buffer);
348}
349
350/// Applies a bit-reversal permutation to groups of `2^LOG_GROUP` consecutive elements.
351///
352/// # Overview
353///
354/// Group `i` moves to the group whose index is `i` with its bits reversed.
355/// A group of one element permutes single elements.
356/// Wider groups permute whole cache lines instead, which is what a strided caller wants.
357///
358/// # Arguments
359///
360/// * `buffer` - Mutable slice of elements to permute
361///
362/// # Panics
363///
364/// Panics if the group count is not a power of two.
365fn bit_reverse_groups<T, const LOG_GROUP: usize>(buffer: &mut [T]) {
366	let n_groups = buffer.len() >> LOG_GROUP;
367	let bits = checked_log_2(n_groups) as u32;
368
369	// We need to use UnsafeCell-like semantics here to get proper Sync behavior.
370	// Creating a raw pointer from the slice inside the closure avoids Sync issues.
371	let buffer_ptr = buffer.as_mut_ptr() as usize;
372
373	// Half the iterations swap a pair of groups and half do nothing.
374	// So one iteration moves one group of elements on average.
375	// The cost is the memory it moves, not the index arithmetic around it.
376	let min_len = (min_len_for_bytes::<T>() >> LOG_GROUP).max(1);
377	(0..n_groups)
378		.into_par_iter()
379		.with_min_len(min_len)
380		.for_each(|i| {
381			let i_rev = reverse_bits(i, bits);
382			if i < i_rev {
383				// SAFETY: The i < i_rev condition guarantees that:
384				// 1. Each (i, i_rev) pair is processed by exactly one thread (the one with i <
385				//    i_rev)
386				// 2. Since bit-reversal is bijective, no two threads access the same pair
387				// 3. The two groups are disjoint runs of `1 << LOG_GROUP` elements
388				// 4. Both runs lie in the buffer, since their group indices are below the count
389				// 5. No data races can occur
390				// 6. buffer_ptr is valid for the lifetime of this closure
391				unsafe {
392					let ptr = buffer_ptr as *mut T;
393					let ptr_i = ptr.add(i << LOG_GROUP);
394					let ptr_i_rev = ptr.add(i_rev << LOG_GROUP);
395					ptr::swap_nonoverlapping(ptr_i, ptr_i_rev, 1 << LOG_GROUP);
396				}
397			}
398		});
399}
400
401#[cfg(test)]
402mod tests {
403	use binius_field::{Field, PackedGhash1x128b, PackedGhash2x128b, PackedGhash4x128b};
404	use proptest::prelude::*;
405	use rand::{RngExt, SeedableRng, rngs::StdRng};
406
407	use super::*;
408	use crate::{
409		FieldBuffer,
410		test_utils::{random_field_buffer, random_scalars},
411	};
412
413	// Packings of one, two and four scalars per element, at a 16-byte scalar.
414	// Each drives the tile choice down a different branch, so every property covers all three:
415	//
416	//     1 scalar  (16 B) -> tile widens to a cache line, sub-block transpose is empty
417	//     2 scalars (32 B) -> keeps its own width, tile is one sub-block
418	//     4 scalars (64 B) -> keeps its own width, tile is one sub-block
419	type P1 = PackedGhash1x128b;
420	type P2 = PackedGhash2x128b;
421	type P4 = PackedGhash4x128b;
422
423	fn check_equivalence<P: PackedField>(log_d: usize, seed: u64) {
424		// Two copies of one random buffer, so each implementation permutes the same input.
425		let mut rng = StdRng::seed_from_u64(seed);
426		let data_orig = random_field_buffer::<P>(&mut rng, log_d);
427		let mut blocked = data_orig.clone();
428		let mut naive = data_orig;
429
430		// Invariant: moving whole tiles lands every element where the definition puts it.
431		bit_reverse_packed(blocked.as_mut_view());
432		bit_reverse_packed_naive(naive.as_mut_view());
433
434		assert_eq!(blocked, naive, "mismatch at log_d={log_d}");
435	}
436
437	// Lengths chosen to straddle every branch the tile choice can take:
438	//
439	//     0, 1   -> the length leaves room for no tile at all
440	//     3      -> under the square of a four-scalar packing, so its simple path runs
441	//     4      -> tile clamped by half the length
442	//     7      -> odd length, tile still clamped by half of it
443	//     8      -> full tile, exactly one middle index
444	//     9, 13  -> odd length with middle indices left over
445	//     12     -> several middle indices
446	#[rstest::rstest]
447	#[case::single_element(0)]
448	#[case::one_bit_is_identity(1)]
449	#[case::naive_fallback_of_wide_packings(3)]
450	#[case::tile_clamped_by_length(4)]
451	#[case::odd_length_clamped_tile(7)]
452	#[case::one_middle_index(8)]
453	#[case::odd_length_full_tile(9)]
454	#[case::several_middle_indices(12)]
455	#[case::odd_length_several_middle_indices(13)]
456	fn test_bit_reverse_packed_equivalence(#[case] log_d: usize) {
457		check_equivalence::<P1>(log_d, 0);
458		check_equivalence::<P2>(log_d, 0);
459		check_equivalence::<P4>(log_d, 0);
460	}
461
462	proptest! {
463		#[test]
464		fn prop_bit_reverse_packed_matches_naive(log_d in 0..14usize, seed: u64) {
465			// Sweeps the same equivalence over random lengths and random contents.
466			check_equivalence::<P1>(log_d, seed);
467			check_equivalence::<P2>(log_d, seed);
468			check_equivalence::<P4>(log_d, seed);
469		}
470
471		#[test]
472		fn prop_bit_reverse_packed_is_an_involution(log_d in 0..14usize, seed: u64) {
473			let mut rng = StdRng::seed_from_u64(seed);
474			let orig = random_field_buffer::<P1>(&mut rng, log_d);
475
476			// Invariant: reversing the bits of an index twice is the identity.
477			// So a buffer permuted twice has to come back exactly as it went in.
478			let mut twice = orig.clone();
479			bit_reverse_packed(twice.as_mut_view());
480			bit_reverse_packed(twice.as_mut_view());
481
482			prop_assert_eq!(twice, orig);
483		}
484	}
485
486	fn transpose_reference<F: Field>(log_n: usize, scalars: &[F]) -> Vec<F> {
487		// Output row `r` is input column `r`, taking one element from each input row.
488		let n = 1 << log_n;
489		(0..n)
490			.flat_map(|r| (0..n).map(move |c| (r, c)))
491			.map(|(r, c)| scalars[(c << log_n) | r])
492			.collect()
493	}
494
495	fn check_transpose_tile<P: PackedField, const LOG_TILE_PACKED: usize>() {
496		// A tile is a square of this many scalars per side.
497		let log_n = P::LOG_WIDTH + LOG_TILE_PACKED;
498		let mut rng = StdRng::seed_from_u64(log_n as u64);
499		let scalars = random_scalars::<P::Scalar>(&mut rng, 1 << (2 * log_n));
500
501		// Pack the square, transpose it in place, then read it back out as scalars.
502		let mut tile = FieldBuffer::<P>::from_values(&scalars);
503		let mut block = zeroed_vec::<P>(P::WIDTH);
504		transpose_tile::<P, LOG_TILE_PACKED>(tile.as_mut(), &mut block);
505
506		// Invariant: exchanging sub-blocks and then lanes equals transposing scalar by scalar.
507		let expected = transpose_reference(log_n, &scalars);
508		assert_eq!(tile.iter_scalars().collect::<Vec<_>>(), expected, "mismatch at log_n={log_n}");
509	}
510
511	#[test]
512	fn test_transpose_tile_matches_scalar_transpose() {
513		// Fixture state: every tile width the dispatch can select, on every packing.
514		// A width of one element takes the single-sub-block return.
515		// Wider ones run both the mirror swaps and the per-sub-block lane transpose.
516		check_transpose_tile::<P1, 0>();
517		check_transpose_tile::<P1, 1>();
518		check_transpose_tile::<P1, 2>();
519		check_transpose_tile::<P1, 3>();
520		check_transpose_tile::<P2, 0>();
521		check_transpose_tile::<P2, 1>();
522		check_transpose_tile::<P2, 2>();
523		check_transpose_tile::<P2, 3>();
524		check_transpose_tile::<P4, 0>();
525		check_transpose_tile::<P4, 1>();
526		check_transpose_tile::<P4, 2>();
527		check_transpose_tile::<P4, 3>();
528	}
529
530	fn bit_reverse_groups_reference<T: Copy, const LOG_GROUP: usize>(buffer: &[T]) -> Vec<T> {
531		// Output group `i` is the input group whose index is `i` with its bits reversed.
532		let n_groups = buffer.len() >> LOG_GROUP;
533		let bits = checked_log_2(n_groups) as u32;
534		(0..n_groups)
535			.flat_map(|i| {
536				let src = reverse_bits(i, bits) << LOG_GROUP;
537				buffer[src..src + (1 << LOG_GROUP)].to_vec()
538			})
539			.collect()
540	}
541
542	fn check_bit_reverse_groups<const LOG_GROUP: usize>() {
543		let mut rng = StdRng::seed_from_u64(LOG_GROUP as u64);
544
545		// Group counts from one up to 32, so both the no-op and the split cases run.
546		for log_len in LOG_GROUP..LOG_GROUP + 6 {
547			let orig = (0..1usize << log_len)
548				.map(|_| rng.random::<u64>())
549				.collect::<Vec<_>>();
550
551			// Invariant: the parallel swap loop agrees with the definition, group for group.
552			let mut permuted = orig.clone();
553			bit_reverse_groups::<u64, LOG_GROUP>(&mut permuted);
554
555			assert_eq!(permuted, bit_reverse_groups_reference::<u64, LOG_GROUP>(&orig));
556		}
557	}
558
559	#[test]
560	fn test_bit_reverse_groups_matches_reference() {
561		// Fixture state: group widths of 1, 2, 4 and 8 elements.
562		// Only the first moves single elements; the rest move runs.
563		check_bit_reverse_groups::<0>();
564		check_bit_reverse_groups::<1>();
565		check_bit_reverse_groups::<2>();
566		check_bit_reverse_groups::<3>();
567	}
568
569	#[test]
570	fn test_bit_reverse_indices_is_the_single_element_group_case() {
571		let mut rng = StdRng::seed_from_u64(0);
572		let orig = (0..1usize << 10)
573			.map(|_| rng.random::<u64>())
574			.collect::<Vec<_>>();
575
576		// Invariant: permuting elements one at a time is the width-one group permutation.
577		// Fixture state: 1024 elements, so 512 index pairs are candidates for a swap.
578		let mut by_indices = orig.clone();
579		bit_reverse_indices(&mut by_indices);
580
581		assert_eq!(by_indices, bit_reverse_groups_reference::<u64, 0>(&orig));
582	}
583}