use crate::{DataType, IndexOutOfBounds, Model, MultiArray, TensorError, f16};
use crate::model::contract::{
Checked, ContractViolation, Dim, FeatureContract, LoadContract, Rendered, StateContract,
};
use crate::audio::whisper::{
backend::{
AlignmentView, AudioLength, BackendError, ContractMismatch, InferenceBackend, MissingFeature,
ModelDims,
},
constants::MAX_TOKEN_CONTEXT,
model::manager::LoadedModels,
};
use crate::ModelDescription;
#[cfg(test)]
mod tests;
mod names {
pub const AUDIO: &str = "audio";
pub const MEL: &str = "melspectrogram_features";
pub const ENCODER: &str = "encoder_output_embeds";
pub const INPUT_IDS: &str = "input_ids";
pub const CACHE_LENGTH: &str = "cache_length";
pub const KEY_CACHE: &str = "key_cache";
pub const VALUE_CACHE: &str = "value_cache";
pub const KV_UPDATE_MASK: &str = "kv_cache_update_mask";
pub const PADDING_MASK: &str = "decoder_key_padding_mask";
pub const LOGITS: &str = "logits";
pub const KEY_UPDATES: &str = "key_cache_updates";
pub const VALUE_UPDATES: &str = "value_cache_updates";
pub const ALIGNMENT: &str = "alignment_heads_weights";
}
const PADDING_MASK_HIDDEN: f32 = -10000.0;
fn input_dim(
description: &ModelDescription,
model_name: &'static str,
feature: &'static str,
position: usize,
) -> Result<usize, BackendError> {
description
.input(feature)
.and_then(|f| f.shape().get(position).copied())
.ok_or(BackendError::MissingFeature(MissingFeature::new(
model_name, feature,
)))
}
fn output_dim(
description: &ModelDescription,
model_name: &'static str,
feature: &'static str,
position: usize,
) -> Result<usize, BackendError> {
description
.output(feature)
.and_then(|f| f.shape().get(position).copied())
.ok_or(BackendError::MissingFeature(MissingFeature::new(
model_name, feature,
)))
}
fn audio_input(audio: &[f32], expected: usize) -> Result<MultiArray, BackendError> {
if audio.len() != expected {
return Err(BackendError::AudioLength(AudioLength::new(
audio.len(),
expected,
)));
}
let f16_max = f32::from(f16::MAX);
for (index, &sample) in audio.iter().enumerate() {
if !sample.is_finite() {
return Err(BackendError::NonFiniteAudio(index));
}
if sample.abs() > f16_max {
return Err(BackendError::F16OverflowAudio(index));
}
}
Ok(MultiArray::from_slice(&[expected], audio)?)
}
fn contract_violation(model: &'static str) -> impl Fn(ContractViolation) -> BackendError {
move |violation| {
let (feature, expected, actual) = match violation.rendered() {
Rendered::UnsatisfiableInput(name) => (
name,
"an input this backend sends".to_string(),
"a required input the contract does not name".to_string(),
),
Rendered::UnsatisfiableState(name) => (
name,
"no state buffer".to_string(),
"a declared state buffer".to_string(),
),
Rendered::Feature(feature) => (
feature.feature().to_string(),
feature.clone().expected(),
feature.actual(),
),
};
BackendError::Contract(ContractMismatch::new(model, feature, expected, actual))
}
}
#[derive(Debug)]
pub struct CoreMlDecoderState {
input_ids: MultiArray,
cache_length: MultiArray,
key_cache: MultiArray,
value_cache: MultiArray,
kv_cache_update_mask: MultiArray,
decoder_key_padding_mask: MultiArray,
alignment: Vec<f32>,
pending_alignment: Option<usize>,
window_has_alignment: bool,
kv_scratch: Vec<f16>,
logits_scratch: Vec<f16>,
align_scratch: Vec<f16>,
}
fn append_kv(
cache: &mut MultiArray,
update: &MultiArray,
scratch: &mut Vec<f16>,
kv_dim: usize,
max_ctx: usize,
position: usize,
) -> Result<(), BackendError> {
scratch.resize(kv_dim, f16::ZERO);
update.copy_into::<f16>(scratch)?;
let dst = cache.as_slice_mut::<f16>()?;
for (j, &value) in scratch.iter().enumerate() {
dst[j * max_ctx + position] = value;
}
Ok(())
}
#[derive(Debug)]
pub struct CoreMlBackend {
mel: Checked,
encoder: Checked,
decoder: Checked,
dims: ModelDims,
supports_alignment: bool,
}
impl CoreMlBackend {
pub fn new(
mel: Model,
encoder: Model,
decoder: Model,
vocab: usize,
) -> Result<Self, BackendError> {
let mel = Checked::new(mel, &mel_contract()).map_err(contract_violation("mel"))?;
let window_samples = input_dim(mel.description(), "mel", names::AUDIO, 0)?;
let n_mels = output_dim(mel.description(), "mel", names::MEL, 1)?;
let mel_frames = output_dim(mel.description(), "mel", names::MEL, 3)?;
let encoder = Checked::new(encoder, &encoder_contract(n_mels, mel_frames))
.map_err(contract_violation("encoder"))?;
let embed_dim = output_dim(encoder.description(), "encoder", names::ENCODER, 1)?;
let n_audio_ctx = output_dim(encoder.description(), "encoder", names::ENCODER, 3)?;
let supports_alignment =
output_dim(decoder.description(), "decoder", names::ALIGNMENT, 0).is_ok();
let kv_dim = input_dim(decoder.description(), "decoder", names::KEY_CACHE, 1)?;
let max_token_context = input_dim(decoder.description(), "decoder", names::KEY_CACHE, 3)?;
let decoder = Checked::new(
decoder,
&decoder_contract(
embed_dim,
n_audio_ctx,
kv_dim,
max_token_context,
vocab,
supports_alignment,
),
)
.map_err(contract_violation("decoder"))?;
let dims = ModelDims::new()
.with_window_samples(window_samples)
.with_n_mels(n_mels)
.with_embed_dim(embed_dim)
.with_n_audio_ctx(n_audio_ctx)
.with_kv_dim(kv_dim)
.with_max_token_context(max_token_context)
.with_vocab(vocab);
Ok(Self {
mel,
encoder,
decoder,
dims,
supports_alignment,
})
}
pub fn from_loaded(models: LoadedModels, vocab: usize) -> Result<Self, BackendError> {
let (mel, encoder, decoder) = models.into_parts();
Self::new(mel, encoder, decoder, vocab)
}
#[inline(always)]
pub const fn supports_word_timestamps(&self) -> bool {
self.supports_alignment
}
}
impl InferenceBackend for CoreMlBackend {
type Features = MultiArray;
type EncoderOutput = MultiArray;
type DecoderState = CoreMlDecoderState;
fn extract_features(&self, audio: &[f32]) -> Result<Self::Features, BackendError> {
let array = audio_input(audio, self.dims.window_samples())?;
let mut outputs = self.mel.predict_with(&[(names::AUDIO, &array)])?;
outputs
.take(names::MEL)
.ok_or(BackendError::MissingFeature(MissingFeature::new(
"mel",
names::MEL,
)))
}
fn encode(&self, features: &Self::Features) -> Result<Self::EncoderOutput, BackendError> {
let mut outputs = self.encoder.predict_with(&[(names::MEL, features)])?;
outputs
.take(names::ENCODER)
.ok_or(BackendError::MissingFeature(MissingFeature::new(
"encoder",
names::ENCODER,
)))
}
fn new_decoder_state(&self) -> Result<Self::DecoderState, BackendError> {
let kv_dim = self.dims.kv_dim();
let max_ctx = self.dims.max_token_context();
let input_ids = MultiArray::zeros(&[1], DataType::I32)?;
let cache_length = MultiArray::zeros(&[1], DataType::I32)?;
let key_cache = MultiArray::zeros(&[1, kv_dim, 1, max_ctx], DataType::F16)?;
let value_cache = MultiArray::zeros(&[1, kv_dim, 1, max_ctx], DataType::F16)?;
let mut kv_cache_update_mask = MultiArray::zeros(&[1, max_ctx], DataType::F16)?;
let mut decoder_key_padding_mask = MultiArray::zeros(&[1, max_ctx], DataType::F16)?;
decoder_key_padding_mask
.as_slice_mut::<f16>()?
.fill(f16::from_f32(PADDING_MASK_HIDDEN));
decoder_key_padding_mask.fill_at(&[0, 0], f16::ZERO)?;
kv_cache_update_mask.fill_at(&[0, 0], f16::ONE)?;
Ok(CoreMlDecoderState {
input_ids,
cache_length,
key_cache,
value_cache,
kv_cache_update_mask,
decoder_key_padding_mask,
alignment: vec![0.0; (max_ctx + 1) * self.dims.n_audio_ctx()],
pending_alignment: None,
window_has_alignment: false,
kv_scratch: vec![f16::ZERO; kv_dim],
logits_scratch: vec![f16::ZERO; self.dims.vocab()],
align_scratch: vec![f16::ZERO; self.dims.n_audio_ctx()],
})
}
fn reset_decoder_state(&self, state: &mut Self::DecoderState) {
state
.cache_length
.fill_at(&[0], 0_i32)
.expect("cache_length is a self-allocated contiguous [1] i32 array");
let padding = state
.decoder_key_padding_mask
.as_slice_mut::<f16>()
.expect("padding mask is a self-allocated contiguous f16 array");
padding.fill(f16::from_f32(PADDING_MASK_HIDDEN));
padding[0] = f16::ZERO;
let update = state
.kv_cache_update_mask
.as_slice_mut::<f16>()
.expect("update mask is a self-allocated contiguous f16 array");
update.fill(f16::ZERO);
update[0] = f16::ONE;
state.window_has_alignment = false;
state.pending_alignment = None;
}
fn decode_step(
&self,
token: u32,
position: usize,
encoder_output: &Self::EncoderOutput,
state: &mut Self::DecoderState,
logits: &mut Vec<f32>,
) -> Result<(), BackendError> {
let max_ctx = self.dims.max_token_context();
if position >= max_ctx {
return Err(BackendError::Tensor(TensorError::IndexOutOfBounds(
IndexOutOfBounds::new(position, max_ctx),
)));
}
state.input_ids.fill_at(&[0], token as i32)?;
state.cache_length.fill_at(&[0], position as i32)?;
let mut outputs = self.decoder.predict_with(&[
(names::INPUT_IDS, &state.input_ids),
(names::CACHE_LENGTH, &state.cache_length),
(names::KEY_CACHE, &state.key_cache),
(names::VALUE_CACHE, &state.value_cache),
(names::KV_UPDATE_MASK, &state.kv_cache_update_mask),
(names::ENCODER, encoder_output),
(names::PADDING_MASK, &state.decoder_key_padding_mask),
])?;
let logits_array = outputs
.take(names::LOGITS)
.ok_or(BackendError::MissingFeature(MissingFeature::new(
"decoder",
names::LOGITS,
)))?;
state.logits_scratch.resize(self.dims.vocab(), f16::ZERO);
logits_array.copy_into::<f16>(&mut state.logits_scratch)?;
logits.clear();
logits.extend(state.logits_scratch.iter().map(|v| v.to_f32()));
let key_updates = outputs
.take(names::KEY_UPDATES)
.ok_or(BackendError::MissingFeature(MissingFeature::new(
"decoder",
names::KEY_UPDATES,
)))?;
let value_updates = outputs
.take(names::VALUE_UPDATES)
.ok_or(BackendError::MissingFeature(MissingFeature::new(
"decoder",
names::VALUE_UPDATES,
)))?;
let kv_dim = self.dims.kv_dim();
append_kv(
&mut state.key_cache,
&key_updates,
&mut state.kv_scratch,
kv_dim,
max_ctx,
position,
)?;
append_kv(
&mut state.value_cache,
&value_updates,
&mut state.kv_scratch,
kv_dim,
max_ctx,
position,
)?;
if position + 1 < max_ctx {
state
.decoder_key_padding_mask
.fill_at(&[0, position + 1], f16::ZERO)?;
state
.kv_cache_update_mask
.fill_at(&[0, position], f16::ZERO)?;
state
.kv_cache_update_mask
.fill_at(&[0, position + 1], f16::ONE)?;
}
if position == 0 {
state.window_has_alignment = false;
state.pending_alignment = None;
}
if self.supports_alignment
&& let Some(alignment) = outputs.take(names::ALIGNMENT)
{
let cols = self.dims.n_audio_ctx();
state.align_scratch.resize(cols, f16::ZERO);
alignment.copy_into::<f16>(&mut state.align_scratch)?;
state.pending_alignment = Some(position);
} else {
state.pending_alignment = None;
}
Ok(())
}
fn commit_alignment_row(&self, state: &mut Self::DecoderState) {
let Some(position) = state.pending_alignment.take() else {
return;
};
let cols = self.dims.n_audio_ctx();
let start = (position + 1) * cols;
for (dst, src) in state.alignment[start..start + cols]
.iter_mut()
.zip(&state.align_scratch)
{
*dst = src.to_f32();
}
state.window_has_alignment = true;
}
fn alignment_weights<'state>(
&self,
state: &'state Self::DecoderState,
) -> Option<AlignmentView<'state>> {
(self.supports_alignment && 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
}
}
fn mel_contract() -> LoadContract {
LoadContract::new(
vec![FeatureContract::new(
names::AUDIO,
DataType::F16,
vec![Dim::AnyFixed],
)],
vec![FeatureContract::new(
names::MEL,
DataType::F16,
vec![
Dim::Exactly(1),
Dim::AnyFixed,
Dim::Exactly(1),
Dim::AnyFixed,
],
)],
StateContract::None,
)
}
fn encoder_contract(n_mels: usize, mel_frames: usize) -> LoadContract {
LoadContract::new(
vec![FeatureContract::new(
names::MEL,
DataType::F16,
vec![
Dim::Exactly(1),
Dim::Exactly(n_mels),
Dim::Exactly(1),
Dim::Exactly(mel_frames),
],
)],
vec![FeatureContract::new(
names::ENCODER,
DataType::F16,
vec![
Dim::Exactly(1),
Dim::AnyFixed,
Dim::Exactly(1),
Dim::AnyFixed,
],
)],
StateContract::None,
)
}
fn decoder_contract(
embed_dim: usize,
n_audio_ctx: usize,
kv_dim: usize,
max_token_context: usize,
vocab: usize,
supports_alignment: bool,
) -> LoadContract {
let inputs = vec![
FeatureContract::new(names::INPUT_IDS, DataType::I32, vec![Dim::Exactly(1)]),
FeatureContract::new(names::CACHE_LENGTH, DataType::I32, vec![Dim::Exactly(1)]),
FeatureContract::new(
names::KEY_CACHE,
DataType::F16,
vec![
Dim::Exactly(1),
Dim::AnyFixed,
Dim::Exactly(1),
Dim::AtLeast(MAX_TOKEN_CONTEXT),
],
),
FeatureContract::new(
names::VALUE_CACHE,
DataType::F16,
vec![
Dim::Exactly(1),
Dim::Exactly(kv_dim),
Dim::Exactly(1),
Dim::Exactly(max_token_context),
],
),
FeatureContract::new(
names::KV_UPDATE_MASK,
DataType::F16,
vec![Dim::Exactly(1), Dim::Exactly(max_token_context)],
),
FeatureContract::new(
names::ENCODER,
DataType::F16,
vec![
Dim::Exactly(1),
Dim::Exactly(embed_dim),
Dim::Exactly(1),
Dim::Exactly(n_audio_ctx),
],
),
FeatureContract::new(
names::PADDING_MASK,
DataType::F16,
vec![Dim::Exactly(1), Dim::Exactly(max_token_context)],
),
];
let mut outputs = vec![
FeatureContract::new(
names::LOGITS,
DataType::F16,
vec![Dim::Exactly(1), Dim::Exactly(1), Dim::Exactly(vocab)],
),
FeatureContract::new(
names::KEY_UPDATES,
DataType::F16,
vec![
Dim::Exactly(1),
Dim::Exactly(kv_dim),
Dim::Exactly(1),
Dim::Exactly(1),
],
),
FeatureContract::new(
names::VALUE_UPDATES,
DataType::F16,
vec![
Dim::Exactly(1),
Dim::Exactly(kv_dim),
Dim::Exactly(1),
Dim::Exactly(1),
],
),
];
if supports_alignment {
outputs.push(FeatureContract::new(
names::ALIGNMENT,
DataType::F16,
vec![Dim::Exactly(1), Dim::Exactly(n_audio_ctx)],
));
}
LoadContract::new(inputs, outputs, StateContract::None)
}