use crate::audio::whisper::model::is_model_multilingual;
pub mod coreml;
pub mod mock;
#[cfg(test)]
mod tests;
pub const DEFAULT_VOCAB: usize = 51_865;
pub const DEFAULT_N_MELS: usize = 80;
pub const DEFAULT_EMBED_DIM: usize = 384;
pub const DEFAULT_KV_DIM: usize = 1536;
pub const DEFAULT_MAX_TOKEN_CONTEXT: usize = crate::audio::whisper::constants::MAX_TOKEN_CONTEXT;
pub const DEFAULT_N_AUDIO_CTX: usize = 1500;
pub const DEFAULT_WINDOW_SAMPLES: usize = crate::audio::whisper::constants::WINDOW_SAMPLES;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ModelDims {
vocab: usize,
n_mels: usize,
embed_dim: usize,
kv_dim: usize,
max_token_context: usize,
n_audio_ctx: usize,
window_samples: usize,
}
impl Default for ModelDims {
fn default() -> Self {
Self::new()
}
}
impl ModelDims {
pub const fn new() -> Self {
Self {
vocab: DEFAULT_VOCAB,
n_mels: DEFAULT_N_MELS,
embed_dim: DEFAULT_EMBED_DIM,
kv_dim: DEFAULT_KV_DIM,
max_token_context: DEFAULT_MAX_TOKEN_CONTEXT,
n_audio_ctx: DEFAULT_N_AUDIO_CTX,
window_samples: DEFAULT_WINDOW_SAMPLES,
}
}
#[inline(always)]
pub const fn vocab(&self) -> usize {
self.vocab
}
#[must_use]
#[inline(always)]
pub const fn with_vocab(mut self, vocab: usize) -> Self {
self.set_vocab(vocab);
self
}
#[inline(always)]
pub const fn set_vocab(&mut self, vocab: usize) -> &mut Self {
self.vocab = vocab;
self
}
#[inline(always)]
pub const fn n_mels(&self) -> usize {
self.n_mels
}
#[must_use]
#[inline(always)]
pub const fn with_n_mels(mut self, n_mels: usize) -> Self {
self.set_n_mels(n_mels);
self
}
#[inline(always)]
pub const fn set_n_mels(&mut self, n_mels: usize) -> &mut Self {
self.n_mels = n_mels;
self
}
#[inline(always)]
pub const fn embed_dim(&self) -> usize {
self.embed_dim
}
#[must_use]
#[inline(always)]
pub const fn with_embed_dim(mut self, embed_dim: usize) -> Self {
self.set_embed_dim(embed_dim);
self
}
#[inline(always)]
pub const fn set_embed_dim(&mut self, embed_dim: usize) -> &mut Self {
self.embed_dim = embed_dim;
self
}
#[inline(always)]
pub const fn kv_dim(&self) -> usize {
self.kv_dim
}
#[must_use]
#[inline(always)]
pub const fn with_kv_dim(mut self, kv_dim: usize) -> Self {
self.set_kv_dim(kv_dim);
self
}
#[inline(always)]
pub const fn set_kv_dim(&mut self, kv_dim: usize) -> &mut Self {
self.kv_dim = kv_dim;
self
}
#[inline(always)]
pub const fn max_token_context(&self) -> usize {
self.max_token_context
}
#[must_use]
#[inline(always)]
pub const fn with_max_token_context(mut self, max_token_context: usize) -> Self {
self.set_max_token_context(max_token_context);
self
}
#[inline(always)]
pub const fn set_max_token_context(&mut self, max_token_context: usize) -> &mut Self {
self.max_token_context = max_token_context;
self
}
#[inline(always)]
pub const fn n_audio_ctx(&self) -> usize {
self.n_audio_ctx
}
#[must_use]
#[inline(always)]
pub const fn with_n_audio_ctx(mut self, n_audio_ctx: usize) -> Self {
self.set_n_audio_ctx(n_audio_ctx);
self
}
#[inline(always)]
pub const fn set_n_audio_ctx(&mut self, n_audio_ctx: usize) -> &mut Self {
self.n_audio_ctx = n_audio_ctx;
self
}
#[inline(always)]
pub const fn window_samples(&self) -> usize {
self.window_samples
}
#[must_use]
#[inline(always)]
pub const fn with_window_samples(mut self, window_samples: usize) -> Self {
self.set_window_samples(window_samples);
self
}
#[inline(always)]
pub const fn set_window_samples(&mut self, window_samples: usize) -> &mut Self {
self.window_samples = window_samples;
self
}
#[inline(always)]
pub const fn is_multilingual(&self) -> bool {
is_model_multilingual(self.vocab)
}
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct AlignmentView<'a> {
data: &'a [f32],
rows: usize,
cols: usize,
}
impl<'a> AlignmentView<'a> {
pub fn new(data: &'a [f32], rows: usize, cols: usize) -> Self {
assert_eq!(
data.len(),
rows * cols,
"AlignmentView: data.len() ({}) != rows ({rows}) * cols ({cols})",
data.len()
);
Self { data, rows, cols }
}
#[inline(always)]
pub const fn rows(&self) -> usize {
self.rows
}
#[inline(always)]
pub const fn cols(&self) -> usize {
self.cols
}
#[inline(always)]
pub fn row(&self, index: usize) -> &'a [f32] {
&self.data[index * self.cols..(index + 1) * self.cols]
}
#[inline(always)]
pub const fn data(&self) -> &'a [f32] {
self.data
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct AlignmentMatrix {
data: Vec<f32>,
rows: usize,
cols: usize,
}
impl AlignmentMatrix {
pub fn new(data: Vec<f32>, rows: usize, cols: usize) -> Self {
assert_eq!(
data.len(),
rows * cols,
"AlignmentMatrix: data.len() ({}) != rows ({rows}) * cols ({cols})",
data.len()
);
Self { data, rows, cols }
}
#[inline(always)]
pub const fn rows(&self) -> usize {
self.rows
}
#[inline(always)]
pub const fn cols(&self) -> usize {
self.cols
}
#[inline(always)]
pub fn view(&self) -> AlignmentView<'_> {
AlignmentView::new(&self.data, self.rows, self.cols)
}
}
impl AlignmentView<'_> {
pub fn to_matrix(&self) -> AlignmentMatrix {
AlignmentMatrix::new(self.data.to_vec(), self.rows, self.cols)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct MissingFeature {
model: &'static str,
name: &'static str,
}
impl MissingFeature {
#[inline(always)]
pub const fn new(model: &'static str, name: &'static str) -> Self {
Self { model, name }
}
#[inline(always)]
pub const fn model(&self) -> &'static str {
self.model
}
#[inline(always)]
pub const fn name(&self) -> &'static str {
self.name
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct AudioLength {
got: usize,
expected: usize,
}
impl AudioLength {
#[inline(always)]
pub const fn new(got: usize, expected: usize) -> Self {
Self { got, expected }
}
#[inline(always)]
pub const fn got(&self) -> usize {
self.got
}
#[inline(always)]
pub const fn expected(&self) -> usize {
self.expected
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ContractMismatch {
model: &'static str,
feature: String,
expected: String,
actual: String,
}
impl ContractMismatch {
#[inline(always)]
pub const fn new(model: &'static str, feature: String, expected: String, actual: String) -> Self {
Self {
model,
feature,
expected,
actual,
}
}
#[inline(always)]
pub const fn model(&self) -> &'static str {
self.model
}
#[inline(always)]
pub fn feature(&self) -> &str {
&self.feature
}
#[inline(always)]
pub fn expected(&self) -> &str {
&self.expected
}
#[inline(always)]
pub fn actual(&self) -> &str {
&self.actual
}
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum BackendError {
#[error("backend prediction failed: {0}")]
Prediction(#[from] crate::PredictionError),
#[error("backend tensor failed: {0}")]
Tensor(#[from] crate::TensorError),
#[error("{} model output is missing feature `{}`", .0.model(), .0.name())]
MissingFeature(MissingFeature),
#[error(
"{} model contract mismatch on `{}`: expected {}, got {}",
.0.model(), .0.feature(), .0.expected(), .0.actual()
)]
Contract(ContractMismatch),
#[error("audio window has {} samples, backend expects {}", .0.got(), .0.expected())]
AudioLength(AudioLength),
#[error("audio window has a non-finite sample at index {0}")]
NonFiniteAudio(usize),
#[error(
"audio window has a sample at index {0} that is finite in f32 but overflows the mel model's f16 input domain (|x| > f16::MAX)"
)]
F16OverflowAudio(usize),
#[error("mock script exhausted at step {0}")]
ScriptExhausted(usize),
#[error("scripted decode-step failure on call {0}")]
ScriptedFailure(usize),
}
pub trait InferenceBackend {
type Features;
type EncoderOutput;
type DecoderState;
fn extract_features(&self, audio: &[f32]) -> Result<Self::Features, BackendError>;
fn encode(&self, features: &Self::Features) -> Result<Self::EncoderOutput, BackendError>;
fn new_decoder_state(&self) -> Result<Self::DecoderState, BackendError>;
fn reset_decoder_state(&self, state: &mut Self::DecoderState);
fn decode_step(
&self,
token: u32,
position: usize,
encoder_output: &Self::EncoderOutput,
state: &mut Self::DecoderState,
logits: &mut Vec<f32>,
) -> Result<(), BackendError>;
fn commit_alignment_row(&self, state: &mut Self::DecoderState);
fn alignment_weights<'state>(
&self,
state: &'state Self::DecoderState,
) -> Option<AlignmentView<'state>>;
fn dims(&self) -> ModelDims;
}