use std::sync::{Arc, Mutex};
use ferrox_models::PrefixCache;
use crate::generate::{FinishReason, GenerationParams};
use crate::{
budget, decode_error_response, generate, join_error_response, serving, ActiveModel, ApiError,
AppState, Model,
};
pub(crate) struct DecodeHandles {
model: Arc<Model>,
kv_pool: Option<generate::KvPoolConfig>,
paged_kv: Option<generate::PagedKvConfig>,
prefix_cache: Option<Arc<Mutex<PrefixCache>>>,
batcher: Option<serving::batch::ContinuousBatcher>,
ceiling: Option<Arc<budget::ContextCeiling>>,
}
impl DecodeHandles {
pub(crate) fn take(state: &AppState, active: &ActiveModel) -> Result<Self, ApiError> {
Ok(DecodeHandles {
model: Arc::clone(active.generative()?),
kv_pool: state.kv_pool.clone(),
paged_kv: state.paged_kv.clone(),
prefix_cache: state.prefix_cache.clone(),
batcher: active.batcher.clone(),
ceiling: active.ceiling.clone(),
})
}
pub(crate) fn model(&self) -> &Model {
&self.model
}
pub(crate) fn has_prefix_cache(&self) -> bool {
self.prefix_cache.is_some()
}
pub(crate) fn run(
&self,
prompt: &str,
params: &GenerationParams,
) -> Result<(Vec<String>, FinishReason, generate::Usage), generate::DecodeError> {
crate::run_generation(
&self.model,
prompt,
params,
self.kv_pool.as_ref(),
self.paged_kv.as_ref(),
self.prefix_cache.as_deref(),
self.batcher.as_ref(),
self.ceiling.as_deref(),
)
}
pub(crate) fn run_emit(
&self,
prompt: &str,
params: &GenerationParams,
emit: impl FnMut(&str),
) -> Result<(FinishReason, generate::Usage, String), generate::DecodeError> {
crate::run_generation_emit(
&self.model,
prompt,
params,
self.kv_pool.as_ref(),
self.paged_kv.as_ref(),
self.prefix_cache.as_deref(),
self.batcher.as_ref(),
self.ceiling.as_deref(),
emit,
)
}
}
pub(crate) async fn buffered(
handles: DecodeHandles,
prompt: String,
params: GenerationParams,
) -> Result<(Vec<String>, FinishReason, generate::Usage), ApiError> {
tokio::task::spawn_blocking(move || handles.run(&prompt, ¶ms))
.await
.map_err(join_error_response)?
.map_err(decode_error_response)
}