Skip to main content

hekate_math/fft/
additive.rs

1// SPDX-License-Identifier: Apache-2.0
2// This file is part of the hekate-math project.
3// Copyright (C) 2026 Andrei Kochergin <andrei@oumuamua.dev>
4// Copyright (C) 2026 Oumuamua Labs <info@oumuamua.dev>.
5//
6// Licensed under the Apache License, Version 2.0 (the "License");
7// you may not use this file except in compliance with the License.
8// You may obtain a copy of the License at
9//
10//     http://www.apache.org/licenses/LICENSE-2.0
11//
12// Unless required by applicable law or agreed to in writing, software
13// distributed under the License is distributed on an "AS IS" BASIS,
14// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
15// See the License for the specific language governing permissions and
16// limitations under the License.
17
18//! Gao–Mateer additive FFT (Cantor basis).
19
20use 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/// Error returned by the additive-FFT transforms.
39#[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
57/// In-place additive FFT over a 2^log_n subspace of a binary
58/// tower field. Transforms return `Err(FftError::BadLength)`
59/// unless data.len() == 2^log_n; on success the buffer
60/// is overwritten in place.
61pub struct AdditiveFft<F> {
62    log_n: u32,
63
64    // twiddles[t] = Σ_{bit i of t} β_{i+1},
65    // flat basis.
66    twiddles: Box<[Flat<F>]>,
67}
68
69impl<F: BinaryFieldExtras + HardwareField> AdditiveFft<F> {
70    /// Derives the Cantor basis (via solve_quadratic) and
71    /// the twiddle schedule for transform size 2^log_n.
72    /// This one-time allocation is the only heap use;
73    /// the transforms are in-place.
74    ///
75    /// # Panics
76    /// If log_n is not in 1..=min(F::BITS, 63),
77    /// or F admits no Cantor basis of that size.
78    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    /// Forward: novel-basis coefficients to evaluations.
117    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    /// Inverse: evaluations to novel-basis coefficients.
122    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    /// Forward over the coset offset + W_log_n.
127    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    /// Inverse over the coset offset + W_log_n.
139    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    /// Forward, F::WIDTH column-lanes per element in lockstep.
151    pub fn forward(&self, data: &mut [PackedFlat<F>]) -> Result<(), FftError> {
152        self.forward_coset(data, Flat::from_raw(F::ZERO))
153    }
154
155    /// Inverse, F::WIDTH column-lanes per element in lockstep.
156    pub fn inverse(&self, data: &mut [PackedFlat<F>]) -> Result<(), FftError> {
157        self.inverse_coset(data, Flat::from_raw(F::ZERO))
158    }
159
160    /// Packed forward over the coset offset + W_log_n.
161    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    /// Packed inverse over the coset offset + W_log_n.
173    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    // Every depth-ℓ node shares the coset σ^ℓ(offset),
194    // σ(x) = x^2 + x; a level's butterflies tile into
195    // contiguous 2s-blocks (s = 2^ℓ), block b pairing
196    // (blk[r], blk[r+s]) with twiddle coset + twiddles[b].
197    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    // No β^-1 anywhere:
218    // paired points differ by β_0 = 1.
219    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
232// data.len() is 2^log_n (check_len), every level
233// tiles exactly; kernel gets a block's aligned halves.
234fn 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            // Many small blocks: tile-spans
253            // of whole blocks per work item.
254            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            // Mid levels: one block per work item.
259            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            // Deepest levels, too few blocks to spread:
267            // split each block's halves into aligned tiles.
268            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}