mod mel;
use std::path::Path;
use crate::{
ComputeUnits, DataType, Model, MultiArray,
model::contract::{Checked, Dim, FeatureContract, LoadContract, StateContract},
};
use crate::embeddings::clap::{
embedding::{EMBEDDING_DIM, Embedding, check_finite_output},
error::{AudioTooLong, Error, OutputShape, Result, WinditError, contract_violation},
window::{WindowEmbedding, WindowPlan},
};
pub use self::mel::{N_MELS, SAMPLE_RATE_HZ, T_FRAMES, TARGET_SAMPLES};
mod names {
pub const INPUT_FEATURES: &str = "input_features";
pub const AUDIO_EMBEDS: &str = "audio_embeds";
}
pub const DEFAULT_AUDIO_COMPUTE: ComputeUnits = ComputeUnits::All;
#[cfg(feature = "serde")]
fn default_audio_compute() -> ComputeUnits {
DEFAULT_AUDIO_COMPUTE
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
pub struct AudioEncoderOptions {
#[cfg_attr(feature = "serde", serde(default = "default_audio_compute"))]
compute: ComputeUnits,
}
impl Default for AudioEncoderOptions {
fn default() -> Self {
Self::new()
}
}
impl AudioEncoderOptions {
pub const fn new() -> Self {
Self {
compute: DEFAULT_AUDIO_COMPUTE,
}
}
#[inline]
pub const fn compute(&self) -> ComputeUnits {
self.compute
}
#[must_use]
#[inline]
pub const fn with_compute(mut self, compute: ComputeUnits) -> Self {
self.set_compute(compute);
self
}
#[inline]
pub const fn set_compute(&mut self, compute: ComputeUnits) -> &mut Self {
self.compute = compute;
self
}
}
#[derive(Debug)]
pub struct AudioEncoder {
model: Checked,
mel: mel::MelExtractor,
}
impl AudioEncoder {
pub fn from_file(path: impl AsRef<Path>) -> Result<Self> {
Self::from_file_with(path, AudioEncoderOptions::new())
}
pub fn from_file_with(path: impl AsRef<Path>, options: AudioEncoderOptions) -> Result<Self> {
let model = Model::load(path, options.compute())?;
let model = Checked::new(model, &audio_contract()).map_err(contract_violation)?;
Ok(Self {
model,
mel: mel::MelExtractor::new(),
})
}
pub fn embed_window(&self, samples: &[f32]) -> Result<Embedding> {
if samples.is_empty() {
return Err(Error::EmptyAudio);
}
check_window_len(samples.len())?;
if let Some(index) = first_non_finite(samples) {
return Err(Error::NonFiniteInput(index));
}
let mut features = vec![0.0f32; N_MELS * T_FRAMES];
self.mel.extract_into(samples, &mut features)?;
let input = MultiArray::from_slice(&[1, 1, T_FRAMES, N_MELS], &features)?;
let mut outputs = self
.model
.predict_with(&[(names::INPUT_FEATURES, &input)])?;
let embeds = outputs
.take(names::AUDIO_EMBEDS)
.ok_or_else(|| crate::PredictionError::MissingOutput(names::AUDIO_EMBEDS.to_string()))?;
if embeds.shape() != [1, EMBEDDING_DIM] {
return Err(Error::OutputShape(OutputShape::new(
embeds.shape().to_vec(),
vec![1, EMBEDDING_DIM],
)));
}
let mut row = [0.0f32; EMBEDDING_DIM];
embeds.copy_into::<f32>(&mut row)?;
check_finite_output(&row)?;
Embedding::from_slice_normalizing(&row)
}
pub fn embed_windows(&self, samples: &[f32], plan: &WindowPlan) -> Result<Vec<WindowEmbedding>> {
if samples.is_empty() {
return Err(Error::EmptyAudio);
}
let spans = plan.spans(samples.len())?;
let mut out = Vec::new();
out.try_reserve_exact(spans.len()).map_err(|_| {
Error::Windowing(WinditError::AllocFailed {
elements: spans.len(),
})
})?;
for span in spans {
let embedding = self.embed_window(&samples[span.start()..span.end()])?;
out.push(WindowEmbedding::new(embedding, span));
}
Ok(out)
}
pub fn prewarm(&self) -> Result<()> {
let signal: Vec<f32> = (0..SAMPLE_RATE_HZ)
.map(|i| 0.5 * (std::f32::consts::TAU * 440.0 * (i as f32 / SAMPLE_RATE_HZ as f32)).sin())
.collect();
self.embed_window(&signal)?;
Ok(())
}
}
fn audio_contract() -> LoadContract {
LoadContract::new(
vec![FeatureContract::new(
names::INPUT_FEATURES,
DataType::F32,
vec![
Dim::Exactly(1),
Dim::Exactly(1),
Dim::Exactly(T_FRAMES),
Dim::Exactly(N_MELS),
],
)],
vec![FeatureContract::new(
names::AUDIO_EMBEDS,
DataType::F32,
vec![Dim::Exactly(1), Dim::Exactly(EMBEDDING_DIM)],
)],
StateContract::None,
)
}
fn first_non_finite(samples: &[f32]) -> Option<usize> {
samples.iter().position(|v| !v.is_finite())
}
fn check_window_len(len: usize) -> Result<()> {
if len > TARGET_SAMPLES {
return Err(Error::AudioTooLong(AudioTooLong::new(len, TARGET_SAMPLES)));
}
Ok(())
}
#[cfg(test)]
mod tests;