1use crate::{BinaryFieldExtras, Flat, HardwareField, PackedFlat};
21use alloc::boxed::Box;
22use alloc::vec::Vec;
23use core::ops::{Add, AddAssign, Mul};
24#[cfg(feature = "parallel")]
25use rayon::prelude::*;
26
27const MAX_LEVELS: usize = 64;
28
29#[cfg(feature = "parallel")]
30const TILE: usize = 1024;
31
32#[cfg(feature = "parallel")]
33const PARALLEL_THRESHOLD_BYTES: usize = 1 << 20;
34
35#[cfg(feature = "parallel")]
36const MIN_PAR_BLOCKS: usize = 16;
37
38#[derive(Debug, Clone, Copy, PartialEq, Eq)]
40#[non_exhaustive]
41pub enum FftError {
42 BadLength { expected: usize, got: usize },
43}
44
45impl core::fmt::Display for FftError {
46 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
47 match self {
48 FftError::BadLength { expected, got } => {
49 write!(f, "AdditiveFft data length {got}, expected {expected}")
50 }
51 }
52 }
53}
54
55impl core::error::Error for FftError {}
56
57pub struct AdditiveFft<F> {
62 log_n: u32,
63
64 twiddles: Box<[Flat<F>]>,
67}
68
69impl<F: BinaryFieldExtras + HardwareField> AdditiveFft<F> {
70 pub fn new(log_n: u32) -> Self {
79 assert!(
80 (1..=F::BITS).contains(&(log_n as usize)) && log_n < usize::BITS,
81 "AdditiveFft: log_n must be in 1..=min(F::BITS, 63)"
82 );
83
84 let dim = log_n as usize;
85
86 let mut lift: Vec<Flat<F>> = Vec::with_capacity(dim - 1);
87 let mut beta = F::ONE;
88
89 for _ in 1..dim {
90 beta = F::solve_quadratic(beta).expect("field admits no Cantor basis of this size");
91 lift.push(beta.to_hardware());
92 }
93
94 let half = 1usize << (log_n - 1);
95
96 let mut twiddles = Vec::with_capacity(half);
97 for t in 0..half {
98 let mut acc = Flat::from_raw(F::ZERO);
99 let mut bits = t;
100
101 while bits != 0 {
102 let j = bits.trailing_zeros() as usize;
103 acc += lift[j];
104 bits &= bits - 1;
105 }
106
107 twiddles.push(acc);
108 }
109
110 Self {
111 log_n,
112 twiddles: twiddles.into_boxed_slice(),
113 }
114 }
115
116 pub fn forward_scalar(&self, data: &mut [Flat<F>]) -> Result<(), FftError> {
118 self.forward_coset_scalar(data, Flat::from_raw(F::ZERO))
119 }
120
121 pub fn inverse_scalar(&self, data: &mut [Flat<F>]) -> Result<(), FftError> {
123 self.inverse_coset_scalar(data, Flat::from_raw(F::ZERO))
124 }
125
126 pub fn forward_coset_scalar(
128 &self,
129 data: &mut [Flat<F>],
130 offset: Flat<F>,
131 ) -> Result<(), FftError> {
132 self.check_len(data.len())?;
133 self.fwd_levels(data, offset, fwd_butterflies);
134
135 Ok(())
136 }
137
138 pub fn inverse_coset_scalar(
140 &self,
141 data: &mut [Flat<F>],
142 offset: Flat<F>,
143 ) -> Result<(), FftError> {
144 self.check_len(data.len())?;
145 self.inv_levels(data, offset, inv_butterflies);
146
147 Ok(())
148 }
149
150 pub fn forward(&self, data: &mut [PackedFlat<F>]) -> Result<(), FftError> {
152 self.forward_coset(data, Flat::from_raw(F::ZERO))
153 }
154
155 pub fn inverse(&self, data: &mut [PackedFlat<F>]) -> Result<(), FftError> {
157 self.inverse_coset(data, Flat::from_raw(F::ZERO))
158 }
159
160 pub fn forward_coset(
162 &self,
163 data: &mut [PackedFlat<F>],
164 offset: Flat<F>,
165 ) -> Result<(), FftError> {
166 self.check_len(data.len())?;
167 self.fwd_levels(data, offset, fwd_butterflies);
168
169 Ok(())
170 }
171
172 pub fn inverse_coset(
174 &self,
175 data: &mut [PackedFlat<F>],
176 offset: Flat<F>,
177 ) -> Result<(), FftError> {
178 self.check_len(data.len())?;
179 self.inv_levels(data, offset, inv_butterflies);
180
181 Ok(())
182 }
183
184 fn check_len(&self, got: usize) -> Result<(), FftError> {
185 let expected = 1usize << self.log_n;
186 if got != expected {
187 return Err(FftError::BadLength { expected, got });
188 }
189
190 Ok(())
191 }
192
193 fn fwd_levels<T, K>(&self, data: &mut [T], offset: Flat<F>, kernel: K)
198 where
199 T: Send,
200 K: Fn(&mut [T], &mut [T], Flat<F>) + Sync,
201 {
202 let levels = self.log_n as usize;
203
204 let mut chain = [Flat::from_raw(F::ZERO); MAX_LEVELS];
205 let mut c = offset;
206
207 for slot in chain.iter_mut().take(levels) {
208 *slot = c;
209 c = c * c + c;
210 }
211
212 for l in (0..levels).rev() {
213 pass(data, &self.twiddles, chain[l], 1usize << l, &kernel);
214 }
215 }
216
217 fn inv_levels<T, K>(&self, data: &mut [T], offset: Flat<F>, kernel: K)
220 where
221 T: Send,
222 K: Fn(&mut [T], &mut [T], Flat<F>) + Sync,
223 {
224 let mut c = offset;
225 for l in 0..self.log_n as usize {
226 pass(data, &self.twiddles, c, 1usize << l, &kernel);
227 c = c * c + c;
228 }
229 }
230}
231
232fn pass<F, T, K>(data: &mut [T], twiddles: &[Flat<F>], coset: Flat<F>, s: usize, kernel: &K)
235where
236 F: HardwareField,
237 T: Send,
238 K: Fn(&mut [T], &mut [T], Flat<F>) + Sync,
239{
240 let block = 2 * s;
241 let tws = &twiddles[..data.len() / block];
242
243 #[cfg(feature = "parallel")]
244 if size_of_val(data) >= PARALLEL_THRESHOLD_BYTES {
245 if block <= TILE {
246 assert!(
247 data.len().is_multiple_of(TILE),
248 "parallel tiling drops a tail: data.len()={} not a multiple of TILE={TILE}",
249 data.len()
250 );
251
252 data.par_chunks_exact_mut(TILE)
255 .zip(tws.par_chunks_exact(TILE / block))
256 .for_each(|(span, span_tws)| blocks_serial(span, span_tws, coset, s, kernel));
257 } else if tws.len() >= MIN_PAR_BLOCKS {
258 data.par_chunks_exact_mut(block)
260 .zip(tws.par_iter())
261 .for_each(|(blk, &t)| {
262 let (lo, hi) = blk.split_at_mut(s);
263 kernel(lo, hi, coset + t);
264 });
265 } else {
266 for (blk, &t) in data.chunks_exact_mut(block).zip(tws) {
269 let tw = coset + t;
270 let (lo, hi) = blk.split_at_mut(s);
271
272 lo.par_chunks_mut(TILE)
273 .zip(hi.par_chunks_mut(TILE))
274 .for_each(|(l, h)| kernel(l, h, tw));
275 }
276 }
277
278 return;
279 }
280
281 blocks_serial(data, tws, coset, s, kernel);
282}
283
284fn blocks_serial<F, T, K>(data: &mut [T], tws: &[Flat<F>], coset: Flat<F>, s: usize, kernel: &K)
285where
286 F: HardwareField,
287 K: Fn(&mut [T], &mut [T], Flat<F>) + Sync,
288{
289 for (blk, &t) in data.chunks_exact_mut(2 * s).zip(tws) {
290 let (lo, hi) = blk.split_at_mut(s);
291 kernel(lo, hi, coset + t);
292 }
293}
294
295fn fwd_butterflies<F, T>(lo: &mut [T], hi: &mut [T], tw: Flat<F>)
296where
297 F: HardwareField,
298 T: Copy + Add<Output = T> + Mul<Flat<F>, Output = T>,
299{
300 for (p, q) in lo.iter_mut().zip(hi.iter_mut()) {
301 let qv = *q;
302 let v = *p + qv * tw;
303
304 *p = v;
305 *q = v + qv;
306 }
307}
308
309fn inv_butterflies<F, T>(lo: &mut [T], hi: &mut [T], tw: Flat<F>)
310where
311 F: HardwareField,
312 T: Copy + AddAssign + Add<Output = T> + Mul<Flat<F>, Output = T>,
313{
314 for (p, q) in lo.iter_mut().zip(hi.iter_mut()) {
315 let qv = *p + *q;
316 *p += qv * tw;
317 *q = qv;
318 }
319}