use std::sync::{Arc, Mutex};
use crate::audio::whisper::backend::{AlignmentView, BackendError, InferenceBackend, ModelDims};
#[cfg(test)]
mod tests;
#[derive(Debug, Clone, PartialEq)]
struct ScriptedStep {
logits: Vec<f32>,
alignment_row: Option<Vec<f32>>,
}
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub struct MockCounters {
extract_calls: usize,
encode_calls: usize,
decode_steps: usize,
resets: usize,
}
impl MockCounters {
#[inline(always)]
pub const fn extract_calls(&self) -> usize {
self.extract_calls
}
#[inline(always)]
pub const fn encode_calls(&self) -> usize {
self.encode_calls
}
#[inline(always)]
pub const fn decode_steps(&self) -> usize {
self.decode_steps
}
#[inline(always)]
pub const fn resets(&self) -> usize {
self.resets
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct MockDecoderState {
step: usize,
consumed: Vec<(u32, usize)>,
alignment: Vec<f32>,
pending: Option<(usize, Vec<f32>)>,
window_has_alignment: bool,
}
impl MockDecoderState {
#[inline(always)]
pub fn consumed_slice(&self) -> &[(u32, usize)] {
self.consumed.as_slice()
}
}
#[derive(Debug)]
pub struct MockBackend {
dims: ModelDims,
script: Vec<ScriptedStep>,
counters: Arc<Mutex<MockCounters>>,
fail_calls: Vec<usize>,
attempted: Arc<Mutex<usize>>,
continuous_script: bool,
}
impl Default for MockBackend {
fn default() -> Self {
Self::new()
}
}
impl MockBackend {
pub fn new() -> Self {
Self {
dims: ModelDims::new(),
script: Vec::new(),
counters: Arc::new(Mutex::new(MockCounters::default())),
fail_calls: Vec::new(),
attempted: Arc::new(Mutex::new(0)),
continuous_script: false,
}
}
#[must_use]
#[inline(always)]
pub fn with_continuous_script(mut self) -> Self {
self.set_continuous_script(true);
self
}
#[inline(always)]
pub fn set_continuous_script(&mut self, continuous: bool) -> &mut Self {
self.continuous_script = continuous;
self
}
pub fn fail_on_call(&mut self, call: usize) -> &mut Self {
self.fail_calls.push(call);
self
}
#[must_use]
#[inline(always)]
pub fn with_dims(mut self, dims: ModelDims) -> Self {
self.set_dims(dims);
self
}
#[inline(always)]
pub fn set_dims(&mut self, dims: ModelDims) -> &mut Self {
self.dims = dims;
self
}
pub fn push_step(&mut self, logits: Vec<f32>) -> &mut Self {
assert_eq!(
logits.len(),
self.dims.vocab(),
"scripted logits len {} != dims.vocab() {}",
logits.len(),
self.dims.vocab()
);
self.script.push(ScriptedStep {
logits,
alignment_row: None,
});
self
}
pub fn push_token_step(&mut self, token: u32) -> &mut Self {
let mut logits = vec![0.0_f32; self.dims.vocab()];
logits[token as usize] = 10.0;
self.push_step(logits)
}
pub fn push_token_steps(&mut self, tokens: &[u32]) -> &mut Self {
for &token in tokens {
self.push_token_step(token);
}
self
}
pub fn push_step_with_alignment(
&mut self,
logits: Vec<f32>,
alignment_row: Vec<f32>,
) -> &mut Self {
self.push_step(logits);
self
.script
.last_mut()
.expect("push_step just appended one entry")
.alignment_row = Some(alignment_row);
self
}
pub fn counters(&self) -> MockCounters {
*self
.counters
.lock()
.expect("mock backend counters lock poisoned")
}
}
impl InferenceBackend for MockBackend {
type Features = Vec<f32>;
type EncoderOutput = Vec<f32>;
type DecoderState = MockDecoderState;
fn extract_features(&self, audio: &[f32]) -> Result<Self::Features, BackendError> {
self
.counters
.lock()
.expect("mock backend counters lock poisoned")
.extract_calls += 1;
Ok(audio.to_vec())
}
fn encode(&self, features: &Self::Features) -> Result<Self::EncoderOutput, BackendError> {
self
.counters
.lock()
.expect("mock backend counters lock poisoned")
.encode_calls += 1;
Ok(features.clone())
}
fn new_decoder_state(&self) -> Result<Self::DecoderState, BackendError> {
Ok(MockDecoderState {
step: 0,
consumed: Vec::new(),
alignment: vec![0.0; (self.dims.max_token_context() + 1) * self.dims.n_audio_ctx()],
pending: None,
window_has_alignment: false,
})
}
fn reset_decoder_state(&self, state: &mut Self::DecoderState) {
if !self.continuous_script {
state.step = 0;
}
state.consumed.clear();
state.window_has_alignment = false;
state.pending = None;
self
.counters
.lock()
.expect("mock backend counters lock poisoned")
.resets += 1;
}
fn decode_step(
&self,
token: u32,
position: usize,
_encoder_output: &Self::EncoderOutput,
state: &mut Self::DecoderState,
logits: &mut Vec<f32>,
) -> Result<(), BackendError> {
let call = {
let mut attempted = self
.attempted
.lock()
.expect("mock backend attempted-call lock poisoned");
*attempted += 1;
*attempted
};
if self.fail_calls.contains(&call) {
return Err(BackendError::ScriptedFailure(call));
}
let Some(scripted) = self.script.get(state.step) else {
return Err(BackendError::ScriptExhausted(state.step));
};
state.consumed.push((token, position));
logits.clear();
logits.extend_from_slice(&scripted.logits);
if position == 0 {
state.window_has_alignment = false;
state.pending = None;
}
state.pending = scripted
.alignment_row
.as_ref()
.map(|row| (position, row.clone()));
state.step += 1;
self
.counters
.lock()
.expect("mock backend counters lock poisoned")
.decode_steps += 1;
Ok(())
}
fn commit_alignment_row(&self, state: &mut Self::DecoderState) {
let Some((position, row)) = state.pending.take() else {
return;
};
let cols = self.dims.n_audio_ctx();
let start = (position + 1) * cols;
state.alignment[start..start + cols].copy_from_slice(&row);
state.window_has_alignment = true;
}
fn alignment_weights<'state>(
&self,
state: &'state Self::DecoderState,
) -> Option<AlignmentView<'state>> {
state.window_has_alignment.then(|| {
let cols = self.dims.n_audio_ctx();
AlignmentView::new(&state.alignment, self.dims.max_token_context() + 1, cols)
})
}
fn dims(&self) -> ModelDims {
self.dims
}
}