use std::num::NonZeroUsize;
use non_empty_slice::NonEmptySlice;
use pyo3::prelude::*;
use super::dtype::{Dtype, array1_to_py, array2_to_py, parse_dtype, real_1d_vec};
use super::spectrogram::{PyArrayData, PyScalar, real_dlpack};
use crate::binaural::{
ILDSpectrogramParams, ILRSpectrogramParams, IPDSpectrogramParams, ITDSpectrogramParams,
IldSpectrogram, IlrSpectrogram, IpdSpectrogram, ItdSpectrogram, compute_ild_spectrogram,
compute_ilr_spectrogram, compute_ilr_spectrogram_diff, compute_ipd_spectrogram,
compute_itd_spectrogram, compute_itd_spectrogram_diff,
};
use crate::{StftPlan, python::PySpectrogramParams};
fn read_stereo<T: PyScalar>(
py: Python<'_>,
audio: &[Bound<'_, PyAny>; 2],
) -> PyResult<(Vec<T>, Vec<T>)> {
let left = real_1d_vec::<T>(py, &audio[0])?;
let right = real_1d_vec::<T>(py, &audio[1])?;
Ok((left, right))
}
fn as_stereo_slices<'a, T>(
left: &'a [T],
right: &'a [T],
) -> PyResult<(&'a NonEmptySlice<T>, &'a NonEmptySlice<T>)> {
let left = NonEmptySlice::new(left).ok_or_else(|| {
pyo3::exceptions::PyValueError::new_err("left audio array must not be empty")
})?;
let right = NonEmptySlice::new(right).ok_or_else(|| {
pyo3::exceptions::PyValueError::new_err("right audio array must not be empty")
})?;
Ok((left, right))
}
macro_rules! binaural_result_class {
(
$enum:ident,
$py_struct:ident,
$py_name:literal,
$rust:ident,
$py_params:ident,
$hist:item
) => {
pub(crate) enum $enum {
F32($rust<f32>),
F64($rust<f64>),
}
impl From<$rust<f32>> for $enum {
#[inline]
fn from(s: $rust<f32>) -> Self {
Self::F32(s)
}
}
impl From<$rust<f64>> for $enum {
#[inline]
fn from(s: $rust<f64>) -> Self {
Self::F64(s)
}
}
#[doc = concat!($py_name, " computation result.")]
#[pyclass(name = $py_name, skip_from_py_object)]
pub struct $py_struct {
inner: $enum,
}
impl $py_struct {
pub(crate) fn from_result<T>(_py: Python<'_>, result: $rust<T>) -> Self
where
$rust<T>: Into<$enum>,
{
Self {
inner: result.into(),
}
}
fn data_any<'py>(&self, py: Python<'py>) -> Bound<'py, PyAny> {
match &self.inner {
$enum::F32(s) => {
numpy::PyArray2::from_owned_array(py, s.data.clone()).into_any()
}
$enum::F64(s) => {
numpy::PyArray2::from_owned_array(py, s.data.clone()).into_any()
}
}
}
fn array_data(&self, py: Python<'_>) -> PyArrayData {
match &self.inner {
$enum::F32(s) => <f32 as PyScalar>::into_array_data(py, s.data.clone()),
$enum::F64(s) => <f64 as PyScalar>::into_array_data(py, s.data.clone()),
}
}
}
#[pymethods]
impl $py_struct {
#[getter]
fn data<'py>(&self, py: Python<'py>) -> Bound<'py, PyAny> {
self.data_any(py)
}
#[getter]
fn dtype(&self) -> &'static str {
match &self.inner {
$enum::F32(_) => "float32",
$enum::F64(_) => "float64",
}
}
#[getter]
fn n_bins(&self) -> usize {
match &self.inner {
$enum::F32(s) => s.n_bins().get(),
$enum::F64(s) => s.n_bins().get(),
}
}
#[getter]
fn n_frames(&self) -> usize {
match &self.inner {
$enum::F32(s) => s.n_frames().get(),
$enum::F64(s) => s.n_frames().get(),
}
}
#[getter]
fn shape(&self) -> (usize, usize) {
(self.n_bins(), self.n_frames())
}
#[getter]
fn frequencies(&self) -> Vec<f64> {
match &self.inner {
$enum::F32(s) => s.frequencies().to_vec(),
$enum::F64(s) => s.frequencies().to_vec(),
}
}
#[getter]
fn times(&self) -> Vec<f64> {
match &self.inner {
$enum::F32(s) => s.times().to_vec(),
$enum::F64(s) => s.times().to_vec(),
}
}
fn frequency_range(&self) -> (f64, f64) {
match &self.inner {
$enum::F32(s) => s.frequency_range(),
$enum::F64(s) => s.frequency_range(),
}
}
fn duration(&self) -> f64 {
match &self.inner {
$enum::F32(s) => s.duration(),
$enum::F64(s) => s.duration(),
}
}
#[getter]
fn params(&self) -> $py_params {
match &self.inner {
$enum::F32(s) => $py_params {
inner: s.params().clone(),
},
$enum::F64(s) => $py_params {
inner: s.params().clone(),
},
}
}
#[pyo3(signature = (dtype=None))]
fn __array__<'py>(
&self,
py: Python<'py>,
dtype: Option<&Bound<'py, PyAny>>,
) -> PyResult<Bound<'py, PyAny>> {
let arr = self.data_any(py);
if let Some(dt) = dtype {
arr.call_method1("astype", (dt,))
} else {
Ok(arr)
}
}
#[staticmethod]
const fn __dlpack_device__() -> (i32, i32) {
(1, 0) }
#[pyo3(signature = (*, stream=None, max_version=None, dl_device=None, copy=None))]
fn __dlpack__<'py>(
&self,
py: Python<'py>,
stream: Option<&Bound<'py, PyAny>>,
max_version: Option<(u32, u32)>,
dl_device: Option<(i32, i32)>,
copy: Option<bool>,
) -> PyResult<Bound<'py, pyo3::types::PyCapsule>> {
let data = self.array_data(py);
real_dlpack(py, &data, stream, max_version, dl_device, copy)
}
fn __repr__(&self) -> String {
format!(
"{}(shape=({}, {}), dtype={})",
$py_name,
self.n_bins(),
self.n_frames(),
self.dtype(),
)
}
$hist
}
};
}
binaural_result_class!(
ItdInner,
PyItdSpectrogram,
"ItdSpectrogram",
ItdSpectrogram,
PyITDSpectrogramParams,
#[pyo3(signature = (num_bins=None, delay_range=None, energy_weighted=false, normalize=false))]
fn histogram<'py>(
&self,
py: Python<'py>,
num_bins: Option<usize>,
delay_range: Option<(f64, f64)>,
energy_weighted: bool,
normalize: bool,
) -> Bound<'py, PyAny> {
let num_bins = num_bins.and_then(NonZeroUsize::new);
let hist = match &self.inner {
ItdInner::F32(s) => s.histogram(num_bins, delay_range, energy_weighted, normalize),
ItdInner::F64(s) => s.histogram(num_bins, delay_range, energy_weighted, normalize),
};
array2_to_py(py, hist).into_bound(py)
}
);
binaural_result_class!(
IpdInner,
PyIpdSpectrogram,
"IpdSpectrogram",
IpdSpectrogram,
PyIPDSpectrogramParams,
#[pyo3(signature = (num_bins=None, phase_range=None, energy_weighted=false, normalize=false))]
fn histogram<'py>(
&self,
py: Python<'py>,
num_bins: Option<usize>,
phase_range: Option<(f64, f64)>,
energy_weighted: bool,
normalize: bool,
) -> Bound<'py, PyAny> {
let num_bins = num_bins.and_then(NonZeroUsize::new);
let hist = match &self.inner {
IpdInner::F32(s) => s.histogram(num_bins, phase_range, energy_weighted, normalize),
IpdInner::F64(s) => s.histogram(num_bins, phase_range, energy_weighted, normalize),
};
array2_to_py(py, hist).into_bound(py)
}
);
binaural_result_class!(
IldInner,
PyIldSpectrogram,
"IldSpectrogram",
IldSpectrogram,
PyILDSpectrogramParams,
#[pyo3(signature = (num_bins=None, db_range=None, exponent=None, energy_weighted=false, normalize=false))]
fn histogram<'py>(
&self,
py: Python<'py>,
num_bins: Option<usize>,
db_range: Option<(f64, f64)>,
exponent: Option<i32>,
energy_weighted: bool,
normalize: bool,
) -> Bound<'py, PyAny> {
let num_bins = num_bins.and_then(NonZeroUsize::new);
let hist = match &self.inner {
IldInner::F32(s) => {
s.histogram(num_bins, db_range, exponent, energy_weighted, normalize)
}
IldInner::F64(s) => {
s.histogram(num_bins, db_range, exponent, energy_weighted, normalize)
}
};
array2_to_py(py, hist).into_bound(py)
}
);
binaural_result_class!(
IlrInner,
PyIlrSpectrogram,
"IlrSpectrogram",
IlrSpectrogram,
PyILRSpectrogramParams,
#[pyo3(signature = (num_bins=None, ratio_range=None, exponent=None, energy_weighted=false, normalize=false))]
fn histogram<'py>(
&self,
py: Python<'py>,
num_bins: Option<usize>,
ratio_range: Option<(f64, f64)>,
exponent: Option<i32>,
energy_weighted: bool,
normalize: bool,
) -> Bound<'py, PyAny> {
let num_bins = num_bins.and_then(NonZeroUsize::new);
let hist = match &self.inner {
IlrInner::F32(s) => {
s.histogram(num_bins, ratio_range, exponent, energy_weighted, normalize)
}
IlrInner::F64(s) => {
s.histogram(num_bins, ratio_range, exponent, energy_weighted, normalize)
}
};
array2_to_py(py, hist).into_bound(py)
}
);
#[pyclass(name = "ITDSpectrogramParams", from_py_object)]
#[derive(Debug, Clone)]
pub struct PyITDSpectrogramParams {
pub(crate) inner: ITDSpectrogramParams,
}
#[pymethods]
impl PyITDSpectrogramParams {
#[new]
#[pyo3(signature = (spectrogram_params: "SpectrogramParams", start_freq: "float" = 50.0, end_freq: "float" = 620.0, magphase_power: "Optional[int]" = 1), text_signature = "(spectrogram_params: SpectrogramParams, start_freq: float = 50.0, end_freq: float = 620.0, magphase_power: Optional[int] = 1) -> ITDSpectrogramParams")]
fn new(
spectrogram_params: PySpectrogramParams,
start_freq: Option<f64>,
end_freq: Option<f64>,
magphase_power: Option<usize>,
) -> Self {
let inner = ITDSpectrogramParams {
spectrogram_params: spectrogram_params.into(),
start_freq: start_freq.unwrap_or(50.0),
end_freq: end_freq.unwrap_or(620.0),
magphase_power: magphase_power
.and_then(NonZeroUsize::new)
.unwrap_or_else(|| crate::nzu!(1)),
};
Self { inner }
}
#[getter]
fn spectrogram_params(&self) -> PySpectrogramParams {
PySpectrogramParams::from(self.inner.spectrogram_params.clone())
}
#[getter]
const fn start_freq(&self) -> f64 {
self.inner.start_freq
}
#[getter]
const fn end_freq(&self) -> f64 {
self.inner.end_freq
}
#[getter]
const fn magphase_power(&self) -> NonZeroUsize {
self.inner.magphase_power
}
}
impl From<ITDSpectrogramParams> for PyITDSpectrogramParams {
#[inline]
fn from(inner: ITDSpectrogramParams) -> Self {
Self { inner }
}
}
impl From<PyITDSpectrogramParams> for ITDSpectrogramParams {
#[inline]
fn from(val: PyITDSpectrogramParams) -> Self {
val.inner
}
}
#[pyfunction(name = "compute_itd_spectrogram")]
#[pyo3(signature = (audio: "list[numpy.typing.NDArray[numpy.float64]]", params: "ITDSpectrogramParams", dtype: "str" = None), text_signature = "(audio: list[numpy.typing.NDArray[numpy.float64]], params: ITDSpectrogramParams, dtype: str = \"float64\") -> ItdSpectrogram")]
fn py_compute_itd_spectrogram(
py: Python<'_>,
audio: [Bound<'_, PyAny>; 2],
params: &PyITDSpectrogramParams,
dtype: Option<&str>,
) -> PyResult<PyItdSpectrogram> {
fn run<T: PyScalar>(
py: Python<'_>,
audio: &[Bound<'_, PyAny>; 2],
params: &PyITDSpectrogramParams,
) -> PyResult<PyItdSpectrogram>
where
ItdSpectrogram<T>: Into<ItdInner>,
{
let mut plan = StftPlan::<T>::new(¶ms.inner.spectrogram_params).map_err(|e| {
PyErr::new::<pyo3::exceptions::PyRuntimeError, _>(format!(
"Failed to create STFT plan: {e}"
))
})?;
let (left, right) = read_stereo::<T>(py, audio)?;
let (left_s, right_s) = as_stereo_slices(&left, &right)?;
let itd =
compute_itd_spectrogram([left_s, right_s], ¶ms.inner, &mut plan).map_err(|e| {
PyErr::new::<pyo3::exceptions::PyRuntimeError, _>(format!(
"Failed to compute ITD spectrogram: {e}"
))
})?;
Ok(PyItdSpectrogram::from_result(py, itd))
}
match parse_dtype(dtype)? {
Dtype::F32 => run::<f32>(py, &audio, params),
Dtype::F64 => run::<f64>(py, &audio, params),
}
}
#[pyclass(name = "IPDSpectrogramParams", from_py_object)]
#[derive(Debug, Clone)]
pub struct PyIPDSpectrogramParams {
pub(crate) inner: IPDSpectrogramParams,
}
#[pymethods]
impl PyIPDSpectrogramParams {
#[new]
#[pyo3(signature = (spectrogram_params, start_freq = 50.0, end_freq = 620.0, wrapped = false), text_signature = "(spectrogram_params: SpectrogramParams, start_freq: float = 50.0, end_freq: float = 620.0, wrapped: bool = False) -> IPDSpectrogramParams")]
fn new(
spectrogram_params: PySpectrogramParams,
start_freq: Option<f64>,
end_freq: Option<f64>,
wrapped: Option<bool>,
) -> PyResult<Self> {
let inner = IPDSpectrogramParams::new(
spectrogram_params.into(),
start_freq.unwrap_or(50.0),
end_freq.unwrap_or(620.0),
wrapped.unwrap_or(false),
)
.map_err(|e| PyErr::new::<pyo3::exceptions::PyValueError, _>(format!("{e}")))?;
Ok(Self { inner })
}
#[getter]
fn spectrogram_params(&self) -> PySpectrogramParams {
PySpectrogramParams::from(self.inner.spectrogram_params.clone())
}
#[getter]
const fn start_freq(&self) -> f64 {
self.inner.start_freq
}
#[getter]
const fn end_freq(&self) -> f64 {
self.inner.end_freq
}
#[getter]
const fn wrapped(&self) -> bool {
self.inner.wrapped
}
}
#[pyfunction(name = "compute_ipd_spectrogram")]
#[pyo3(signature = (audio: "list[numpy.typing.NDArray[numpy.float64]]", params: "IPDSpectrogramParams", dtype: "str" = None), text_signature = "(audio: list[numpy.typing.NDArray[numpy.float64]], params: IPDSpectrogramParams, dtype: str = \"float64\") -> IpdSpectrogram")]
fn py_compute_ipd_spectrogram(
py: Python<'_>,
audio: [Bound<'_, PyAny>; 2],
params: &PyIPDSpectrogramParams,
dtype: Option<&str>,
) -> PyResult<PyIpdSpectrogram> {
fn run<T: PyScalar>(
py: Python<'_>,
audio: &[Bound<'_, PyAny>; 2],
params: &PyIPDSpectrogramParams,
) -> PyResult<PyIpdSpectrogram>
where
IpdSpectrogram<T>: Into<IpdInner>,
{
let mut plan = StftPlan::<T>::new(¶ms.inner.spectrogram_params).map_err(|e| {
PyErr::new::<pyo3::exceptions::PyRuntimeError, _>(format!(
"Failed to create STFT plan: {e}"
))
})?;
let (left, right) = read_stereo::<T>(py, audio)?;
let (left_s, right_s) = as_stereo_slices(&left, &right)?;
let ipd =
compute_ipd_spectrogram([left_s, right_s], ¶ms.inner, &mut plan).map_err(|e| {
PyErr::new::<pyo3::exceptions::PyRuntimeError, _>(format!(
"Failed to compute IPD spectrogram: {e}"
))
})?;
Ok(PyIpdSpectrogram::from_result(py, ipd))
}
match parse_dtype(dtype)? {
Dtype::F32 => run::<f32>(py, &audio, params),
Dtype::F64 => run::<f64>(py, &audio, params),
}
}
#[pyclass(name = "ILDSpectrogramParams", from_py_object)]
#[derive(Debug, Clone)]
pub struct PyILDSpectrogramParams {
pub(crate) inner: ILDSpectrogramParams,
}
#[pymethods]
impl PyILDSpectrogramParams {
#[new]
#[pyo3(signature = (spectrogram_params: "SpectrogramParams", start_freq: "float" = 1700.0, end_freq: "float" = 4600.0), text_signature = "(spectrogram_params: SpectrogramParams, start_freq: float = 1700.0, end_freq: float = 4600.0) -> ILDSpectrogramParams")]
fn new(
spectrogram_params: PySpectrogramParams,
start_freq: Option<f64>,
end_freq: Option<f64>,
) -> PyResult<Self> {
let inner = ILDSpectrogramParams::new(
spectrogram_params.into(),
start_freq.unwrap_or(1700.0),
end_freq.unwrap_or(4600.0),
)
.map_err(|e| PyErr::new::<pyo3::exceptions::PyValueError, _>(format!("{e}")))?;
Ok(Self { inner })
}
#[getter]
fn spectrogram_params(&self) -> PySpectrogramParams {
PySpectrogramParams::from(self.inner.spectrogram_params.clone())
}
#[getter]
const fn start_freq(&self) -> f64 {
self.inner.start_freq
}
#[getter]
const fn end_freq(&self) -> f64 {
self.inner.end_freq
}
}
#[pyfunction(name = "compute_ild_spectrogram")]
#[pyo3(signature = (audio: "list[numpy.typing.NDArray[numpy.float64]]", params: "ILDSpectrogramParams", dtype: "str" = None), text_signature = "(audio: list[numpy.typing.NDArray[numpy.float64]], params: ILDSpectrogramParams, dtype: str = \"float64\") -> IldSpectrogram")]
fn py_compute_ild_spectrogram(
py: Python<'_>,
audio: [Bound<'_, PyAny>; 2],
params: &PyILDSpectrogramParams,
dtype: Option<&str>,
) -> PyResult<PyIldSpectrogram> {
fn run<T: PyScalar>(
py: Python<'_>,
audio: &[Bound<'_, PyAny>; 2],
params: &PyILDSpectrogramParams,
) -> PyResult<PyIldSpectrogram>
where
IldSpectrogram<T>: Into<IldInner>,
{
let mut plan = StftPlan::<T>::new(¶ms.inner.spectrogram_params).map_err(|e| {
PyErr::new::<pyo3::exceptions::PyRuntimeError, _>(format!(
"Failed to create STFT plan: {e}"
))
})?;
let (left, right) = read_stereo::<T>(py, audio)?;
let (left_s, right_s) = as_stereo_slices(&left, &right)?;
let ild =
compute_ild_spectrogram([left_s, right_s], ¶ms.inner, &mut plan).map_err(|e| {
PyErr::new::<pyo3::exceptions::PyRuntimeError, _>(format!(
"Failed to compute ILD spectrogram: {e}"
))
})?;
Ok(PyIldSpectrogram::from_result(py, ild))
}
match parse_dtype(dtype)? {
Dtype::F32 => run::<f32>(py, &audio, params),
Dtype::F64 => run::<f64>(py, &audio, params),
}
}
#[pyclass(name = "ILRSpectrogramParams", from_py_object)]
#[derive(Debug, Clone)]
pub struct PyILRSpectrogramParams {
pub(crate) inner: ILRSpectrogramParams,
}
#[pymethods]
impl PyILRSpectrogramParams {
#[new]
#[pyo3(signature = (spectrogram_params: "SpectrogramParams", start_freq: "float" = 1700.0, end_freq: "float" = 4600.0), text_signature = "(spectrogram_params: SpectrogramParams, start_freq: float = 1700.0, end_freq: float = 4600.0) -> ILRSpectrogramParams")]
fn new(
spectrogram_params: PySpectrogramParams,
start_freq: Option<f64>,
end_freq: Option<f64>,
) -> PyResult<Self> {
let inner = ILRSpectrogramParams::new(
spectrogram_params.into(),
start_freq.unwrap_or(1700.0),
end_freq.unwrap_or(4600.0),
)
.map_err(|e| PyErr::new::<pyo3::exceptions::PyValueError, _>(format!("{e}")))?;
Ok(Self { inner })
}
#[getter]
fn spectrogram_params(&self) -> PySpectrogramParams {
PySpectrogramParams::from(self.inner.spectrogram_params.clone())
}
#[getter]
const fn start_freq(&self) -> f64 {
self.inner.start_freq
}
#[getter]
const fn end_freq(&self) -> f64 {
self.inner.end_freq
}
}
#[pyfunction(name = "compute_ilr_spectrogram")]
#[pyo3(signature = (audio, params, dtype=None), text_signature = "(audio: list[numpy.typing.NDArray[numpy.float64]], params: ILRSpectrogramParams, dtype: str = \"float64\") -> IlrSpectrogram")]
fn py_compute_ilr_spectrogram(
py: Python<'_>,
audio: [Bound<'_, PyAny>; 2],
params: &PyILRSpectrogramParams,
dtype: Option<&str>,
) -> PyResult<PyIlrSpectrogram> {
fn run<T: PyScalar>(
py: Python<'_>,
audio: &[Bound<'_, PyAny>; 2],
params: &PyILRSpectrogramParams,
) -> PyResult<PyIlrSpectrogram>
where
IlrSpectrogram<T>: Into<IlrInner>,
{
let mut plan = StftPlan::<T>::new(¶ms.inner.spectrogram_params).map_err(|e| {
PyErr::new::<pyo3::exceptions::PyRuntimeError, _>(format!(
"Failed to create STFT plan: {e}"
))
})?;
let (left, right) = read_stereo::<T>(py, audio)?;
let (left_s, right_s) = as_stereo_slices(&left, &right)?;
let ilr =
compute_ilr_spectrogram([left_s, right_s], ¶ms.inner, &mut plan).map_err(|e| {
PyErr::new::<pyo3::exceptions::PyRuntimeError, _>(format!(
"Failed to compute ILR spectrogram: {e}"
))
})?;
Ok(PyIlrSpectrogram::from_result(py, ilr))
}
match parse_dtype(dtype)? {
Dtype::F32 => run::<f32>(py, &audio, params),
Dtype::F64 => run::<f64>(py, &audio, params),
}
}
#[pyfunction(name = "compute_itd_spectrogram_diff")]
#[pyo3(signature = (reference, test, params, dtype=None), text_signature = "(reference: list[numpy.typing.NDArray[numpy.float64]], test: list[numpy.typing.NDArray[numpy.float64]], params: ITDSpectrogramParams, dtype: str = \"float64\") -> tuple[numpy.typing.NDArray[numpy.float64], float, float]")]
fn py_compute_itd_spectrogram_diff(
py: Python<'_>,
reference: [Bound<'_, PyAny>; 2],
test: [Bound<'_, PyAny>; 2],
params: &PyITDSpectrogramParams,
dtype: Option<&str>,
) -> PyResult<(Py<PyAny>, f64, f64)> {
fn run<T: PyScalar>(
py: Python<'_>,
reference: &[Bound<'_, PyAny>; 2],
test: &[Bound<'_, PyAny>; 2],
params: &PyITDSpectrogramParams,
) -> PyResult<(Py<PyAny>, f64, f64)> {
let mut plan = StftPlan::<T>::new(¶ms.inner.spectrogram_params).map_err(|e| {
PyErr::new::<pyo3::exceptions::PyRuntimeError, _>(format!(
"Failed to create STFT plan: {e}"
))
})?;
let (lr, rr) = read_stereo::<T>(py, reference)?;
let (lt, rt) = read_stereo::<T>(py, test)?;
let (lr_s, rr_s) = as_stereo_slices(&lr, &rr)?;
let (lt_s, rt_s) = as_stereo_slices(<, &rt)?;
let (time_diff, mean_deg, mean_itd) =
compute_itd_spectrogram_diff([lr_s, rr_s], [lt_s, rt_s], ¶ms.inner, &mut plan)
.map_err(|e| {
PyErr::new::<pyo3::exceptions::PyRuntimeError, _>(format!(
"Failed to compute ITD diff: {e}"
))
})?;
Ok((
array1_to_py(py, time_diff),
mean_deg.to_f64().unwrap_or(f64::NAN),
mean_itd.to_f64().unwrap_or(f64::NAN),
))
}
match parse_dtype(dtype)? {
Dtype::F32 => run::<f32>(py, &reference, &test, params),
Dtype::F64 => run::<f64>(py, &reference, &test, params),
}
}
#[pyfunction(name = "compute_ilr_spectrogram_diff")]
#[pyo3(signature = (reference, test, params, dtype=None), text_signature = "(reference: list[numpy.typing.NDArray[numpy.float64]], test: list[numpy.typing.NDArray[numpy.float64]], params: ILRSpectrogramParams, dtype: str = \"float64\") -> tuple[numpy.typing.NDArray[numpy.float64], float]")]
fn py_compute_ilr_spectrogram_diff(
py: Python<'_>,
reference: [Bound<'_, PyAny>; 2],
test: [Bound<'_, PyAny>; 2],
params: &PyILRSpectrogramParams,
dtype: Option<&str>,
) -> PyResult<(Py<PyAny>, f64)> {
fn run<T: PyScalar>(
py: Python<'_>,
reference: &[Bound<'_, PyAny>; 2],
test: &[Bound<'_, PyAny>; 2],
params: &PyILRSpectrogramParams,
) -> PyResult<(Py<PyAny>, f64)> {
let mut plan = StftPlan::<T>::new(¶ms.inner.spectrogram_params).map_err(|e| {
PyErr::new::<pyo3::exceptions::PyRuntimeError, _>(format!(
"Failed to create STFT plan: {e}"
))
})?;
let (lr, rr) = read_stereo::<T>(py, reference)?;
let (lt, rt) = read_stereo::<T>(py, test)?;
let (lr_s, rr_s) = as_stereo_slices(&lr, &rr)?;
let (lt_s, rt_s) = as_stereo_slices(<, &rt)?;
let (time_diff, mean_diff) =
compute_ilr_spectrogram_diff([lr_s, rr_s], [lt_s, rt_s], ¶ms.inner, &mut plan)
.map_err(|e| {
PyErr::new::<pyo3::exceptions::PyRuntimeError, _>(format!(
"Failed to compute ILR diff: {e}"
))
})?;
Ok((
array1_to_py(py, time_diff),
mean_diff.to_f64().unwrap_or(f64::NAN),
))
}
match parse_dtype(dtype)? {
Dtype::F32 => run::<f32>(py, &reference, &test, params),
Dtype::F64 => run::<f64>(py, &reference, &test, params),
}
}
pub fn register(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_function(wrap_pyfunction!(py_compute_itd_spectrogram, m)?)?;
m.add_function(wrap_pyfunction!(py_compute_itd_spectrogram_diff, m)?)?;
m.add_class::<PyITDSpectrogramParams>()?;
m.add_class::<PyItdSpectrogram>()?;
m.add_function(wrap_pyfunction!(py_compute_ipd_spectrogram, m)?)?;
m.add_class::<PyIPDSpectrogramParams>()?;
m.add_class::<PyIpdSpectrogram>()?;
m.add_function(wrap_pyfunction!(py_compute_ild_spectrogram, m)?)?;
m.add_class::<PyILDSpectrogramParams>()?;
m.add_class::<PyIldSpectrogram>()?;
m.add_function(wrap_pyfunction!(py_compute_ilr_spectrogram, m)?)?;
m.add_function(wrap_pyfunction!(py_compute_ilr_spectrogram_diff, m)?)?;
m.add_class::<PyILRSpectrogramParams>()?;
m.add_class::<PyIlrSpectrogram>()?;
Ok(())
}