use std::num::{NonZeroU32, NonZeroUsize};
use crate::operations::traits::AudioTransforms;
use crate::{
AudioSampleError, AudioSampleResult, AudioSamples, AudioTypeConversion, ChannelRequirement,
LayoutError, ParameterError, StandardSample,
};
use ndarray::{Array1, Array2, Axis};
use non_empty_slice::NonEmptySlice;
use num_complex::Complex;
use spectrograms::{
AmpScaleSpec, ChromaParams, Chromagram, CqtParams, CqtResult, CqtSpectrogram, FftPlanner,
Gammatone, GammatoneParams, GammatoneSpectrogram, LinearHz, LinearSpectrogram, LogHz,
LogHzParams, LogHzSpectrogram, LogParams, MelParams, MelSpectrogram, Mfcc, MfccParams,
Spectrogram, SpectrogramParams, StftParams, StftPlan, StftResult, WindowType, fft_convolve,
fft_deconvolve,
};
use std::cell::RefCell;
thread_local! {
static FFT_PLANNER: RefCell<FftPlanner> = RefCell::new(FftPlanner::new());
}
#[inline]
fn with_fft_planner<R>(f: impl FnOnce(&mut FftPlanner) -> R) -> R {
FFT_PLANNER.with(|p| f(&mut p.borrow_mut()))
}
#[derive(Clone, Debug, PartialEq)]
pub struct Psd {
frequencies: Vec<f64>,
density: Vec<f64>,
}
impl Psd {
#[inline]
pub(crate) fn new(frequencies: Vec<f64>, density: Vec<f64>) -> Self {
Self {
frequencies,
density,
}
}
#[inline]
#[must_use]
pub fn frequencies(&self) -> &[f64] {
&self.frequencies
}
#[inline]
#[must_use]
pub fn density(&self) -> &[f64] {
&self.density
}
#[inline]
#[must_use]
pub fn into_parts(self) -> (Vec<f64>, Vec<f64>) {
(self.frequencies, self.density)
}
}
impl<T> AudioTransforms for AudioSamples<'_, T>
where
T: StandardSample,
{
#[inline]
fn fft(&self, n_fft: NonZeroUsize) -> AudioSampleResult<Array2<Complex<f64>>> {
let working_samples = self.to_format::<f64>();
let mut channel_ffts: Vec<Array1<Complex<f64>>> =
Vec::with_capacity(self.num_channels().get() as usize);
with_fft_planner(|fft_planner| -> AudioSampleResult<()> {
for ch in working_samples.channels() {
let ch_slice = ch
.as_slice()
.expect("Can always get a mono channel as a slice");
let working_ch_slice = unsafe { NonEmptySlice::new_unchecked(ch_slice) };
let channel_fft: Array1<Complex<f64>> = fft_planner.fft(working_ch_slice, n_fft)?;
channel_ffts.push(channel_fft);
}
Ok(())
})?;
ndarray::stack(
Axis(0),
&channel_ffts.iter().map(|a| a.view()).collect::<Vec<_>>(),
)
.map_err(Into::into)
}
#[inline]
fn stft(&self, params: &StftParams) -> AudioSampleResult<StftResult> {
if self.is_multi_channel() {
return Err(AudioSampleError::Layout(
LayoutError::channel_count_unsupported(
"stft",
ChannelRequirement::Mono,
self.num_channels().get(),
),
));
}
let samples = self.to_format::<f64>();
let samples_slice = samples
.as_slice()
.expect("Safe since we have ensured mono audio samples");
let samples_slice = unsafe { NonEmptySlice::new_unchecked(samples_slice) };
let spect_params = SpectrogramParams::new(params.to_owned(), self.sample_rate_hz())?;
let mut stft_plan = StftPlan::new(&spect_params)?;
stft_plan
.compute(samples_slice, &spect_params)
.map_err(Into::into)
}
#[inline]
fn istft(stft: StftResult) -> AudioSampleResult<AudioSamples<'static, T>> {
let sample_rate = stft.sample_rate;
let sample_rate = unsafe { NonZeroU32::new_unchecked(sample_rate as u32) };
let data = &stft.data;
let n_fft = stft.params.n_fft();
let hop_size = stft.params.hop_size();
let window = stft.params.window();
let centre = stft.params.centre();
let data = spectrograms::istft(data, n_fft, hop_size, window, centre)?;
let audio = AudioSamples::from_mono_vec::<f64>(data, sample_rate);
Ok(audio)
}
#[inline]
fn linear_spectrogram<AmpScale>(
&self,
params: &SpectrogramParams,
db: Option<&LogParams>,
) -> AudioSampleResult<Spectrogram<LinearHz, AmpScale>>
where
AmpScale: AmpScaleSpec + 'static,
{
if self.is_multi_channel() {
return Err(AudioSampleError::Layout(
LayoutError::channel_count_unsupported(
"linear_frequency_spectrogram",
ChannelRequirement::Mono,
self.num_channels().get(),
),
));
}
let samples = self.cast_as_f64();
let samples_slice = samples.as_slice().expect("Safe since we have ensured mono");
let samples_slice = unsafe { NonEmptySlice::new_unchecked(samples_slice) };
LinearSpectrogram::<AmpScale>::compute(samples_slice, params, db).map_err(Into::into)
}
#[inline]
fn log_frequency_spectrogram<AmpScale>(
&self,
params: &SpectrogramParams,
loghz: &LogHzParams,
db: Option<&LogParams>,
) -> AudioSampleResult<Spectrogram<LogHz, AmpScale>>
where
AmpScale: AmpScaleSpec,
{
if self.is_multi_channel() {
return Err(AudioSampleError::Layout(
LayoutError::channel_count_unsupported(
"log_frequency_spectrogram",
ChannelRequirement::Mono,
self.num_channels().get(),
),
));
}
let samples: AudioSamples<'static, f64> = self.cast_as_f64();
let samples_slice = samples.as_slice().expect("Safe since we have ensured mono");
let samples_slice = unsafe { NonEmptySlice::new_unchecked(samples_slice) };
LogHzSpectrogram::<AmpScale>::compute(samples_slice, params, loghz, db).map_err(Into::into)
}
#[inline]
fn mel_spectrogram<AmpScale>(
&self,
params: &SpectrogramParams,
mel: &MelParams,
db: Option<&LogParams>, ) -> AudioSampleResult<MelSpectrogram<AmpScale>>
where
AmpScale: AmpScaleSpec,
{
if self.is_multi_channel() {
return Err(AudioSampleError::Layout(
LayoutError::channel_count_unsupported(
"mel_spectrogram",
ChannelRequirement::Mono,
self.num_channels().get(),
),
));
}
let samples = self.cast_as_f64();
let samples_slice = samples.as_slice().expect("Safe since we have ensured mono");
let samples_slice = unsafe { NonEmptySlice::new_unchecked(samples_slice) };
MelSpectrogram::<AmpScale>::compute(samples_slice, params, mel, db).map_err(Into::into)
}
#[inline]
fn mfcc(
&self,
stft_params: &StftParams,
n_mels: NonZeroUsize,
mfcc_params: &MfccParams,
) -> AudioSampleResult<Mfcc> {
if self.is_multi_channel() {
return Err(AudioSampleError::Layout(
LayoutError::channel_count_unsupported(
"mfcc",
ChannelRequirement::Mono,
self.num_channels().get(),
),
));
}
let samples = self.cast_as_f64();
let samples_slice = samples.as_slice().expect("Safe since we have ensured mono");
let samples_slice = unsafe { NonEmptySlice::new_unchecked(samples_slice) };
let sample_rate_f = self.sample_rate_hz();
spectrograms::mfcc(
samples_slice,
stft_params,
sample_rate_f,
n_mels,
mfcc_params,
)
.map_err(Into::into)
}
#[inline]
fn chromagram(
&self,
stft_params: &StftParams,
cfg: &ChromaParams,
) -> AudioSampleResult<Chromagram> {
if self.is_multi_channel() {
return Err(AudioSampleError::Layout(
LayoutError::channel_count_unsupported(
"chromagram",
ChannelRequirement::Mono,
self.num_channels().get(),
),
));
}
let samples = self.cast_as_f64();
let samples_slice = samples.as_slice().expect("Safe since we have ensured mono");
let samples_slice = unsafe { NonEmptySlice::new_unchecked(samples_slice) };
let sample_rate_f = self.sample_rate_hz();
spectrograms::chromagram(samples_slice, stft_params, sample_rate_f, cfg).map_err(Into::into)
}
#[inline]
fn gammatone_spectrogram<AmpScale>(
&self,
params: &SpectrogramParams,
gammatone_params: &GammatoneParams,
db: Option<&LogParams>,
) -> AudioSampleResult<Spectrogram<Gammatone, AmpScale>>
where
AmpScale: AmpScaleSpec,
{
if self.is_multi_channel() {
return Err(AudioSampleError::Layout(
LayoutError::channel_count_unsupported(
"gammatone_spectrogram",
ChannelRequirement::Mono,
self.num_channels().get(),
),
));
}
let samples = self.cast_as_f64();
let samples_slice = samples.as_slice().expect("Safe since we have ensured mono");
let samples_slice = unsafe { NonEmptySlice::new_unchecked(samples_slice) };
GammatoneSpectrogram::<AmpScale>::compute(samples_slice, params, gammatone_params, db)
.map_err(Into::into)
}
#[inline]
fn constant_q_transform(
&self,
params: &CqtParams,
hop_size: NonZeroUsize,
) -> AudioSampleResult<CqtResult> {
if self.is_multi_channel() {
return Err(AudioSampleError::Layout(
LayoutError::channel_count_unsupported(
"constant_q_transform",
ChannelRequirement::Mono,
self.num_channels().get(),
),
));
}
let samples = self.cast_as_f64();
let samples_slice = samples.as_slice().expect("Safe since we have ensured mono");
let samples_slice = unsafe { NonEmptySlice::new_unchecked(samples_slice) };
spectrograms::cqt(
samples_slice,
f64::from(self.sample_rate().get()),
params,
hop_size,
)
.map_err(Into::into)
}
#[inline]
fn cqt_spectrogram<AmpScale>(
&self,
params: &SpectrogramParams,
cqt: &CqtParams,
db: Option<&LogParams>,
) -> AudioSampleResult<CqtSpectrogram<AmpScale>>
where
AmpScale: AmpScaleSpec,
{
if self.is_multi_channel() {
return Err(AudioSampleError::Layout(
LayoutError::channel_count_unsupported(
"cqt",
ChannelRequirement::Mono,
self.num_channels().get(),
),
));
}
let working_samples = self.as_f64();
let working_samples_slice = working_samples
.as_slice()
.expect("Safe since we have ensured mono");
let working_samples_slice = unsafe { NonEmptySlice::new_unchecked(working_samples_slice) };
CqtSpectrogram::<AmpScale>::compute(working_samples_slice, params, cqt, db)
.map_err(Into::into)
}
#[inline]
fn power_spectral_density(
&self,
window_size: NonZeroUsize,
overlap: f64,
) -> AudioSampleResult<Psd> {
if self.is_multi_channel() {
return Err(AudioSampleError::Layout(
LayoutError::channel_count_unsupported(
"power_spectral_density",
ChannelRequirement::Mono,
self.num_channels().get(),
),
));
}
if !(0.0..1.0).contains(&overlap) {
return Err(AudioSampleError::Parameter(ParameterError::out_of_range(
"overlap",
overlap.to_string(),
"0.0",
"1.0",
"overlap must be in [0, 1)",
)));
}
if self.samples_per_channel() < window_size {
return Err(AudioSampleError::Parameter(ParameterError::invalid_value(
"window_size",
"window_size must not exceed the signal length",
)));
}
let samples = self.to_format::<f64>();
let signal = samples
.as_slice()
.expect("Safe since we have ensured mono audio samples");
let win = window_size.get();
let hop = ((1.0 - overlap) * win as f64).floor().max(1.0) as usize;
let sample_rate = self.sample_rate_hz();
let mut sum: Vec<f64> = Vec::new();
let mut segment_count = 0u64;
with_fft_planner(|planner| -> AudioSampleResult<()> {
let mut start = 0;
while start + win <= signal.len() {
let segment = &signal[start..start + win];
let segment_slice = unsafe { NonEmptySlice::new_unchecked(segment) };
let power = planner.power_spectrum(
segment_slice,
window_size,
Some(WindowType::Hanning),
)?;
if sum.is_empty() {
sum = power.to_vec();
} else {
for (acc, &val) in sum.iter_mut().zip(power.iter()) {
*acc += val;
}
}
segment_count += 1;
start += hop;
}
Ok(())
})?;
let count = segment_count as f64;
for val in &mut sum {
*val /= count;
}
let freq_resolution = sample_rate / win as f64;
for val in &mut sum {
*val /= freq_resolution;
}
let frequencies: Vec<f64> = (0..sum.len())
.map(|i| i as f64 * sample_rate / win as f64)
.collect();
Ok(Psd::new(frequencies, sum))
}
#[inline]
fn linear_magnitude_spectrogram(
&self,
params: &SpectrogramParams,
) -> AudioSampleResult<spectrograms::LinearMagnitudeSpectrogram> {
self.linear_spectrogram::<spectrograms::Magnitude>(params, None)
}
#[inline]
fn linear_power_spectrogram(
&self,
params: &SpectrogramParams,
) -> AudioSampleResult<spectrograms::LinearPowerSpectrogram> {
self.linear_spectrogram::<spectrograms::Power>(params, None)
}
#[inline]
fn linear_db_spectrogram(
&self,
params: &SpectrogramParams,
db: &LogParams,
) -> AudioSampleResult<spectrograms::LinearDbSpectrogram> {
self.linear_spectrogram::<spectrograms::Decibels>(params, Some(db))
}
#[inline]
fn loghz_power_spectrogram(
&self,
params: &SpectrogramParams,
loghz: &LogHzParams,
) -> AudioSampleResult<spectrograms::LogHzPowerSpectrogram> {
self.log_frequency_spectrogram::<spectrograms::Power>(params, loghz, None)
}
#[inline]
fn loghz_magnitude_spectrogram(
&self,
params: &SpectrogramParams,
loghz: &LogHzParams,
) -> AudioSampleResult<spectrograms::LogHzMagnitudeSpectrogram> {
self.log_frequency_spectrogram::<spectrograms::Magnitude>(params, loghz, None)
}
#[inline]
fn loghz_db_spectrogram(
&self,
params: &SpectrogramParams,
loghz: &LogHzParams,
db: &LogParams,
) -> AudioSampleResult<spectrograms::LogHzDbSpectrogram> {
self.log_frequency_spectrogram::<spectrograms::Decibels>(params, loghz, Some(db))
}
#[inline]
fn mel_mag_spectrogram(
&self,
params: &SpectrogramParams,
mel: &MelParams,
) -> AudioSampleResult<spectrograms::MelMagnitudeSpectrogram> {
self.mel_spectrogram(params, mel, None)
}
#[inline]
fn mel_db_spectrogram(
&self,
params: &SpectrogramParams,
mel: &MelParams,
db: &LogParams,
) -> AudioSampleResult<spectrograms::LogMelSpectrogram> {
self.mel_spectrogram(params, mel, Some(db))
}
#[inline]
fn mel_power_spectrogram(
&self,
params: &SpectrogramParams,
mel: &MelParams,
) -> AudioSampleResult<spectrograms::MelPowerSpectrogram> {
self.mel_spectrogram(params, mel, None)
}
#[inline]
fn gammatone_magnitude_spectrogram(
&self,
params: &SpectrogramParams,
gammatone_params: &GammatoneParams,
) -> AudioSampleResult<spectrograms::GammatoneMagnitudeSpectrogram> {
self.gammatone_spectrogram::<spectrograms::Magnitude>(params, gammatone_params, None)
}
#[inline]
fn gammatone_power_spectrogram(
&self,
params: &SpectrogramParams,
gammatone_params: &GammatoneParams,
) -> AudioSampleResult<spectrograms::GammatonePowerSpectrogram> {
self.gammatone_spectrogram::<spectrograms::Power>(params, gammatone_params, None)
}
#[inline]
fn gammatone_db_spectrogram(
&self,
params: &SpectrogramParams,
gammatone_params: &GammatoneParams,
db: &LogParams,
) -> AudioSampleResult<spectrograms::GammatoneDbSpectrogram> {
self.gammatone_spectrogram::<spectrograms::Decibels>(params, gammatone_params, Some(db))
}
#[inline]
fn cqt_magnitude_spectrogram(
&self,
params: &SpectrogramParams,
cqt: &CqtParams,
) -> AudioSampleResult<spectrograms::CqtMagnitudeSpectrogram> {
self.cqt_spectrogram::<spectrograms::Magnitude>(params, cqt, None)
}
#[inline]
fn cqt_power_spectrogram(
&self,
params: &SpectrogramParams,
cqt: &CqtParams,
) -> AudioSampleResult<spectrograms::CqtPowerSpectrogram> {
self.cqt_spectrogram::<spectrograms::Power>(params, cqt, None)
}
#[inline]
fn cqt_db_spectrogram(
&self,
params: &SpectrogramParams,
cqt: &CqtParams,
db: &LogParams,
) -> AudioSampleResult<spectrograms::CqtDbSpectrogram> {
self.cqt_spectrogram::<spectrograms::Decibels>(params, cqt, Some(db))
}
#[inline]
fn convolve(&self, other: &Self) -> AudioSampleResult<AudioSamples<'static, Self::Sample>> {
let result_f64 = convolve_or_deconvolve_mono(self, other, |a, b| Ok(fft_convolve(a, b)?))?;
from_f64_mono::<T>(result_f64, self.sample_rate())
}
#[inline]
fn deconvolve(
&self,
denominator: &Self,
regularization: f64,
) -> AudioSampleResult<AudioSamples<'static, Self::Sample>> {
let result_f64 = convolve_or_deconvolve_mono(self, denominator, |a, b| {
Ok(fft_deconvolve(a, b, regularization)?)
})?;
from_f64_mono::<T>(result_f64, self.sample_rate())
}
#[inline]
fn magphase(
complex_spect: &Array2<Complex<f64>>,
power: Option<NonZeroUsize>,
) -> (Array2<f64>, Array2<Complex<f64>>) {
let mut mag = complex_spect.mapv(num_complex::Complex::norm);
let zeros_to_ones = mag.mapv(|x| if x == 0.0 { 1.0 } else { 0.0 });
let mag_nonzero = &mag + &zeros_to_ones;
let mut phase = complex_spect.clone();
let power = power.map_or(1.0, |p| p.get() as f64);
ndarray::Zip::from(&mut phase)
.and(&mag_nonzero)
.and(&zeros_to_ones)
.for_each(|p, &m_nz, &z| {
let div = Complex {
re: p.re / m_nz + z, im: p.im / m_nz,
};
*p = div;
});
mag.mapv_inplace(|x| x.powf(power));
(mag, phase)
}
}
fn mono_f64<T: StandardSample>(audio: &AudioSamples<'_, T>) -> AudioSampleResult<Vec<f64>> {
if audio.num_channels().get() != 1 {
return Err(AudioSampleError::Layout(
LayoutError::channel_count_unsupported(
"convolve/deconvolve (mono only)",
ChannelRequirement::Mono,
audio.num_channels().get(),
),
));
}
let working = audio.to_format::<f64>();
let ch = working.as_slice().expect("validated single channel");
Ok(ch.to_vec())
}
fn convolve_or_deconvolve_mono<T, F>(
a: &AudioSamples<'_, T>,
b: &AudioSamples<'_, T>,
op: F,
) -> AudioSampleResult<Vec<f64>>
where
T: StandardSample,
F: FnOnce(
&non_empty_slice::NonEmptySlice<f64>,
&non_empty_slice::NonEmptySlice<f64>,
) -> AudioSampleResult<non_empty_slice::NonEmptyVec<f64>>,
{
let a_vec = mono_f64(a)?;
let b_vec = mono_f64(b)?;
let a_ne = non_empty_slice::NonEmptySlice::new(a_vec.as_slice()).ok_or_else(|| {
AudioSampleError::Parameter(ParameterError::invalid_value("self", "empty signal"))
})?;
let b_ne = non_empty_slice::NonEmptySlice::new(b_vec.as_slice()).ok_or_else(|| {
AudioSampleError::Parameter(ParameterError::invalid_value("other", "empty signal"))
})?;
Ok(op(a_ne, b_ne)?.into_vec())
}
fn from_f64_mono<T: StandardSample>(
data: Vec<f64>,
sample_rate: std::num::NonZeroU32,
) -> AudioSampleResult<AudioSamples<'static, T>> {
use non_empty_slice::NonEmptyVec;
let ne_vec = NonEmptyVec::new(data).map_err(|_| {
AudioSampleError::Parameter(ParameterError::invalid_value("data", "empty result"))
})?;
Ok(AudioSamples::<T>::from_mono_vec::<f64>(ne_vec, sample_rate))
}
#[cfg(test)]
mod convolution_tests {
use crate::{AudioSampleError, AudioSamples, AudioTransforms, LayoutError, sample_rate};
use ndarray::array;
#[test]
fn convolve_with_unit_impulse_is_identity() {
let signal =
AudioSamples::new_mono(array![1.0_f64, 2.0, 3.0], sample_rate!(48000)).unwrap();
let impulse = AudioSamples::new_mono(array![1.0_f64], sample_rate!(48000)).unwrap();
let out = signal.convolve(&impulse).unwrap();
assert_eq!(out.sample_rate().get(), 48000);
let ch = out.as_slice().unwrap();
assert!((ch[0] - 1.0).abs() < 1e-9);
assert!((ch[1] - 2.0).abs() < 1e-9);
assert!((ch[2] - 3.0).abs() < 1e-9);
}
#[test]
fn deconvolve_inverts_convolve() {
let excitation = AudioSamples::new_mono(
array![1.0_f64, 0.7, -0.3, 0.2, 0.9, -0.5, 0.1, 0.4],
sample_rate!(48000),
)
.unwrap();
let system =
AudioSamples::new_mono(array![0.0_f64, 0.0, 1.0, 0.5], sample_rate!(48000)).unwrap();
let recorded = excitation.convolve(&system).unwrap();
let recovered = recorded.deconvolve(&excitation, 0.0).unwrap();
let r = recovered.as_slice().unwrap();
for (i, want) in [0.0, 0.0, 1.0, 0.5].iter().enumerate() {
assert!(
(r[i] - want).abs() < 1e-6,
"tap {i}: got {}, want {want}",
r[i]
);
}
}
#[test]
fn multichannel_input_is_rejected() {
let stereo = AudioSamples::new_multi_channel(
array![[1.0_f64, 2.0, 3.0, 4.0], [5.0, 6.0, 7.0, 8.0]],
sample_rate!(48000),
)
.unwrap();
let mono = AudioSamples::new_mono(array![1.0_f64], sample_rate!(48000)).unwrap();
let err = stereo.convolve(&mono).unwrap_err();
assert!(
matches!(
err,
AudioSampleError::Layout(LayoutError::ChannelCountUnsupported { .. })
),
"expected Layout(ChannelCountUnsupported), got: {err:?}"
);
}
}
#[cfg(test)]
mod psd_tests {
use super::Psd;
#[test]
fn psd_accessors_and_into_parts_round_trip() {
let psd = Psd::new(vec![0.0, 10.0, 20.0], vec![1.0, 0.5, 0.25]);
assert_eq!(psd.frequencies(), &[0.0, 10.0, 20.0]);
assert_eq!(psd.density(), &[1.0, 0.5, 0.25]);
let (freqs, density) = psd.into_parts();
assert_eq!(freqs, vec![0.0, 10.0, 20.0]);
assert_eq!(density, vec![1.0, 0.5, 0.25]);
}
}