mod classical;
use std::path::PathBuf;
pub use classical::process_classical;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Backend {
Classical,
#[cfg(feature = "rnnoise")]
Rnnoise,
#[cfg(feature = "deepfilter")]
DeepFilter,
#[cfg(feature = "onnx")]
Onnx,
#[cfg(feature = "mpsenet")]
MpSenet,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct OnnxModelConfig {
pub path: PathBuf,
pub sample_rate: u32,
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct BackendOptions {
pub onnx: Option<OnnxModelConfig>,
}
impl Backend {
pub fn parse(s: &str) -> Option<Self> {
Some(match s.to_ascii_lowercase().as_str() {
"classical" | "dsp" | "stft" => Backend::Classical,
#[cfg(feature = "rnnoise")]
"rnnoise" | "rnn" => Backend::Rnnoise,
#[cfg(feature = "deepfilter")]
"deepfilter" | "deepfilternet" | "dfn" | "dfn3" => Backend::DeepFilter,
#[cfg(feature = "onnx")]
"onnx" | "model" => Backend::Onnx,
#[cfg(feature = "mpsenet")]
"mpsenet" | "mp-senet" | "mp_senet" => Backend::MpSenet,
#[cfg(not(feature = "rnnoise"))]
"rnnoise" | "rnn" => return None,
#[cfg(not(feature = "deepfilter"))]
"deepfilter" | "deepfilternet" | "dfn" | "dfn3" => return None,
#[cfg(not(feature = "onnx"))]
"onnx" | "model" => return None,
#[cfg(not(feature = "mpsenet"))]
"mpsenet" | "mp-senet" | "mp_senet" => return None,
_ => return None,
})
}
pub fn available_names() -> &'static [&'static str] {
&[
"classical",
#[cfg(feature = "rnnoise")]
"rnnoise",
#[cfg(feature = "deepfilter")]
"deepfilter",
#[cfg(feature = "onnx")]
"onnx",
#[cfg(feature = "mpsenet")]
"mpsenet",
]
}
}
pub fn process_channels(
backend: Backend,
channels: &[Vec<f64>],
sample_rate: u32,
classical_cfg: &crate::denoiser::DenoiserConfig,
backend_options: &BackendOptions,
) -> Result<Vec<Vec<f64>>, String> {
let _ = sample_rate; let _ = backend_options; match backend {
Backend::Classical => Ok(process_classical(channels, classical_cfg)),
#[cfg(feature = "rnnoise")]
Backend::Rnnoise => rnnoise::process(channels, sample_rate),
#[cfg(feature = "deepfilter")]
Backend::DeepFilter => deepfilter::process(channels, sample_rate),
#[cfg(feature = "onnx")]
Backend::Onnx => {
let config = backend_options.onnx.as_ref().ok_or_else(|| {
"ONNX backend requires a model path (CLI: --onnx-model <PATH>)".to_string()
})?;
onnx::process(channels, sample_rate, config)
}
#[cfg(feature = "mpsenet")]
Backend::MpSenet => {
let config = backend_options.onnx.as_ref().ok_or_else(|| {
"MP-SENet backend requires a converted model (CLI: --onnx-model <PATH>)".to_string()
})?;
mpsenet::process(channels, sample_rate, config)
}
}
}
#[cfg(feature = "rnnoise")]
pub mod rnnoise;
#[cfg(feature = "deepfilter")]
pub mod deepfilter;
#[cfg(feature = "onnx")]
pub mod onnx;
#[cfg(feature = "mpsenet")]
pub mod mpsenet;
#[cfg(all(test, feature = "mpsenet"))]
mod tests {
use super::*;
#[test]
fn parses_mp_senet_aliases() {
assert_eq!(Backend::parse("mpsenet"), Some(Backend::MpSenet));
assert_eq!(Backend::parse("mp-senet"), Some(Backend::MpSenet));
assert!(Backend::available_names().contains(&"mpsenet"));
}
}