fft.rs

11.4 kB · rust · 326 lines

1use std::f64::consts::PI;23/// Transforms parallel real and imaginary slices in place over a power-of-two length, unscaled in either direction.4///5/// ```6/// let (mut re, mut im) = (vec![1.0, 0.0, 0.0, 0.0], vec![0.0; 4]);7/// mrlynum::fft::fft(&mut re, &mut im, false);8/// assert_eq!(re, vec![1.0; 4]);9/// ```10pub fn fft(re: &mut [f64], im: &mut [f64], inverse: bool) {11    let n = re.len();12    assert_eq!(n, im.len(), "re and im must be equal length");13    if n <= 1 {14        return;15    }16    assert!(n.is_power_of_two(), "fft length must be a power of two");17    let mut j = 0usize;18    for i in 1..n {19        let mut bit = n >> 1;20        while j & bit != 0 {21            j ^= bit;22            bit >>= 1;23        }24        j |= bit;25        if i < j {26            re.swap(i, j);27            im.swap(i, j);28        }29    }30    let sign = if inverse { 1.0 } else { -1.0 };31    let mut len = 2;32    while len <= n {33        let ang = sign * 2.0 * PI / len as f64;34        let (wre, wim) = (ang.cos(), ang.sin());35        let half = len / 2;36        let mut start = 0;37        while start < n {38            let (mut cre, mut cim) = (1.0f64, 0.0f64);39            for k in 0..half {40                let i = start + k;41                let j = i + half;42                let (ure, uim) = (re[i], im[i]);43                let (vre, vim) = (re[j] * cre - im[j] * cim, re[j] * cim + im[j] * cre);44                re[i] = ure + vre;45                im[i] = uim + vim;46                re[j] = ure - vre;47                im[j] = uim - vim;48                let nre = cre * wre - cim * wim;49                cim = cre * wim + cim * wre;50                cre = nre;51            }52            start += len;53        }54        len <<= 1;55    }56}5758/// Transforms a size-square field in place, rows first and then columns.59pub fn fft2(re: &mut [f64], im: &mut [f64], size: usize, inverse: bool) {60    assert_eq!(re.len(), size * size, "buffer must be size*size");61    for r in 0..size {62        let s = r * size;63        fft(&mut re[s..s + size], &mut im[s..s + size], inverse);64    }65    let mut cre = vec![0.0; size];66    let mut cim = vec![0.0; size];67    for c in 0..size {68        for r in 0..size {69            cre[r] = re[r * size + c];70            cim[r] = im[r * size + c];71        }72        fft(&mut cre, &mut cim, inverse);73        for r in 0..size {74            re[r * size + c] = cre[r];75            im[r * size + c] = cim[r];76        }77    }78}7980/// Returns the magnitudes of a square field's transform, shifted so zero frequency sits at the centre.81pub fn magnitude_spectrum(field: &[f64], size: usize) -> Vec<f64> {82    let mut re = field.to_vec();83    let mut im = vec![0.0; field.len()];84    fft2(&mut re, &mut im, size, false);85    let mut out = vec![0.0; size * size];86    let half = size / 2;87    for r in 0..size {88        for c in 0..size {89            let mag = (re[r * size + c].powi(2) + im[r * size + c].powi(2)).sqrt();90            let rr = (r + half) % size;91            let cc = (c + half) % size;92            out[rr * size + cc] = mag;93        }94    }95    out96}9798/// Transforms a real size-square field forward by fft2, returning the real and imaginary parts.99pub fn transform(field: &[f64], size: usize) -> (Vec<f64>, Vec<f64>) {100    let mut re = field.to_vec();101    let mut im = vec![0.0; field.len()];102    fft2(&mut re, &mut im, size, false);103    (re, im)104}105106/// Lays an odd-side mask into a size-square kernel with the mask centre at index (0, 0) and negative offsets wrapped; the cell at offset (dr, dc) lands at (-dr, -dc) modulo size, so convolving a field by the kernel reads at every site the mask-weighted sum over its neighbours, the neighbour count the life step counts.107pub fn embed_kernel(mask: &[u8], side: usize, size: usize) -> Vec<f64> {108    assert!(side % 2 == 1, "the mask side must be odd");109    assert!(side <= size, "the mask must fit the field");110    assert_eq!(mask.len(), side * side, "mask must be side*side");111    let centre = side / 2;112    let mut out = vec![0.0; size * size];113    for r in 0..side {114        for c in 0..side {115            let value = mask[r * side + c];116            if value == 0 {117                continue;118            }119            let rr = (size + centre - r) % size;120            let cc = (size + centre - c) % size;121            out[rr * size + cc] = f64::from(value);122        }123    }124    out125}126127/// Convolves a size-square field on the torus by a kernel already transformed by fft2, the inverse scaled back by size squared.128pub fn convolve_with(field: &[f64], kernel_re: &[f64], kernel_im: &[f64], size: usize) -> Vec<f64> {129    let n = size * size;130    assert_eq!(field.len(), n, "field must be size*size");131    assert_eq!(kernel_re.len(), n, "kernel must be size*size");132    assert_eq!(kernel_im.len(), n, "kernel must be size*size");133    let (mut re, mut im) = transform(field, size);134    for ((a, b), (&kr, &ki)) in re135        .iter_mut()136        .zip(im.iter_mut())137        .zip(kernel_re.iter().zip(kernel_im))138    {139        let (fr, fi) = (*a, *b);140        *a = fr * kr - fi * ki;141        *b = fr * ki + fi * kr;142    }143    fft2(&mut re, &mut im, size, true);144    let scale = 1.0 / n as f64;145    re.iter_mut().for_each(|v| *v *= scale);146    re147}148149/// Circularly convolves a size-square field on the torus by a kernel of the same shape through fft2 both ways.150pub fn convolve(field: &[f64], kernel: &[f64], size: usize) -> Vec<f64> {151    let (kernel_re, kernel_im) = transform(kernel, size);152    convolve_with(field, &kernel_re, &kernel_im, size)153}154155/// Returns the centred magnitude spectrum of a size-square field through log(1 + magnitude), the DC bin included at the centre.156pub fn log_spectrum(field: &[f64], size: usize) -> Vec<f64> {157    magnitude_spectrum(field, size)158        .into_iter()159        .map(f64::ln_1p)160        .collect()161}162163/// Averages a centred size-square spectrum over rings of integer radius from the centre bin, a bin joining the ring its distance rounds to, rings 0 through size over two; ring k holds the frequencies near k cycles per field.164pub fn radial_profile(spectrum: &[f64], size: usize) -> Vec<f64> {165    assert_eq!(spectrum.len(), size * size, "spectrum must be size*size");166    let half = size / 2;167    let mut sums = vec![0.0; half + 1];168    let mut counts = vec![0usize; half + 1];169    for r in 0..size {170        for c in 0..size {171            let dr = r as f64 - half as f64;172            let dc = c as f64 - half as f64;173            let ring = dr.hypot(dc).round() as usize;174            if ring <= half {175                sums[ring] += spectrum[r * size + c];176                counts[ring] += 1;177            }178        }179    }180    sums.iter()181        .zip(&counts)182        .map(|(&sum, &count)| if count == 0 { 0.0 } else { sum / count as f64 })183        .collect()184}185186/// Finds the ring past the centre where a radial profile peaks, a tie broken at the smaller ring; zero when the profile holds no ring past ring 0.187pub fn peak_ring(profile: &[f64]) -> usize {188    let mut best = 0usize;189    let mut top = f64::NEG_INFINITY;190    for (ring, &value) in profile.iter().enumerate().skip(1) {191        if value > top {192            top = value;193            best = ring;194        }195    }196    best197}198199/// Reads the wavelength in cells at a radial profile's peak, size over the peak ring with a tie broken at the smaller ring; zero when the profile holds no ring past ring 0.200pub fn peak_wavelength(profile: &[f64], size: usize) -> f64 {201    match peak_ring(profile) {202        0 => 0.0,203        ring => size as f64 / ring as f64,204    }205}206207#[cfg(test)]208mod tests {209    use super::*;210    #[test]211    fn impulse_has_flat_spectrum() {212        let mut re = vec![0.0; 8];213        let mut im = vec![0.0; 8];214        re[0] = 1.0;215        fft(&mut re, &mut im, false);216        for k in 0..8 {217            let mag = (re[k].powi(2) + im[k].powi(2)).sqrt();218            assert!((mag - 1.0).abs() < 1e-9, "bin {k} mag {mag}");219        }220    }221    #[test]222    fn forward_then_inverse_round_trips() {223        let orig: Vec<f64> = (0..16).map(|i| (i as f64 * 0.7).sin()).collect();224        let mut re = orig.clone();225        let mut im = vec![0.0; 16];226        fft(&mut re, &mut im, false);227        fft(&mut re, &mut im, true);228        for (i, v) in orig.iter().enumerate() {229            assert!((re[i] / 16.0 - v).abs() < 1e-9, "index {i}");230        }231    }232    #[test]233    fn pure_sinusoid_has_two_symmetric_peaks() {234        let n = 16;235        let signal: Vec<f64> = (0..n)236            .map(|i| (2.0 * PI * 3.0 * i as f64 / n as f64).cos())237            .collect();238        let mut re = signal.clone();239        let mut im = vec![0.0; n];240        fft(&mut re, &mut im, false);241        let mag: Vec<f64> = (0..n)242            .map(|k| (re[k].powi(2) + im[k].powi(2)).sqrt())243            .collect();244        let peak = mag245            .iter()246            .enumerate()247            .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())248            .unwrap()249            .0;250        assert!(peak == 3 || peak == n - 3, "peak at {peak}");251    }252    #[test]253    fn spectrum2d_centres_dc() {254        let size = 8;255        let field = vec![2.0; size * size];256        let spec = magnitude_spectrum(&field, size);257        let centre = (size / 2) * size + size / 2;258        let max = spec.iter().cloned().fold(0.0f64, f64::max);259        assert!((spec[centre] - max).abs() < 1e-9);260        assert!(spec[centre] > 0.0);261    }262    fn moore() -> Vec<u8> {263        vec![1, 1, 1, 1, 0, 1, 1, 1, 1]264    }265    #[test]266    fn the_embedded_kernel_wraps_the_mask_about_the_origin() {267        let kernel = embed_kernel(&moore(), 3, 8);268        assert_eq!(kernel.iter().sum::<f64>(), 8.0);269        assert_eq!(kernel[0], 0.0);270        for (r, c) in [271            (0, 1),272            (0, 7),273            (1, 0),274            (7, 0),275            (1, 1),276            (1, 7),277            (7, 1),278            (7, 7),279        ] {280            assert_eq!(kernel[r * 8 + c], 1.0, "({r},{c})");281        }282    }283    #[test]284    fn an_off_centre_mask_cell_reads_the_neighbour_the_life_step_reads() {285        let kernel = embed_kernel(&[0, 1, 0, 0, 0, 0, 0, 0, 0], 3, 8);286        let mut field = vec![0.0; 64];287        field[3 * 8 + 3] = 1.0;288        let read = convolve(&field, &kernel, 8);289        for (i, v) in read.iter().enumerate() {290            let want = if i == 4 * 8 + 3 { 1.0 } else { 0.0 };291            assert!((v - want).abs() < 1e-9, "index {i} read {v}");292        }293    }294    #[test]295    fn the_moore_count_of_a_full_torus_is_eight_everywhere() {296        let kernel = embed_kernel(&moore(), 3, 16);297        let field = vec![1.0; 256];298        let counts = convolve(&field, &kernel, 16);299        assert!(counts.iter().all(|v| (v - 8.0).abs() < 1e-9));300        let (re, im) = transform(&kernel, 16);301        assert_eq!(convolve_with(&field, &re, &im, 16), counts);302    }303    #[test]304    fn the_log_spectrum_of_a_flat_field_is_one_centred_bin() {305        let spec = log_spectrum(&[2.0; 16], 4);306        assert!((spec[2 * 4 + 2] - 32f64.ln_1p()).abs() < 1e-9);307        assert!(spec308            .iter()309            .enumerate()310            .all(|(i, v)| i == 10 || v.abs() < 1e-9));311    }312    #[test]313    fn the_radial_profile_averages_rounded_rings_up_to_the_half_size() {314        let mut spec = vec![1.0; 16];315        spec[2 * 4 + 2] = 5.0;316        assert_eq!(radial_profile(&spec, 4), vec![5.0, 1.0, 1.0]);317    }318    #[test]319    fn the_peak_ring_skips_the_centre_and_breaks_ties_low() {320        assert_eq!(peak_ring(&[9.0, 1.0, 3.0, 3.0]), 2);321        assert_eq!(peak_wavelength(&[9.0, 1.0, 3.0, 3.0], 16), 8.0);322        assert_eq!(peak_ring(&[9.0]), 0);323        assert_eq!(peak_wavelength(&[9.0], 16), 0.0);324        assert_eq!(peak_ring(&[]), 0);325    }326}