mod fft;
mod filter;
mod shift;
use rustfft::num_complex::Complex;
use std::path::Path;
#[derive(Clone, Debug)]
pub struct FreqImage {
pub width: u32,
pub height: u32,
pub data: Vec<Complex<f64>>,
}
impl FreqImage {
pub fn open(path: impl AsRef<Path>) -> Result<Self, image::ImageError> {
let img = image::open(path)?;
Ok(Self::from_image(img))
}
#[must_use]
pub fn from_image(img: image::DynamicImage) -> Self {
let gray = img.into_luma8();
let (width, height) = gray.dimensions();
let data = gray
.as_raw()
.iter()
.map(|&pix| Complex::new(pix as f64 / 255.0, 0.0))
.collect();
FreqImage {
width,
height,
data,
}
}
#[must_use]
pub fn to_image(&self) -> image::GrayImage {
let pixels: Vec<u8> = self
.data
.iter()
.map(|c| (c.re.clamp(0.0, 1.0) * 255.0) as u8)
.collect();
image::GrayImage::from_raw(self.width, self.height, pixels).unwrap()
}
#[must_use]
pub fn view_fft_norm(&self) -> image::GrayImage {
let log_norms: Vec<f64> = self.data.iter().map(|c| (1.0 + c.norm()).ln()).collect();
let max = log_norms.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
let pixels: Vec<u8> = log_norms
.into_iter()
.map(|x| {
if max > 0.0 {
(x / max * 255.0) as u8
} else {
0
}
})
.collect();
image::GrayImage::from_raw(self.width, self.height, pixels).unwrap()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_roundtrip() {
for file in &["data/sjb-aerial.png", "data/mandrill.jpg"] {
let original = image::open(file).unwrap().into_luma8();
let original_pixels = original.as_raw().clone();
let mut fi = FreqImage::open(file).unwrap();
fi.fft_forward();
fi.fft_inverse();
let recovered = fi.to_image();
for (&orig, &rec) in original_pixels.iter().zip(recovered.as_raw().iter()) {
assert!(orig.abs_diff(rec) <= 1, "pixel mismatch: {orig} vs {rec}");
}
}
}
#[test]
fn test_fftshift_double_shift_is_identity() {
let fi = FreqImage::open("data/mandrill.jpg").unwrap();
let shifted = fi.fftshift();
let restored = shifted.fftshift();
for (a, b) in fi.data.iter().zip(restored.data.iter()) {
assert!((a.re - b.re).abs() < 1e-10);
}
}
#[test]
fn test_ifftshift_inverts_fftshift() {
let fi = FreqImage::open("data/mandrill.jpg").unwrap();
let shifted = fi.fftshift();
let restored = shifted.ifftshift();
for (a, b) in fi.data.iter().zip(restored.data.iter()) {
assert!((a.re - b.re).abs() < 1e-10);
}
}
#[test]
fn test_low_high_pass_masks_sum_to_one() {
let fi = FreqImage {
width: 64,
height: 64,
data: vec![Complex::default(); 64 * 64],
};
let lp = fi.low_pass_mask(0.10, 0.02);
let hp = fi.high_pass_mask(0.10, 0.02);
for (l, h) in lp.iter().zip(hp.iter()) {
assert!(
(l + h - 1.0).abs() < 1e-10,
"masks don't sum to 1: {l} + {h}"
);
}
}
#[test]
fn test_band_pass_mask_bounded_by_low_and_high() {
let fi = FreqImage {
width: 64,
height: 64,
data: vec![Complex::default(); 64 * 64],
};
let bp = fi.band_pass_mask(0.05, 0.15, 0.0);
let lp = fi.low_pass_mask(0.15, 0.0);
let hp = fi.high_pass_mask(0.05, 0.0);
for ((&b, &l), &h) in bp.iter().zip(lp.iter()).zip(hp.iter()) {
assert!(
(b - l * h).abs() < 1e-10,
"band-pass != low*high: {b} vs {}",
l * h
);
}
}
#[test]
fn test_view_fft_norm_produces_correct_size() {
let mut fi = FreqImage::open("data/mandrill.jpg").unwrap();
fi.fft_forward();
let vis = fi.view_fft_norm();
assert_eq!(vis.dimensions(), (fi.width, fi.height));
}
#[test]
fn test_open_nonexistent_returns_error() {
assert!(FreqImage::open("nonexistent.png").is_err());
}
}