Skip to main content

binius_math/ntt/
neighbors_last.rs

1// Copyright 2024-2025 Irreducible Inc.
2
3use std::{
4	cmp::{max, min},
5	iter,
6	ops::Range,
7	slice::from_raw_parts_mut,
8};
9
10use binius_field::{BinaryField, PackedField};
11use binius_utils::rayon::{
12	iter::{IndexedParallelIterator, IntoParallelIterator, ParallelIterator},
13	slice::ParallelSliceMut,
14};
15use itertools::izip;
16
17use super::{
18	AdditiveNTT, DomainContext,
19	reference::{NeighborsLastReference, input_check},
20};
21use crate::field_buffer::FieldSliceMut;
22
23// This value is chosen assuming 128-bit field elements.
24//
25// Empirically it performs well and is small enough for the buffer to fit comfortably in L1 cache.
26const DEFAULT_LOG_BASE_LEN: usize = 10;
27
28/// Runs a **part** of an NTT butterfly network, in depth-first order.
29///
30/// Concretely, it processes a specific memory block in the butterfly network, which is given by
31/// `layer` and `block`. For this memory block, it processes the layers given by `layer_range`.
32///
33/// For example, suppose `layer=2` and `block=2`.
34/// That means we are in an NTT butterfly network in layer 2 (the third layer) and block 2 (the
35/// third block in this layer, there are four blocks in total in this layer). `data` contains the
36/// data of this block, so it's only a chunk of the total data used in the NTT. Now suppose
37/// `layer_range=2..5`. Then we will process the following butterfly blocks:
38/// - `layer=2` `block=2`
39/// - `layer=3` `block=4`
40/// - `layer=3` `block=5`
41/// - `layer=4` `block=8`
42/// - `layer=4` `block=9`
43/// - `layer=4` `block=10`
44/// - `layer=4` `block=11`
45///
46/// (Just in a different order. We listed breadth-first order, we would process them in
47/// depth-first order.)
48///
49/// The argument `log_base_len` determines for which `log_d` we call the breadth-first
50/// implementation as a base case.
51///
52/// ## Preconditions
53///
54/// - `2^(log_d) == data.len() * packing_width`
55/// - `data.len() >= 2`
56/// - `domain_context` holds all the twiddles up to `layer_range.end` (exclusive)
57/// - `layer <= layer_range.start`
58fn forward_depth_first<P: PackedField>(
59	domain_context: &impl DomainContext<Field = P::Scalar>,
60	data: &mut [P],
61	log_d: usize,
62	layer: usize,
63	block: usize,
64	mut layer_range: Range<usize>,
65	log_base_len: usize,
66) {
67	// check preconditions
68	debug_assert!(P::LOG_WIDTH < log_d);
69	debug_assert_eq!(data.len(), 1 << (log_d - P::LOG_WIDTH));
70	debug_assert!(layer_range.end <= domain_context.log_domain_size());
71	debug_assert!(layer <= layer_range.start);
72	debug_assert!(log_base_len > P::LOG_WIDTH);
73
74	let log_n = log_d + layer;
75	debug_assert!(layer_range.end <= log_n);
76
77	if layer >= layer_range.end {
78		return;
79	}
80
81	// if the problem size is small, we just do breadth_first (to get rid of the stack overhead)
82	if log_d <= log_base_len {
83		forward_breadth_first(domain_context, data, log_d, layer, block, layer_range);
84		return;
85	}
86
87	let block_size_half = 1 << (log_d - 1 - P::LOG_WIDTH);
88	if layer >= layer_range.start {
89		// process only one layer of this block
90		let (block0, block1) = data.split_at_mut(block_size_half);
91		if block == 0 {
92			// `domain_context.twiddle(layer, 0)` is always zero (see `DomainContext::twiddle`).
93			// So the butterfly collapses to `v += u`, with `u` left unchanged.
94			for (u, v) in iter::zip(block0, block1) {
95				*v += *u;
96			}
97		} else {
98			let twiddle = domain_context.twiddle(layer, block);
99			let packed_twiddle = P::broadcast(twiddle);
100			for (u, v) in iter::zip(block0, block1) {
101				// perform butterfly
102				*u += *v * packed_twiddle;
103				*v += *u;
104			}
105		}
106
107		layer_range.start += 1;
108	}
109
110	// then recurse
111	forward_depth_first(
112		domain_context,
113		&mut data[..block_size_half],
114		log_d - 1,
115		layer + 1,
116		block << 1,
117		layer_range.clone(),
118		log_base_len,
119	);
120	forward_depth_first(
121		domain_context,
122		&mut data[block_size_half..],
123		log_d - 1,
124		layer + 1,
125		(block << 1) + 1,
126		layer_range,
127		log_base_len,
128	);
129}
130
131/// Same as [`forward_depth_first`], but runs in breadth-first order.
132///
133/// ## Preconditions
134///
135/// - `P::LOG_WIDTH < log_d`
136/// - `2^(log_d) == data.len() * packing_width`
137/// - `data.len() >= 2`
138/// - `domain_context` holds all the twiddles up to `layer_bound` (exclusive)
139/// - `layer <= layer_range.start`
140fn forward_breadth_first<P: PackedField>(
141	domain_context: &impl DomainContext<Field = P::Scalar>,
142	data: &mut [P],
143	log_d: usize,
144	base_layer: usize,
145	base_block: usize,
146	layer_range: Range<usize>,
147) {
148	// check preconditions
149	debug_assert!(P::LOG_WIDTH < log_d);
150	debug_assert_eq!(data.len(), 1 << (log_d - P::LOG_WIDTH));
151	debug_assert!(layer_range.end <= domain_context.log_domain_size());
152	debug_assert!(base_layer <= layer_range.start);
153
154	let log_n = log_d + base_layer;
155	debug_assert!(layer_range.end <= log_n);
156
157	let packed_cutoff = (log_n - P::LOG_WIDTH).clamp(layer_range.start, layer_range.end);
158
159	// In these rounds, layer <= log_n - P::LOG_WIDTH. All butterflies are between values in
160	// separate packed elements, and all butterflies within a block share the same twiddle factor.
161	for layer in layer_range.start..packed_cutoff {
162		// log_block_size is log2 the number of packed elements forming one block.
163		let log_block_size = log_n - P::LOG_WIDTH - layer;
164		let log_half_block_size = log_block_size - 1;
165
166		// log2 the number of blocks to process in this layer
167		let log_blocks = layer - base_layer;
168		let mut layer_twiddles = domain_context
169			.iter_twiddles(layer, 0)
170			.skip(base_block << log_blocks)
171			.take(1 << log_blocks);
172		let mut blocks = data.chunks_exact_mut(1 << log_block_size);
173
174		// `domain_context.twiddle(layer, 0)` is always zero (see `DomainContext::twiddle`).
175		// `base_block == 0` is the only case where this call's first block is the domain's block 0.
176		// Peel it off once per layer instead of branching per element.
177		if base_block == 0 {
178			layer_twiddles.next();
179			if let Some(block) = blocks.next() {
180				let (block0, block1) = block.split_at_mut(1 << log_half_block_size);
181				for (u, v) in iter::zip(block0, block1) {
182					*v += *u;
183				}
184			}
185		}
186
187		for (block, twiddle) in iter::zip(blocks, layer_twiddles) {
188			let packed_twiddle = P::broadcast(twiddle);
189			let (block0, block1) = block.split_at_mut(1 << log_half_block_size);
190			for (u, v) in iter::zip(block0, block1) {
191				// perform butterfly
192				*u += *v * packed_twiddle;
193				*v += *u;
194			}
195		}
196	}
197
198	// In these rounds, layer > log_n - P::LOG_WIDTH. The butterflies operate on elements within
199	// packed field elements. We solve this problem by interleaving the packed elements with each
200	// other.
201	for layer in packed_cutoff..layer_range.end {
202		// log_block_size is log2 the number of single elements forming one block.
203		let log_block_size = log_n - layer;
204		let log_half_block_size = log_block_size - 1;
205		let log_blocks_per_packed = P::LOG_WIDTH - log_block_size;
206		let log_half_blocks_per_packed = log_blocks_per_packed + 1;
207
208		// calculate packed_twiddle_offset
209		let mut packed_twiddle_offset = P::zero();
210		for block in 0..1 << log_blocks_per_packed {
211			let twiddle0 = domain_context.twiddle(layer, block);
212			let twiddle1 = domain_context.twiddle(layer, (1 << log_blocks_per_packed) | block);
213
214			let block_start = block << log_block_size;
215			for j in 0..1 << log_half_block_size {
216				packed_twiddle_offset.set(block_start | j, twiddle0);
217				packed_twiddle_offset.set(block_start | j | (1 << log_half_block_size), twiddle1);
218			}
219		}
220
221		// log2 the number of packed element pairs to process in this layer.
222		// This call's data is `2^(log_d - P::LOG_WIDTH)` packed elements, hence half that in pairs.
223		let log_packed_pairs = log_d - P::LOG_WIDTH - 1;
224		let layer_twiddles = domain_context
225			.iter_twiddles(layer, log_half_blocks_per_packed)
226			.skip(base_block << log_packed_pairs)
227			.take(1 << log_packed_pairs);
228
229		let (data_pairs, rest) = data.as_chunks_mut::<2>();
230		debug_assert!(
231			rest.is_empty(),
232			"data_packed length is a power of two; \
233				data_packed length is greater than 1 (checked at beginning of method)"
234		);
235		debug_assert_eq!(data_pairs.len(), 1 << log_packed_pairs);
236
237		for ([packed0, packed1], first_twiddle) in iter::zip(data_pairs, layer_twiddles) {
238			let packed_twiddle = P::broadcast(first_twiddle) + packed_twiddle_offset;
239
240			let (mut u, mut v) = (*packed0).interleave(*packed1, log_half_block_size);
241			u += v * packed_twiddle;
242			v += u;
243			(*packed0, *packed1) = u.interleave(v, log_half_block_size);
244		}
245	}
246}
247
248/// Process a layer of the NTT butterfly network in parallel by splitting the work up into
249/// `2^log_num_shares` many shares. This will also split up single *blocks* into multiple shares.
250///
251/// (The latter is the whole purpose of this function. If the number of shares is small enough (and
252/// the number of blocks is big enough) so that we don't need to split up blocks, we could just run
253/// [`forward_depth_first`] on disjoint chunks.)
254///
255/// - `2^(log_d) == data.len() * packing_width`
256/// - **Important:** `2^log_num_shares * 2 <= data.len()` (every share is working with whole packed
257///   elements, so every share needs at least 2 packed elements)
258/// - `domain_context` holds the twiddles of `layer`
259fn forward_shared_layer<P: PackedField>(
260	domain_context: &(impl DomainContext<Field = P::Scalar> + Sync),
261	data: &mut [P],
262	log_d: usize,
263	layer: usize,
264	log_num_shares: usize,
265) {
266	// check preconditions
267	debug_assert_eq!(data.len() * (1 << P::LOG_WIDTH), 1 << log_d);
268	debug_assert!(1 << (log_num_shares + 1) <= data.len());
269	debug_assert!(layer < domain_context.log_domain_size());
270
271	let log_num_chunks = log_num_shares + 1;
272	let log_d_chunk = log_d - log_num_chunks;
273	let data_ptr = data.as_mut_ptr();
274	let shift = log_num_shares - layer;
275	let tasks: Vec<_> = (0..1 << log_num_shares)
276		.map(|k| {
277			let (chunk0, chunk1) = with_middle_bit(k, shift);
278			let block = chunk0 >> (log_num_chunks - layer);
279			assert!(P::LOG_WIDTH <= log_d_chunk);
280			let log_chunk_len = log_d_chunk - P::LOG_WIDTH;
281			let chunk0 = unsafe {
282				from_raw_parts_mut(data_ptr.add(chunk0 << log_chunk_len), 1 << log_chunk_len)
283			};
284			let chunk1 = unsafe {
285				from_raw_parts_mut(data_ptr.add(chunk1 << log_chunk_len), 1 << log_chunk_len)
286			};
287			// `domain_context.twiddle(layer, 0)` is always zero (see `DomainContext::twiddle`).
288			// `None` signals a task whose butterfly collapses to an add, no multiply.
289			let twiddle = (block != 0).then(|| P::broadcast(domain_context.twiddle(layer, block)));
290			(chunk0, chunk1, twiddle)
291		})
292		.collect();
293
294	tasks
295		.into_par_iter()
296		.for_each(|(chunk0, chunk1, twiddle)| match twiddle {
297			Some(twiddle) => {
298				for (u, v) in iter::zip(chunk0, chunk1) {
299					butterfly(u, v, twiddle);
300				}
301			}
302			None => {
303				for (u, v) in iter::zip(chunk0, chunk1) {
304					*v += *u;
305				}
306			}
307		});
308}
309
310/// Applies one butterfly of the network: `u += v * twiddle`, then `v += u`.
311#[inline(always)]
312fn butterfly<P: PackedField>(u: &mut P, v: &mut P, twiddle: P) {
313	*u += *v * twiddle;
314	*v += *u;
315}
316
317/// Applies two consecutive butterfly layers to four planes held in registers.
318///
319/// The four planes are the quarters of one super-block, in plane order.
320/// `twiddle_0` belongs to the first layer, which pairs planes `(0, 2)` and `(1, 3)`.
321/// `twiddle_1_even` and `twiddle_1_odd` belong to the second, which pairs `(0, 1)` and `(2, 3)`.
322///
323/// Each element is loaded once, takes part in both layers, and is stored once.
324fn fused_pair<P: PackedField>(
325	planes: [&mut [P]; 4],
326	twiddle_0: P,
327	twiddle_1_even: P,
328	twiddle_1_odd: P,
329) {
330	let [plane_0, plane_1, plane_2, plane_3] = planes;
331
332	for (x_0, x_1, x_2, x_3) in izip!(plane_0, plane_1, plane_2, plane_3) {
333		butterfly(x_0, x_2, twiddle_0);
334		butterfly(x_1, x_3, twiddle_0);
335		butterfly(x_0, x_1, twiddle_1_even);
336		butterfly(x_2, x_3, twiddle_1_odd);
337	}
338}
339
340/// Same as [`fused_pair`], for the super-block whose index is zero.
341///
342/// There `twiddle_0` and `twiddle_1_even` are both the layer's block-0 twiddle, which is zero.
343/// Each of their butterflies collapses to `v += u`, leaving `u` untouched.
344/// So plane 0 is never written, and its cache lines stay clean.
345fn fused_pair_zero_block<P: PackedField>(planes: [&mut [P]; 4], twiddle_1_odd: P) {
346	let [plane_0, plane_1, plane_2, plane_3] = planes;
347
348	for (x_0, x_1, x_2, x_3) in izip!(plane_0, plane_1, plane_2, plane_3) {
349		*x_2 += *x_0;
350		*x_3 += *x_1;
351		*x_1 += *x_0;
352		butterfly(x_2, x_3, twiddle_1_odd);
353	}
354}
355
356/// Processes layers `first_layer` and `first_layer + 1` in a single pass over the buffer.
357///
358/// Index the buffer by `i` in `[0, 2^log_d)`.
359/// Layer `l` pairs `i` with `i XOR 2^(log_d - 1 - l)`, under the twiddle of block `i >> (log_d -
360/// l)`.
361///
362/// Writing `a` for `first_layer`, split the index three ways:
363///
364/// ```text
365///     i = h * 2^(log_d - a)  +  m * 2^(log_d - a - 2)  +  offset
366///
367///     h < 2^a                    super-block index, which the two layers never move
368///     m < 4                      plane index, the only bits they do move
369///     offset < 2^(log_d - a - 2) position inside a plane, which they never move
370/// ```
371///
372/// The four entries sharing one `(h, offset)` are closed under both layers, and distinct pairs
373/// never interact.
374/// So a quarter of the buffer's elements can be transformed independently of the rest, which is
375/// what lets both layers run without a barrier between them.
376///
377/// Layer `a` reads the twiddle of block `h`, shared by both of its pairs.
378/// Layer `a + 1` reads blocks `2h` and `2h + 1`, one per pair.
379///
380/// ## Preconditions
381///
382/// - `2^log_d == data.len() * packing_width`
383/// - `first_layer + 2 <= log_num_shares`
384/// - `first_layer + 2 + P::LOG_WIDTH < log_d`
385/// - `domain_context` holds the twiddles of layers `first_layer` and `first_layer + 1`
386fn forward_shared_layer_pair<P: PackedField>(
387	domain_context: &(impl DomainContext<Field = P::Scalar> + Sync),
388	data: &mut [P],
389	log_d: usize,
390	first_layer: usize,
391	log_num_shares: usize,
392) {
393	// check preconditions
394	debug_assert_eq!(data.len() << P::LOG_WIDTH, 1 << log_d);
395	debug_assert!(first_layer + 2 <= log_num_shares);
396	debug_assert!(first_layer + 2 + P::LOG_WIDTH < log_d);
397	debug_assert!(first_layer + 2 <= domain_context.log_domain_size());
398
399	let log_plane_len = log_d - first_layer - 2 - P::LOG_WIDTH;
400	let plane_len = 1 << log_plane_len;
401
402	// Cut every plane into equal runs, enough of them to feed each share at least one.
403	// A run never drops below one packed element.
404	let log_run_len = log_plane_len.saturating_sub(log_num_shares - first_layer);
405	let run_len = 1 << log_run_len;
406
407	// One task per (super-block, run) pair, which the borrow checker proves disjoint for us.
408	let tasks = data
409		.chunks_exact_mut(plane_len << 2)
410		.enumerate()
411		.flat_map(|(h, super_block)| {
412			let twiddle = |layer, block| P::broadcast(domain_context.twiddle(layer, block));
413			let twiddles = (
414				twiddle(first_layer, h),
415				twiddle(first_layer + 1, h << 1),
416				twiddle(first_layer + 1, (h << 1) | 1),
417			);
418
419			let (halves_0_1, halves_2_3) = super_block.split_at_mut(plane_len << 1);
420			let (plane_0, plane_1) = halves_0_1.split_at_mut(plane_len);
421			let (plane_2, plane_3) = halves_2_3.split_at_mut(plane_len);
422
423			izip!(
424				plane_0.chunks_exact_mut(run_len),
425				plane_1.chunks_exact_mut(run_len),
426				plane_2.chunks_exact_mut(run_len),
427				plane_3.chunks_exact_mut(run_len),
428			)
429			.map(move |(run_0, run_1, run_2, run_3)| ([run_0, run_1, run_2, run_3], h, twiddles))
430		})
431		.collect::<Vec<_>>();
432
433	tasks
434		.into_par_iter()
435		.for_each(|(planes, h, (twiddle_0, twiddle_1_even, twiddle_1_odd))| {
436			// `domain_context.twiddle(layer, 0)` is always zero (see `DomainContext::twiddle`).
437			// Only super-block 0 reads block 0, and it does so in both of its layers.
438			if h == 0 {
439				fused_pair_zero_block(planes, twiddle_1_odd);
440			} else {
441				fused_pair(planes, twiddle_0, twiddle_1_even, twiddle_1_odd);
442			}
443		});
444}
445
446/// Runs the shared layers of the butterfly network, two layers per pass where possible.
447///
448/// Pairing halves the passes a window of layers costs, so it halves its memory traffic.
449/// A layer that cannot be paired -- an odd one out, or a shape whose planes would fall below one
450/// packed element -- runs alone, exactly as it did before pairing existed.
451///
452/// ## Preconditions
453///
454/// - same as the routines this dispatches to
455/// - `layers.end <= log_num_shares`
456fn forward_shared_layers<P: PackedField>(
457	domain_context: &(impl DomainContext<Field = P::Scalar> + Sync),
458	data: &mut [P],
459	log_d: usize,
460	layers: Range<usize>,
461	log_num_shares: usize,
462) {
463	debug_assert!(layers.end <= log_num_shares);
464
465	let mut layer = layers.start;
466	while layer < layers.end {
467		// A pair needs one more layer to pair with, and planes of at least one packed element.
468		if layers.end - layer >= 2 && layer + 2 + P::LOG_WIDTH < log_d {
469			forward_shared_layer_pair(domain_context, data, log_d, layer, log_num_shares);
470			layer += 2;
471		} else {
472			forward_shared_layer(domain_context, data, log_d, layer, log_num_shares);
473			layer += 1;
474		}
475	}
476}
477
478/// Inserts a bit into `k`. Returns both the version with `0` inserted and `1` inserted.
479///
480/// The first `shift` bits are preserved, then `0` or `1` is inserted, and then the remaining bits
481/// of `k` follow.
482///
483/// ## Preconditions
484///
485/// - `shift` must be strictly greater than 0
486fn with_middle_bit(k: usize, shift: usize) -> (usize, usize) {
487	assert!(shift >= 1);
488
489	// most significant and least significant bits, overlapping in one bit
490	let ms = k >> (shift - 1);
491	let ls = k & ((1 << shift) - 1);
492
493	let k0 = ls | ((ms & !1) << shift);
494	let k1 = ls | ((ms | 1) << shift);
495
496	(k0, k1)
497}
498
499#[derive(Debug)]
500pub struct NeighborsLastBreadthFirst<DC> {
501	/// The domain context from which the twiddles are pulled.
502	pub domain_context: DC,
503}
504
505impl<F, DC> AdditiveNTT for NeighborsLastBreadthFirst<DC>
506where
507	F: BinaryField,
508	DC: DomainContext<Field = F>,
509{
510	type Field = F;
511
512	fn forward_transform<P: PackedField<Scalar = F>>(
513		&self,
514		mut data: FieldSliceMut<'_, P>,
515		skip_early: usize,
516		skip_late: usize,
517	) {
518		let log_d = data.log_len();
519		if log_d <= P::LOG_WIDTH {
520			let fallback_ntt = NeighborsLastReference {
521				domain_context: &self.domain_context,
522			};
523			return fallback_ntt.forward_transform(data, skip_early, skip_late);
524		}
525
526		input_check(&self.domain_context, log_d, skip_early, skip_late);
527
528		forward_breadth_first(
529			self.domain_context(),
530			data.as_mut(),
531			log_d,
532			0,
533			0,
534			skip_early..(log_d - skip_late),
535		);
536	}
537
538	fn inverse_transform<P: PackedField<Scalar = F>>(
539		&self,
540		_data: FieldSliceMut<'_, P>,
541		_skip_early: usize,
542		_skip_late: usize,
543	) {
544		todo!()
545	}
546
547	fn domain_context(&self) -> &impl DomainContext<Field = F> {
548		&self.domain_context
549	}
550}
551
552/// A single-threaded implementation of [`AdditiveNTT`].
553///
554/// The code only makes sure that it's fast for a _large_ data input.
555/// For small inputs, it can be comparatively slow!
556///
557/// The implementation is depth-first, but calls a breadth-first implementation as a base case.
558///
559/// Note that "neighbors last" refers to the memory layout for the NTT: In the _last_ layer of this
560/// NTT algorithm, neighboring elements speak to each other. In the classic FFT that's usually the
561/// case for "decimation in frequency".
562#[derive(Debug)]
563pub struct NeighborsLastSingleThread<DC> {
564	/// The domain context from which the twiddles are pulled.
565	pub domain_context: DC,
566	/// Determines when to switch from depth-first to the breadth-first base case.
567	pub log_base_len: usize,
568}
569
570impl<DC> NeighborsLastSingleThread<DC> {
571	/// Convenience constructor which sets `log_base_len` to a reasonable default.
572	pub const fn new(domain_context: DC) -> Self {
573		Self {
574			domain_context,
575			log_base_len: DEFAULT_LOG_BASE_LEN,
576		}
577	}
578}
579
580impl<DC: DomainContext> AdditiveNTT for NeighborsLastSingleThread<DC> {
581	type Field = DC::Field;
582
583	fn forward_transform<P: PackedField<Scalar = Self::Field>>(
584		&self,
585		mut data: FieldSliceMut<'_, P>,
586		skip_early: usize,
587		skip_late: usize,
588	) {
589		let log_d = data.log_len();
590		if log_d <= P::LOG_WIDTH {
591			let fallback_ntt = NeighborsLastReference {
592				domain_context: &self.domain_context,
593			};
594			return fallback_ntt.forward_transform(data, skip_early, skip_late);
595		}
596
597		input_check(&self.domain_context, log_d, skip_early, skip_late);
598
599		forward_depth_first(
600			&self.domain_context,
601			data.as_mut(),
602			log_d,
603			0,
604			0,
605			skip_early..(log_d - skip_late),
606			// Ensures that log_base_len satisfies precondition
607			self.log_base_len.max(P::LOG_WIDTH + 1),
608		);
609	}
610
611	fn inverse_transform<P: PackedField<Scalar = Self::Field>>(
612		&self,
613		_data_orig: FieldSliceMut<'_, P>,
614		_skip_early: usize,
615		_skip_late: usize,
616	) {
617		unimplemented!()
618	}
619
620	fn domain_context(&self) -> &impl DomainContext<Field = DC::Field> {
621		&self.domain_context
622	}
623}
624
625/// A multi-threaded implementation of [`AdditiveNTT`].
626///
627/// The code only makes sure that it's fast for a _large_ data input.
628/// For small inputs, it can be comparatively slow!
629///
630/// The implementation is depth-first, but calls a breadth-first implementation as a base case.
631///
632/// Note that "neighbors last" refers to the memory layout for the NTT: In the _last_ layer of this
633/// NTT algorithm, neighboring elements speak to each other. In the classic FFT that's usually the
634/// case for "decimation in frequency".
635#[derive(Debug)]
636pub struct NeighborsLastMultiThread<DC> {
637	/// The domain context from which the twiddles are pulled.
638	pub domain_context: DC,
639	/// Determines when to switch from depth-first to the breadth-first base case.
640	pub log_base_len: usize,
641	/// The base-2 logarithm of number of equal-sized shares that the problem should be split into.
642	/// Each share needs to do the same amount of work. If you have equally powered cores
643	/// available, this should be the base-2 logarithm of the number of cores.
644	pub log_num_shares: usize,
645}
646
647impl<DC> NeighborsLastMultiThread<DC> {
648	/// Convenience constructor which sets `log_base_len` to a reasonable default.
649	pub const fn new(domain_context: DC, log_num_shares: usize) -> Self {
650		Self {
651			domain_context,
652			log_base_len: DEFAULT_LOG_BASE_LEN,
653			log_num_shares,
654		}
655	}
656}
657
658impl<DC: DomainContext + Sync> AdditiveNTT for NeighborsLastMultiThread<DC> {
659	type Field = DC::Field;
660
661	fn forward_transform<P: PackedField<Scalar = Self::Field>>(
662		&self,
663		mut data: FieldSliceMut<'_, P>,
664		skip_early: usize,
665		skip_late: usize,
666	) {
667		let log_d = data.log_len();
668		if log_d <= P::LOG_WIDTH {
669			let fallback_ntt = NeighborsLastReference {
670				domain_context: &self.domain_context,
671			};
672			return fallback_ntt.forward_transform(data, skip_early, skip_late);
673		}
674
675		input_check(&self.domain_context, log_d, skip_early, skip_late);
676
677		// Decide on `actual_log_num_shares`, which also determines how many shared rounds we do.
678		// By default this would just be `self.log_num_shares`, but we will potentially decrease it
679		// in order to make sure that `2^log_num_shares * 2 <= data.len()`. This serves two
680		// purposes:
681		// - when we do the shared rounds, each thread should have at least 2 packed elements to
682		//   work with, see the precondition of [`forward_shared_layer`]
683		// - when we do the independent rounds, again each share should have `chunk.len() >= 2`
684		//   because this is required by [`forward_depth_first`]
685		let maximum_log_num_shares = log_d - P::LOG_WIDTH - 1;
686		let actual_log_num_shares = min(self.log_num_shares, maximum_log_num_shares);
687		let first_independent_layer = actual_log_num_shares;
688
689		let last_layer = log_d - skip_late;
690		let shared_layers = skip_early..min(first_independent_layer, last_layer);
691		let independent_layers = max(first_independent_layer, skip_early)..last_layer;
692
693		forward_shared_layers(
694			&self.domain_context,
695			data.as_mut(),
696			log_d,
697			shared_layers,
698			actual_log_num_shares,
699		);
700
701		// One might think that we could just call `forward_depth_first` with
702		// `layer=independent_layers.start`. However, this would mean that the chunk size (that we
703		// split into using `par_chunks_mut`) could be just one packed element, or even less than
704		// one packed element.
705		let layer = min(independent_layers.start, maximum_log_num_shares);
706		let log_d_chunk = log_d - layer;
707		data.as_mut()
708			.par_chunks_exact_mut(1 << (log_d_chunk - P::LOG_WIDTH))
709			.enumerate()
710			.for_each(|(block, chunk)| {
711				forward_depth_first(
712					&self.domain_context,
713					chunk,
714					log_d_chunk,
715					layer,
716					block,
717					independent_layers.clone(),
718					self.log_base_len,
719				);
720			});
721	}
722
723	fn inverse_transform<P: PackedField<Scalar = Self::Field>>(
724		&self,
725		_data_orig: FieldSliceMut<'_, P>,
726		_skip_early: usize,
727		_skip_late: usize,
728	) {
729		unimplemented!()
730	}
731
732	fn domain_context(&self) -> &impl DomainContext<Field = DC::Field> {
733		&self.domain_context
734	}
735}
736
737#[cfg(test)]
738mod tests {
739	use binius_field::{PackedGhash1x128b, PackedGhash2x128b, PackedGhash4x128b};
740	use proptest::prelude::*;
741	use rand::{SeedableRng, rngs::StdRng};
742
743	use super::*;
744	use crate::{ntt::domain_context::GaoMateerPreExpanded, test_utils::random_field_buffer};
745
746	/// The shared-layer window the multithreaded transform would pick for these parameters.
747	///
748	/// Returns the window together with the share count after the transform's own clamp.
749	/// An empty window means the transform runs no shared layers at all.
750	fn shared_window<P: PackedField>(
751		log_d: usize,
752		skip_early: usize,
753		skip_late: usize,
754		log_num_shares: usize,
755	) -> (Range<usize>, usize) {
756		// Every share needs two packed elements, which caps how finely the buffer splits.
757		let actual_log_num_shares = min(log_num_shares, log_d - P::LOG_WIDTH - 1);
758		let layers = skip_early..min(actual_log_num_shares, log_d - skip_late);
759		(layers, actual_log_num_shares)
760	}
761
762	/// Asserts the fused dispatcher and the one-pass-per-layer loop agree bit for bit.
763	fn check_fused_matches_per_layer<P: PackedField<Scalar: BinaryField>>(
764		log_d: usize,
765		skip_early: usize,
766		skip_late: usize,
767		log_num_shares: usize,
768		seed: u64,
769	) {
770		let (layers, actual_log_num_shares) =
771			shared_window::<P>(log_d, skip_early, skip_late, log_num_shares);
772		if layers.is_empty() {
773			return;
774		}
775
776		let domain_context = GaoMateerPreExpanded::<P::Scalar>::generate(log_d);
777		let mut rng = StdRng::seed_from_u64(seed);
778
779		// Same random contents down both paths, so any difference is the fusion's fault.
780		let mut fused = random_field_buffer::<P>(&mut rng, log_d);
781		let mut per_layer = fused.clone();
782
783		forward_shared_layers(
784			&domain_context,
785			fused.as_mut(),
786			log_d,
787			layers.clone(),
788			actual_log_num_shares,
789		);
790		for layer in layers {
791			forward_shared_layer(
792				&domain_context,
793				per_layer.as_mut(),
794				log_d,
795				layer,
796				actual_log_num_shares,
797			);
798		}
799
800		assert_eq!(fused, per_layer);
801	}
802
803	/// Asserts the whole multithreaded transform matches the independent reference transform.
804	fn check_transform_matches_reference<P: PackedField<Scalar: BinaryField>>(
805		log_d: usize,
806		skip_early: usize,
807		skip_late: usize,
808		log_num_shares: usize,
809		seed: u64,
810	) {
811		let domain_context = GaoMateerPreExpanded::<P::Scalar>::generate(log_d);
812		let mut rng = StdRng::seed_from_u64(seed);
813
814		let mut actual = random_field_buffer::<P>(&mut rng, log_d);
815		let mut expected = actual.clone();
816
817		let multi_thread = NeighborsLastMultiThread {
818			domain_context: &domain_context,
819			log_base_len: 3,
820			log_num_shares,
821		};
822		let reference = NeighborsLastReference {
823			domain_context: &domain_context,
824		};
825
826		multi_thread.forward_transform(actual.as_mut_view(), skip_early, skip_late);
827		reference.forward_transform(expected.as_mut_view(), skip_early, skip_late);
828
829		assert_eq!(actual, expected);
830	}
831
832	/// Every window width from one shared layer up to five, at every start offset that fits.
833	///
834	/// Width 1 has nothing to fuse, widths 2 to 4 fuse in one window, width 5 splits into 4 plus 1.
835	/// The start offset is what shifts the super-block count, so it must move too.
836	fn sweep_window_widths<P: PackedField<Scalar: BinaryField>>(log_d: usize, seed: u64) {
837		for skip_early in 0..4 {
838			for width in 1..=5 {
839				// Layer indices only exist below the buffer's own depth.
840				if skip_early + width > log_d {
841					continue;
842				}
843				// Ending the shared phase at `skip_early + width` is what sets the window width.
844				check_fused_matches_per_layer::<P>(log_d, skip_early, 0, skip_early + width, seed);
845			}
846		}
847	}
848
849	#[test]
850	fn fused_shared_layers_match_per_layer_over_window_widths() {
851		// Fixture: 128-bit scalars at three packing widths.
852		//
853		//     1x128b -> LOG_WIDTH 0, planes are always whole packed elements
854		//     2x128b -> LOG_WIDTH 1
855		//     4x128b -> LOG_WIDTH 2, the width that first refuses the widest windows
856		for log_d in 6..13 {
857			sweep_window_widths::<PackedGhash1x128b>(log_d, 0);
858			sweep_window_widths::<PackedGhash2x128b>(log_d, 1);
859			sweep_window_widths::<PackedGhash4x128b>(log_d, 2);
860		}
861	}
862
863	#[test]
864	fn fused_shared_layers_match_per_layer_at_the_narrowest_plane() {
865		// Invariant: a window is fused only while each plane still holds two packed elements.
866		//
867		//     log_d = 8, LOG_WIDTH = 2, skip_early = 1, width = 4
868		//     plane length = 2^(8 - 1 - 4 - 2) = 2 packed elements, the smallest fusable plane
869		check_fused_matches_per_layer::<PackedGhash4x128b>(8, 1, 0, 5, 7);
870
871		// One layer deeper the plane would hold a single packed element, so this shape falls back.
872		//
873		//     log_d = 8, LOG_WIDTH = 2, skip_early = 0, width = 5 -> plane length 2^1
874		//     the same start with width 6 would need plane length 2^0, refused
875		check_fused_matches_per_layer::<PackedGhash4x128b>(8, 0, 0, 5, 8);
876	}
877
878	#[test]
879	fn multi_thread_transform_matches_the_reference() {
880		// Sweep the skips against the share count, so the shared phase starts and ends everywhere.
881		for log_d in [6, 9, 12] {
882			for skip_early in [0, 1, 2, 4] {
883				for skip_late in [0, 1, 3] {
884					if skip_early + skip_late > log_d {
885						continue;
886					}
887					for log_num_shares in [0, 1, 2, 3, 5, 1000] {
888						check_transform_matches_reference::<PackedGhash1x128b>(
889							log_d,
890							skip_early,
891							skip_late,
892							log_num_shares,
893							0,
894						);
895						check_transform_matches_reference::<PackedGhash4x128b>(
896							log_d,
897							skip_early,
898							skip_late,
899							log_num_shares,
900							1,
901						);
902					}
903				}
904			}
905		}
906	}
907
908	proptest! {
909		#[test]
910		fn prop_fused_shared_layers_match_per_layer(
911			log_d in 6..13usize,
912			skip_early in 0..5usize,
913			skip_late in 0..4usize,
914			log_num_shares in 0..8usize,
915			seed: u64,
916		) {
917			// Property: fusing a window of layers is the same linear map as running them one by
918			// one, for every shape the multithreaded transform can hand the shared phase.
919			prop_assume!(skip_early + skip_late <= log_d);
920
921			check_fused_matches_per_layer::<PackedGhash1x128b>(
922				log_d, skip_early, skip_late, log_num_shares, seed,
923			);
924			check_fused_matches_per_layer::<PackedGhash2x128b>(
925				log_d, skip_early, skip_late, log_num_shares, seed,
926			);
927			check_fused_matches_per_layer::<PackedGhash4x128b>(
928				log_d, skip_early, skip_late, log_num_shares, seed,
929			);
930		}
931
932		#[test]
933		fn prop_multi_thread_transform_matches_the_reference(
934			log_d in 1..11usize,
935			skip_early in 0..5usize,
936			skip_late in 0..4usize,
937			log_num_shares in 0..8usize,
938			seed: u64,
939		) {
940			// Property: the fused shared phase leaves the transform equal to the reference one,
941			// which shares no code with it.
942			prop_assume!(skip_early + skip_late <= log_d);
943
944			check_transform_matches_reference::<PackedGhash1x128b>(
945				log_d, skip_early, skip_late, log_num_shares, seed,
946			);
947			check_transform_matches_reference::<PackedGhash4x128b>(
948				log_d, skip_early, skip_late, log_num_shares, seed,
949			);
950		}
951	}
952
953	#[test]
954	fn test_with_middle_bit() {
955		assert_eq!(with_middle_bit(0b000, 1), (0b0000, 0b0010));
956		assert_eq!(with_middle_bit(0b000, 2), (0b0000, 0b0100));
957		assert_eq!(with_middle_bit(0b000, 3), (0b0000, 0b1000));
958
959		assert_eq!(with_middle_bit(0b111, 1), (0b1101, 0b1111));
960		assert_eq!(with_middle_bit(0b111, 2), (0b1011, 0b1111));
961		assert_eq!(with_middle_bit(0b111, 3), (0b0111, 0b1111));
962
963		assert_eq!(with_middle_bit(0b1110110, 2), (0b11101010, 0b11101110));
964	}
965}