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}