#![forbid(unsafe_code)]
use oxifft::api::{Direction, Flags, Plan};
use oxifft::Complex;
#[derive(Clone, Debug)]
pub struct PitchDetection {
pub frequency_hz: f32,
pub confidence: f32,
pub midi_note: u8,
pub cents_deviation: f32,
pub note_name: &'static str,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum PitchAlgorithm {
Yin,
Autocorrelation,
Cepstrum,
Amdf,
}
#[derive(Clone, Debug)]
pub struct PitchDetector {
pub algorithm: PitchAlgorithm,
pub sample_rate: u32,
pub min_frequency: f32,
pub max_frequency: f32,
pub threshold: f32,
}
static NOTE_NAMES: &[&str] = &[
"C-1", "C#-1", "D-1", "D#-1", "E-1", "F-1", "F#-1", "G-1", "G#-1", "A-1", "A#-1", "B-1", "C0",
"C#0", "D0", "D#0", "E0", "F0", "F#0", "G0", "G#0", "A0", "A#0", "B0", "C1", "C#1", "D1",
"D#1", "E1", "F1", "F#1", "G1", "G#1", "A1", "A#1", "B1", "C2", "C#2", "D2", "D#2", "E2", "F2",
"F#2", "G2", "G#2", "A2", "A#2", "B2", "C3", "C#3", "D3", "D#3", "E3", "F3", "F#3", "G3",
"G#3", "A3", "A#3", "B3", "C4", "C#4", "D4", "D#4", "E4", "F4", "F#4", "G4", "G#4", "A4",
"A#4", "B4", "C5", "C#5", "D5", "D#5", "E5", "F5", "F#5", "G5", "G#5", "A5", "A#5", "B5", "C6",
"C#6", "D6", "D#6", "E6", "F6", "F#6", "G6", "G#6", "A6", "A#6", "B6", "C7", "C#7", "D7",
"D#7", "E7", "F7", "F#7", "G7", "G#7", "A7", "A#7", "B7", "C8", "C#8", "D8", "D#8", "E8", "F8",
"F#8", "G8", "G#8", "A8", "A#8", "B8", "C9", "C#9", "D9", "D#9", "E9", "F9", "F#9", "G9",
];
#[must_use]
pub fn frequency_to_midi(freq: f32) -> (u8, f32) {
if freq <= 0.0 {
return (0, 0.0);
}
let fractional_midi = 69.0 + 12.0 * (freq / 440.0).log2();
let nearest = fractional_midi.round().clamp(0.0, 127.0);
let cents = (fractional_midi - nearest) * 100.0;
(nearest as u8, cents)
}
#[must_use]
pub fn midi_to_note_name(midi: u8) -> &'static str {
let idx = (midi as usize).min(NOTE_NAMES.len() - 1);
NOTE_NAMES[idx]
}
#[must_use]
pub fn frequency_to_note_name(freq: f32) -> (&'static str, f32) {
let (midi, cents) = frequency_to_midi(freq);
(midi_to_note_name(midi), cents)
}
impl PitchDetector {
#[must_use]
pub fn new(algorithm: PitchAlgorithm, sample_rate: u32) -> Self {
let threshold = match algorithm {
PitchAlgorithm::Yin | PitchAlgorithm::Amdf => 0.10,
PitchAlgorithm::Autocorrelation | PitchAlgorithm::Cepstrum => 0.15,
};
Self {
algorithm,
sample_rate,
min_frequency: 50.0,
max_frequency: 1000.0,
threshold,
}
}
#[must_use]
pub fn detect(&self, samples: &[f32]) -> Option<PitchDetection> {
if samples.is_empty() {
return None;
}
let sr = self.sample_rate as f32;
let min_lag = (sr / self.max_frequency).ceil() as usize;
let max_lag = (sr / self.min_frequency).floor() as usize;
if max_lag >= samples.len() || min_lag == 0 || min_lag > max_lag {
return None;
}
let (period, confidence) = match self.algorithm {
PitchAlgorithm::Yin => detect_yin(samples, min_lag, max_lag, self.threshold)?,
PitchAlgorithm::Autocorrelation => {
detect_autocorrelation(samples, min_lag, max_lag, self.threshold)?
}
PitchAlgorithm::Amdf => detect_amdf(samples, min_lag, max_lag, self.threshold)?,
PitchAlgorithm::Cepstrum => detect_cepstrum(samples, min_lag, max_lag, self.threshold)?,
};
if period == 0 {
return None;
}
let frequency_hz = sr / period as f32;
let (midi_note, cents_deviation) = frequency_to_midi(frequency_hz);
let note_name = midi_to_note_name(midi_note);
Some(PitchDetection {
frequency_hz,
confidence,
midi_note,
cents_deviation,
note_name,
})
}
}
fn detect_yin(
samples: &[f32],
min_lag: usize,
max_lag: usize,
threshold: f32,
) -> Option<(usize, f32)> {
let n = samples.len();
let max_tau = max_lag.min(n / 2);
if min_lag >= max_tau {
return None;
}
let window = max_tau;
let mut d: Vec<f32> = vec![0.0; window + 1];
for tau in 1..=window {
let mut val = 0.0f32;
for t in 0..(n - tau) {
let diff = samples[t] - samples[t + tau];
val += diff * diff;
}
d[tau] = val;
}
let mut cmndf: Vec<f32> = vec![0.0; window + 1];
cmndf[0] = 1.0;
let mut running_sum = 0.0f32;
for tau in 1..=window {
running_sum += d[tau];
if running_sum > 0.0 {
cmndf[tau] = d[tau] * tau as f32 / running_sum;
} else {
cmndf[tau] = 1.0;
}
}
let search_end = max_lag.min(window);
let mut best_tau = 0usize;
let mut best_val = f32::MAX;
let mut tau = min_lag;
while tau < search_end {
if cmndf[tau] < threshold {
let dip_start = tau;
while tau + 1 < search_end && cmndf[tau + 1] < cmndf[tau] {
tau += 1;
}
if cmndf[tau] < best_val {
best_val = cmndf[tau];
best_tau = tau;
}
if best_val < threshold {
break;
}
tau = dip_start + 1;
} else {
tau += 1;
}
}
if best_tau == 0 {
for t in min_lag..=search_end {
if cmndf[t] < best_val {
best_val = cmndf[t];
best_tau = t;
}
}
}
if best_tau == 0 || best_val > 0.9 {
return None;
}
let confidence = (1.0 - best_val).clamp(0.0, 1.0);
Some((best_tau, confidence))
}
fn detect_autocorrelation(
samples: &[f32],
min_lag: usize,
max_lag: usize,
threshold: f32,
) -> Option<(usize, f32)> {
let n = samples.len();
let r0: f32 = samples.iter().map(|&s| s * s).sum();
if r0 < 1e-10 {
return None; }
let upper = max_lag.min(n / 2);
let mut acf: Vec<f32> = vec![0.0; upper + 1];
acf[0] = 1.0;
for lag in 1..=upper {
let len = n - lag;
let r: f32 = (0..len).map(|i| samples[i] * samples[i + lag]).sum();
acf[lag] = r / r0;
}
let mut best_lag = 0usize;
let mut best_val = threshold;
let mut lag = min_lag;
while lag < upper {
let prev = if lag > 0 { acf[lag - 1] } else { 0.0 };
let curr = acf[lag];
let next = if lag < upper { acf[lag + 1] } else { 0.0 };
if curr >= prev && curr >= next && curr > best_val {
best_val = curr;
best_lag = lag;
}
lag += 1;
}
if best_lag == 0 {
return None;
}
let confidence = best_val.clamp(0.0, 1.0);
Some((best_lag, confidence))
}
fn detect_amdf(
samples: &[f32],
min_lag: usize,
max_lag: usize,
threshold: f32,
) -> Option<(usize, f32)> {
let n = samples.len();
let upper = max_lag.min(n - 1);
if min_lag > upper {
return None;
}
let mut amdf: Vec<f32> = vec![0.0; upper + 1];
for tau in min_lag..=upper {
let len = n - tau;
if len == 0 {
amdf[tau] = f32::MAX;
continue;
}
let sum: f32 = (0..len)
.map(|i| (samples[i] - samples[i + tau]).abs())
.sum();
amdf[tau] = sum / len as f32;
}
let mean_abs: f32 = samples.iter().map(|s| s.abs()).sum::<f32>() / n as f32;
if mean_abs < 1e-8 {
return None;
}
let norm_factor = mean_abs * 2.0;
let mut best_lag = 0usize;
let mut best_val = f32::MAX;
for tau in min_lag..=upper {
if amdf[tau] < best_val {
best_val = amdf[tau];
best_lag = tau;
}
}
if best_lag == 0 {
return None;
}
let rel = (best_val / norm_factor).clamp(0.0, 1.0);
let confidence = (1.0 - rel).clamp(0.0, 1.0);
if confidence < threshold {
return None;
}
Some((best_lag, confidence))
}
fn detect_cepstrum(
samples: &[f32],
min_lag: usize,
max_lag: usize,
threshold: f32,
) -> Option<(usize, f32)> {
let n = samples.len();
let fft_size = n.next_power_of_two();
let mut input: Vec<Complex<f32>> = samples.iter().map(|&s| Complex::new(s, 0.0)).collect();
input.resize(fft_size, Complex::zero());
let mut buffer = vec![Complex::zero(); fft_size];
if let Some(plan) = Plan::dft_1d(fft_size, Direction::Forward, Flags::ESTIMATE) {
plan.execute(&input, &mut buffer);
}
let log_mag_input: Vec<Complex<f32>> = buffer
.iter()
.map(|c| {
let mag = (c.norm() + 1e-10).ln();
Complex::new(mag, 0.0)
})
.collect();
let mut log_mag = vec![Complex::zero(); fft_size];
if let Some(plan) = Plan::dft_1d(fft_size, Direction::Backward, Flags::ESTIMATE) {
plan.execute(&log_mag_input, &mut log_mag);
}
let inv_n = 1.0 / fft_size as f32;
let upper = max_lag.min(fft_size / 2);
if min_lag > upper {
return None;
}
let mut best_quefrency = 0usize;
let mut best_val = f32::NEG_INFINITY;
for q in min_lag..=upper {
let val = (log_mag[q].re * inv_n).abs();
if val > best_val {
best_val = val;
best_quefrency = q;
}
}
if best_quefrency == 0 {
return None;
}
let mean_val: f32 = (min_lag..=upper)
.map(|q| (log_mag[q].re * inv_n).abs())
.sum::<f32>()
/ (upper - min_lag + 1) as f32;
let confidence = if mean_val > 1e-10 {
((best_val / mean_val - 1.0) / 10.0).clamp(0.0, 1.0)
} else {
0.0
};
if confidence < threshold {
return None;
}
Some((best_quefrency, confidence))
}
#[cfg(test)]
mod tests {
use super::*;
use std::f32::consts::PI;
fn sine_wave(freq_hz: f32, sample_rate: u32, num_samples: usize) -> Vec<f32> {
let sr = sample_rate as f32;
(0..num_samples)
.map(|i| (2.0 * PI * freq_hz * i as f32 / sr).sin())
.collect()
}
#[test]
fn test_frequency_to_midi_a4() {
let (note, cents) = frequency_to_midi(440.0);
assert_eq!(note, 69, "A4 should be MIDI 69");
assert!(cents.abs() < 0.01, "cents deviation should be near 0");
}
#[test]
fn test_frequency_to_midi_c4() {
let (note, _) = frequency_to_midi(261.626);
assert_eq!(note, 60, "C4 should be MIDI 60");
}
#[test]
fn test_frequency_to_midi_negative() {
let (note, cents) = frequency_to_midi(-1.0);
assert_eq!(note, 0);
assert_eq!(cents, 0.0);
}
#[test]
fn test_frequency_to_midi_cents_sharp() {
let (note, cents) = frequency_to_midi(442.0);
assert_eq!(note, 69, "still nearest A4");
assert!(cents > 0.0, "should be positive (sharp)");
assert!(cents < 10.0, "should be a small deviation");
}
#[test]
fn test_midi_to_note_name_a4() {
assert_eq!(midi_to_note_name(69), "A4");
}
#[test]
fn test_midi_to_note_name_c4() {
assert_eq!(midi_to_note_name(60), "C4");
}
#[test]
fn test_midi_to_note_name_c_sharp_4() {
assert_eq!(midi_to_note_name(61), "C#4");
}
#[test]
fn test_frequency_to_note_name_a4() {
let (name, cents) = frequency_to_note_name(440.0);
assert_eq!(name, "A4");
assert!(cents.abs() < 0.01);
}
#[test]
fn test_yin_detects_440hz() {
let samples = sine_wave(440.0, 44100, 4096);
let detector = PitchDetector::new(PitchAlgorithm::Yin, 44100);
let result = detector.detect(&samples);
assert!(result.is_some(), "YIN should detect 440 Hz sine wave");
let det = result.expect("detection present");
assert!(
(det.frequency_hz - 440.0).abs() < 10.0,
"frequency should be close to 440 Hz, got {}",
det.frequency_hz
);
assert!(det.confidence > 0.5, "confidence should be high");
}
#[test]
fn test_yin_silence_returns_none() {
let samples = vec![0.0f32; 2048];
let detector = PitchDetector::new(PitchAlgorithm::Yin, 44100);
let result = detector.detect(&samples);
assert!(result.is_none(), "YIN should return None for silence");
}
#[test]
fn test_autocorrelation_detects_220hz() {
let samples = sine_wave(220.0, 44100, 8192);
let mut detector = PitchDetector::new(PitchAlgorithm::Autocorrelation, 44100);
detector.min_frequency = 100.0;
detector.max_frequency = 500.0;
let result = detector.detect(&samples);
assert!(result.is_some(), "ACF should detect 220 Hz sine wave");
let det = result.expect("detection present");
assert!(
(det.frequency_hz - 220.0).abs() < 15.0,
"frequency should be near 220 Hz, got {}",
det.frequency_hz
);
}
#[test]
fn test_autocorrelation_silence_returns_none() {
let samples = vec![0.0f32; 4096];
let detector = PitchDetector::new(PitchAlgorithm::Autocorrelation, 44100);
let result = detector.detect(&samples);
assert!(result.is_none(), "ACF should return None for silence");
}
#[test]
fn test_amdf_detects_440hz() {
let samples = sine_wave(440.0, 44100, 4096);
let mut detector = PitchDetector::new(PitchAlgorithm::Amdf, 44100);
detector.threshold = 0.05;
let result = detector.detect(&samples);
assert!(
result.is_some(),
"AMDF should detect a pitch in 440 Hz sine wave"
);
let det = result.expect("detection present");
let ratios = [1.0, 2.0, 3.0, 4.0, 0.5, 0.25, 1.0 / 3.0, 1.0 / 4.0];
let harmonic_match = ratios
.iter()
.any(|&r| (det.frequency_hz - 440.0 * r).abs() < 20.0);
assert!(
harmonic_match,
"AMDF frequency {} should be a harmonic/sub-harmonic of 440 Hz",
det.frequency_hz
);
}
#[test]
fn test_amdf_silence_returns_none() {
let samples = vec![0.0f32; 2048];
let detector = PitchDetector::new(PitchAlgorithm::Amdf, 44100);
let result = detector.detect(&samples);
assert!(result.is_none(), "AMDF should return None for silence");
}
#[test]
fn test_cepstrum_detects_pitch() {
let sr = 44100u32;
let fundamental = 200.0f32;
let num_samples = 4096;
let samples: Vec<f32> = (0..num_samples)
.map(|i| {
let t = i as f32 / sr as f32;
(1..=7)
.step_by(2)
.map(|h| {
let h = h as f32;
(2.0 * PI * fundamental * h * t).sin() / h
})
.sum::<f32>()
})
.collect();
let mut detector = PitchDetector::new(PitchAlgorithm::Cepstrum, sr);
detector.min_frequency = 100.0;
detector.max_frequency = 800.0;
detector.threshold = 0.02;
let result = detector.detect(&samples);
if let Some(det) = result {
assert!(
det.frequency_hz >= detector.min_frequency
&& det.frequency_hz <= detector.max_frequency,
"cepstrum result {} outside [{}, {}]",
det.frequency_hz,
detector.min_frequency,
detector.max_frequency
);
}
}
#[test]
fn test_detection_has_valid_midi_note() {
let samples = sine_wave(440.0, 44100, 4096);
let detector = PitchDetector::new(PitchAlgorithm::Yin, 44100);
if let Some(det) = detector.detect(&samples) {
assert!(det.midi_note <= 127, "MIDI note must be in 0..=127");
}
}
#[test]
fn test_detection_has_valid_cents_deviation() {
let samples = sine_wave(440.0, 44100, 4096);
let detector = PitchDetector::new(PitchAlgorithm::Yin, 44100);
if let Some(det) = detector.detect(&samples) {
assert!(
det.cents_deviation.abs() <= 50.0,
"cents deviation must be in -50..50"
);
}
}
#[test]
fn test_detection_confidence_is_normalized() {
let samples = sine_wave(440.0, 44100, 4096);
let detector = PitchDetector::new(PitchAlgorithm::Yin, 44100);
if let Some(det) = detector.detect(&samples) {
assert!(
(0.0..=1.0).contains(&det.confidence),
"confidence must be 0..1, got {}",
det.confidence
);
}
}
}