use std::{
path::PathBuf,
time::{Duration, Instant},
};
use crate::audio::whisper::{
error::{InvalidState, ModelError},
model::{LocalModelLoader, ModelLoader, ModelState, StateCallback},
options::ComputeOptions,
};
#[cfg(test)]
mod tests;
#[derive(Debug)]
pub struct LoadedModels {
mel: crate::Model,
encoder: crate::Model,
decoder: crate::Model,
}
impl LoadedModels {
pub fn new(mel: crate::Model, encoder: crate::Model, decoder: crate::Model) -> Self {
Self {
mel,
encoder,
decoder,
}
}
#[inline(always)]
pub const fn mel(&self) -> &crate::Model {
&self.mel
}
#[inline(always)]
pub const fn encoder(&self) -> &crate::Model {
&self.encoder
}
#[inline(always)]
pub const fn decoder(&self) -> &crate::Model {
&self.decoder
}
#[must_use]
pub fn into_parts(self) -> (crate::Model, crate::Model, crate::Model) {
(self.mel, self.encoder, self.decoder)
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct ModelLoadTimings {
encoder_load: Duration,
decoder_load: Duration,
encoder_specialization: Duration,
decoder_specialization: Duration,
}
impl ModelLoadTimings {
#[inline(always)]
pub const fn encoder_load(&self) -> Duration {
self.encoder_load
}
#[inline(always)]
pub const fn decoder_load(&self) -> Duration {
self.decoder_load
}
#[inline(always)]
pub const fn encoder_specialization(&self) -> Duration {
self.encoder_specialization
}
#[inline(always)]
pub const fn decoder_specialization(&self) -> Duration {
self.decoder_specialization
}
}
pub struct ModelManager {
folder: PathBuf,
compute: ComputeOptions,
state: ModelState,
callback: Option<StateCallback>,
models: Option<LoadedModels>,
timings: ModelLoadTimings,
}
impl std::fmt::Debug for ModelManager {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ModelManager")
.field("folder", &self.folder)
.field("compute", &self.compute)
.field("state", &self.state)
.field("callback", &self.callback.as_ref().map(|_| "<installed>"))
.field("models", &self.models)
.field("timings", &self.timings)
.finish()
}
}
impl ModelManager {
pub fn new(folder: impl Into<PathBuf>, compute: ComputeOptions) -> Self {
Self {
folder: folder.into(),
compute,
state: ModelState::Unloaded,
callback: None,
models: None,
timings: ModelLoadTimings::default(),
}
}
#[inline(always)]
pub const fn state(&self) -> ModelState {
self.state
}
#[inline(always)]
pub fn set_state_callback(&mut self, callback: StateCallback) -> &mut Self {
self.callback = Some(callback);
self
}
#[must_use]
#[inline(always)]
pub fn with_state_callback(mut self, callback: StateCallback) -> Self {
self.set_state_callback(callback);
self
}
pub fn prewarm(&mut self) -> Result<(), ModelError> {
match self.state {
ModelState::Prewarmed => return Ok(()),
ModelState::Loaded => {
return Err(ModelError::InvalidState(InvalidState::new(
"unloaded (local models prewarm before loading)",
self.state.as_str(),
)));
}
_ => {}
}
self.transition(ModelState::Prewarming);
match self.resolve_and_prewarm() {
Ok(()) => {
self.transition(ModelState::Prewarmed);
Ok(())
}
Err(err) => {
self.transition(ModelState::Unloaded);
Err(err)
}
}
}
fn resolve_and_prewarm(&mut self) -> Result<(), ModelError> {
let resolved = LocalModelLoader::new().resolve(&self.folder)?;
crate::Model::prewarm(resolved.mel_ref(), self.compute.mel())?;
let decoder_start = Instant::now();
crate::Model::prewarm(resolved.decoder_ref(), self.compute.decoder())?;
self.timings.decoder_specialization = decoder_start.elapsed();
let encoder_start = Instant::now();
crate::Model::prewarm(resolved.encoder_ref(), self.compute.encoder())?;
self.timings.encoder_specialization = encoder_start.elapsed();
Ok(())
}
pub fn ensure_loaded(&mut self) -> Result<&LoadedModels, ModelError> {
if self.state != ModelState::Loaded {
self.load_now()?;
}
Ok(
self
.models
.as_ref()
.expect("state Loaded is only ever set together with models, in load_now()"),
)
}
fn load_now(&mut self) -> Result<(), ModelError> {
self.transition(ModelState::Loading);
match self.resolve_and_load() {
Ok(models) => {
self.models = Some(models);
self.transition(ModelState::Loaded);
Ok(())
}
Err(err) => {
self.transition(ModelState::Unloaded);
Err(err)
}
}
}
fn resolve_and_load(&mut self) -> Result<LoadedModels, ModelError> {
let resolved = LocalModelLoader::new().resolve(&self.folder)?;
let mel = crate::Model::load(resolved.mel_ref(), self.compute.mel())?;
let decoder_start = Instant::now();
let decoder = crate::Model::load(resolved.decoder_ref(), self.compute.decoder())?;
self.timings.decoder_load = decoder_start.elapsed();
let encoder_start = Instant::now();
let encoder = crate::Model::load(resolved.encoder_ref(), self.compute.encoder())?;
self.timings.encoder_load = encoder_start.elapsed();
Ok(LoadedModels::new(mel, encoder, decoder))
}
pub fn unload(&mut self) {
if !matches!(self.state, ModelState::Loaded | ModelState::Prewarmed) {
return;
}
self.transition(ModelState::Unloading);
self.models = None;
self.transition(ModelState::Unloaded);
}
pub fn into_loaded(mut self) -> Result<(LoadedModels, ModelLoadTimings), ModelError> {
self.ensure_loaded()?;
let timings = self.timings;
let models = self
.models
.take()
.expect("ensure_loaded() above returned Ok, so models is populated");
Ok((models, timings))
}
fn transition(&mut self, new: ModelState) {
let old = self.state;
self.state = new;
if let Some(callback) = &self.callback {
callback(Some(old), new);
}
}
}