use crate::EvaluationError;
use oxionnx::{OptLevel, Session, Tensor};
use std::path::{Path, PathBuf};
use std::sync::{Arc, RwLock};
use tracing::{debug, info, warn};
#[derive(Debug, Clone)]
pub struct OnnxMosPredictorConfig {
pub model_path: PathBuf,
pub sample_rate: u32,
pub max_length: usize,
pub opt_level: OptLevel,
pub enable_profiling: bool,
pub enable_memory_pool: bool,
}
impl Default for OnnxMosPredictorConfig {
fn default() -> Self {
Self {
model_path: PathBuf::new(),
sample_rate: 16000,
max_length: 16000 * 10,
opt_level: OptLevel::All,
enable_profiling: false,
enable_memory_pool: false,
}
}
}
pub struct OnnxMosPredictor {
session: Arc<RwLock<Session>>,
config: OnnxMosPredictorConfig,
}
impl OnnxMosPredictor {
pub fn new(config: OnnxMosPredictorConfig) -> Result<Self, EvaluationError> {
info!(
"Loading ONNX MOS predictor model from {:?}",
config.model_path
);
let session = load_session(&config.model_path, &config)?;
info!("ONNX MOS predictor model loaded successfully");
debug!(
"Config: sample_rate={}, max_length={}",
config.sample_rate, config.max_length
);
Ok(Self {
session: Arc::new(RwLock::new(session)),
config,
})
}
pub fn predict_mos(&self, audio_samples: &[f32]) -> Result<f32, EvaluationError> {
if audio_samples.is_empty() {
return Err(EvaluationError::InvalidInput {
message: "Audio samples must not be empty".to_string(),
});
}
let samples = if audio_samples.len() > self.config.max_length {
warn!(
"Audio length {} exceeds max_length {}, truncating",
audio_samples.len(),
self.config.max_length
);
&audio_samples[..self.config.max_length]
} else {
audio_samples
};
let num_samples = samples.len();
let input = Tensor::new(samples.to_vec(), vec![1, num_samples]);
let mut inputs = std::collections::HashMap::new();
inputs.insert("waveform", input);
let session = self
.session
.read()
.map_err(|e| EvaluationError::ModelError {
message: format!("Session lock poisoned: {e}"),
source: None,
})?;
let outputs = session
.run(&inputs)
.map_err(|e| EvaluationError::ModelError {
message: format!("ONNX inference failed: {e}"),
source: None,
})?;
let mos_tensor = outputs
.get("mos_score")
.ok_or_else(|| EvaluationError::ModelError {
message: "Missing 'mos_score' output from ONNX model".to_string(),
source: None,
})?;
let raw_score =
mos_tensor
.data
.first()
.copied()
.ok_or_else(|| EvaluationError::ModelError {
message: "MOS score output tensor is empty".to_string(),
source: None,
})?;
let score = raw_score.clamp(1.0, 5.0);
debug!(
"MOS prediction: raw={:.4}, clamped={:.4}, samples={}",
raw_score, score, num_samples
);
Ok(score)
}
pub fn sample_rate(&self) -> u32 {
self.config.sample_rate
}
pub fn max_length(&self) -> usize {
self.config.max_length
}
}
fn load_session(path: &Path, config: &OnnxMosPredictorConfig) -> Result<Session, EvaluationError> {
let mut builder = Session::builder()
.with_optimization_level(config.opt_level)
.with_memory_pool(config.enable_memory_pool);
if config.enable_profiling {
builder = builder.with_profiling();
}
builder.load(path).map_err(|e| EvaluationError::ModelError {
message: format!("Failed to load ONNX model from {}: {e}", path.display()),
source: None,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_default_config() {
let config = OnnxMosPredictorConfig::default();
assert_eq!(config.sample_rate, 16000);
assert_eq!(config.max_length, 160000);
assert!(!config.enable_profiling);
assert!(!config.enable_memory_pool);
}
}