Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
114 changes: 66 additions & 48 deletions src/algorithm/raders_algorithm.rs
Original file line number Diff line number Diff line change
Expand Up @@ -42,10 +42,15 @@ pub struct RadersAlgorithm<T> {
inner_fft: Arc<dyn Fft<T>>,
inner_fft_data: Box<[Complex<T>]>,

primitive_root: u64,
primitive_root_inverse: u64,

len: StrengthReducedU64,
// Buffer offsets for the input reordering: inner FFT input k reads buffer position
// permutation[k], which is g^(k+1) mod len, less one. The sequence depends only on the

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Interesting, so twiddles generation could also be:

let (first_half, second_half) = inner_fft_input.split_at_mut((len - 1) / 2);
let mut twiddle_input = 1;
for (input_cell, conjugate_input_cell) in first_half.iter_mut().zip(second_half) {
    let twiddle = twiddles::compute_twiddle(twiddle_input, len, direction);
    *input_cell = twiddle * inner_fft_scale;
    *conjugate_input_cell = input_cell.conj();

    twiddle_input =
        ((twiddle_input as u64 * primitive_root_inverse) % reduced_len) as usize;
}

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Nice catch, thanks. Implemented this and benchmarked construction time before/after (the process() benches don't touch this since it's all in new()). Solid win across most sizes, roughly 15-28% faster for mid-range primes (149 to 65537), less at the very small end (fixed overhead dominates) and very large end (the inner FFT's own two process() calls inside the constructor start to dominate). No regressions in the existing tests. Committed as 78396b7.

// primitive root and the length, so building it once here keeps a serial chain of
// strength-reduced modular multiplies out of the hot loops. avx_raders.rs precomputes its
// output mapping for the same reason. This costs 4 * (len - 1) bytes, the same order as the
// twiddles avx_raders.rs and Bluestein's already store at this size.
permutation: Box<[u32]>,

len: usize,
inplace_scratch_len: usize,
outofplace_scratch_len: usize,
immut_scratch_len: usize,
Expand Down Expand Up @@ -86,15 +91,37 @@ impl<T: FftNum> RadersAlgorithm<T> {
// precompute the coefficients to use inside the process method
let inner_fft_scale = T::one() / T::from_usize(inner_fft_len).unwrap();
let mut inner_fft_input = vec![Complex::zero(); inner_fft_len];

// primitive_root() above only returns Some for len > 2, so len is an odd prime here and
// inner_fft_len = len - 1 is even. That makes the multiplicative group mod len cyclic of
// even order, and in a cyclic group of even order, g^(order/2) is the unique element of
// order 2, which mod a prime is -1. That holds for primitive_root_inverse just as much
// as for primitive_root, so the twiddle_input sequence e_p = primitive_root_inverse^p
// mod len satisfies e_{p + (len-1)/2} == len - e_p. Since compute_twiddle(len - x, ..)
// is the complex conjugate of compute_twiddle(x, ..), the second half of this array is
// just the conjugate of the first half: only the first half needs an actual
// compute_twiddle (trig) call and a step through the modular-multiply chain. Idea from
// https://github.com/ejmahler/RustFFT/pull/178#discussion_r3995689422
let (first_half, second_half) = inner_fft_input.split_at_mut(inner_fft_len / 2);
let mut twiddle_input = 1;
for input_cell in &mut inner_fft_input {
for (input_cell, conjugate_input_cell) in first_half.iter_mut().zip(second_half) {
let twiddle = twiddles::compute_twiddle(twiddle_input, len, direction);
*input_cell = twiddle * inner_fft_scale;
*conjugate_input_cell = input_cell.conj();

twiddle_input =
((twiddle_input as u64 * primitive_root_inverse) % reduced_len) as usize;
}

// Precompute the input reordering. Stored already offset by one so the hot loops are a
// plain indexed load with no arithmetic.
let mut permutation = Vec::with_capacity(inner_fft_len);
let mut input_index = 1u64;
for _ in 0..inner_fft_len {
input_index = (input_index * primitive_root) % reduced_len;
permutation.push(u32::try_from(input_index - 1).unwrap());
}

let required_inner_scratch = inner_fft.get_inplace_scratch_len();
let extra_inner_scratch = if required_inner_scratch <= inner_fft_len {
0
Expand All @@ -112,17 +139,41 @@ impl<T: FftNum> RadersAlgorithm<T> {
inner_fft,
inner_fft_data: inner_fft_input.into_boxed_slice(),

primitive_root,
primitive_root_inverse,
permutation: permutation.into_boxed_slice(),

len: reduced_len,
len,
inplace_scratch_len,
outofplace_scratch_len: extra_inner_scratch,
immut_scratch_len,
direction,
}
}

/// Copies `src` into `dest`, applying the input reordering.
#[inline]
fn gather_input(&self, src: &[Complex<T>], dest: &mut [Complex<T>]) {
for (dest_element, &input_index) in dest.iter_mut().zip(self.permutation.iter()) {
*dest_element = src[input_index as usize];
}
}

/// Copies `src` into `dest`, conjugating and applying the output reordering.
///
/// The output reordering walks `g^-1` where the input reordering walks `g`. Since
/// `g^(len - 1) == 1` we have `g^-k == g^(len - 1 - k)`, so the inverse sequence is the
/// forward table read backwards and rotated by one. Peeling the rotated element off keeps
/// the loop a plain zip over two slices instead of a chained iterator.
#[inline]
fn scatter_output(&self, src: &[Complex<T>], dest: &mut [Complex<T>]) {
let (&last_index, head) = self.permutation.split_last().unwrap();
let (last_element, src_head) = src.split_last().unwrap();

for (src_element, &output_index) in src_head.iter().zip(head.iter().rev()) {
dest[output_index as usize] = src_element.conj();
}
dest[last_index as usize] = last_element.conj();
}

fn perform_fft_immut(
&self,
input: &[Complex<T>],
Expand All @@ -135,13 +186,7 @@ impl<T: FftNum> RadersAlgorithm<T> {
let (scratch, extra_scratch) = scratch.split_at_mut(self.len() - 1);

// copy the input into the scratch space, reordering as we go
let mut input_index = 1;
for output_element in scratch.iter_mut() {
input_index = ((input_index as u64 * self.primitive_root) % self.len) as usize;

let input_element = input[input_index - 1];
*output_element = input_element;
}
self.gather_input(input, scratch);

self.inner_fft.process_with_scratch(scratch, extra_scratch);

Expand All @@ -163,12 +208,7 @@ impl<T: FftNum> RadersAlgorithm<T> {
self.inner_fft.process_with_scratch(scratch, extra_scratch);

// copy the final values into the output, reordering as we go
let mut output_index = 1;
for scratch_element in scratch {
output_index =
((output_index as u64 * self.primitive_root_inverse) % self.len) as usize;
output[output_index - 1] = scratch_element.conj();
}
self.scatter_output(scratch, output);
}

fn perform_fft_out_of_place(
Expand All @@ -182,13 +222,7 @@ impl<T: FftNum> RadersAlgorithm<T> {
let (input_first, input) = input.split_first_mut().unwrap();

// copy the input into the output, reordering as we go. also compute a sum of all elements
let mut input_index = 1;
for output_element in output.iter_mut() {
input_index = ((input_index as u64 * self.primitive_root) % self.len) as usize;

let input_element = input[input_index - 1];
*output_element = input_element;
}
self.gather_input(input, output);

// perform the first of two inner FFTs
let inner_scratch = if scratch.len() > 0 {
Expand Down Expand Up @@ -225,12 +259,7 @@ impl<T: FftNum> RadersAlgorithm<T> {
self.inner_fft.process_with_scratch(input, inner_scratch);

// copy the final values into the output, reordering as we go
let mut output_index = 1;
for input_element in input {
output_index =
((output_index as u64 * self.primitive_root_inverse) % self.len) as usize;
output[output_index - 1] = input_element.conj();
}
self.scatter_output(input, output);
}
fn perform_fft_inplace(&self, buffer: &mut [Complex<T>], scratch: &mut [Complex<T>]) {
// The first output element is just the sum of all the input elements, and we need to store off the first input value
Expand All @@ -240,13 +269,7 @@ impl<T: FftNum> RadersAlgorithm<T> {
let (scratch, extra_scratch) = scratch.split_at_mut(self.len() - 1);

// copy the buffer into the scratch, reordering as we go. also compute a sum of all elements
let mut input_index = 1;
for scratch_element in scratch.iter_mut() {
input_index = ((input_index as u64 * self.primitive_root) % self.len) as usize;

let buffer_element = buffer[input_index - 1];
*scratch_element = buffer_element;
}
self.gather_input(buffer, scratch);

// perform the first of two inner FFTs
let inner_scratch = if extra_scratch.len() > 0 {
Expand Down Expand Up @@ -274,17 +297,12 @@ impl<T: FftNum> RadersAlgorithm<T> {
self.inner_fft.process_with_scratch(scratch, inner_scratch);

// copy the final values into the output, reordering as we go
let mut output_index = 1;
for scratch_element in scratch {
output_index =
((output_index as u64 * self.primitive_root_inverse) % self.len) as usize;
buffer[output_index - 1] = scratch_element.conj();
}
self.scatter_output(scratch, buffer);
}
}
boilerplate_fft!(
RadersAlgorithm,
|this: &RadersAlgorithm<_>| this.len.get() as usize,
|this: &RadersAlgorithm<_>| this.len,
|this: &RadersAlgorithm<_>| this.inplace_scratch_len,
|this: &RadersAlgorithm<_>| this.outofplace_scratch_len,
|this: &RadersAlgorithm<_>| this.immut_scratch_len
Expand Down
Loading