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}