use std::sync::Arc;
#[cfg(not(target_family = "wasm"))]
use futures::Stream;
use rig_core::completion::{
CompletionError, CompletionModel, CompletionRequest, CompletionResponse,
};
#[cfg(test)]
use rig_core::message::{Message, UserContent};
use rig_core::streaming::{RawStreamingChoice, StreamingCompletionResponse, StreamingResult};
#[cfg(test)]
use tokenizers::Tokenizer;
use crate::artifacts::{GgufModelData, ModelArtifacts, ModelData};
use crate::generation::{GenerationConfig, infer, stream_generate, validate_generation};
#[cfg(test)]
use crate::generation::{
IncrementalTextDecoder, effective_generation, effective_output_limit, max_tokens_to_usize,
next_cache_position, recent_tokens, sampling,
};
#[cfg(test)]
use crate::loader::*;
use crate::loader::{LoadedModel, load_gguf_model, load_model_with_family};
#[cfg(test)]
use crate::profile::{ArtifactFormat, LoaderBackend, definition_for};
#[cfg(test)]
use crate::profile::{BEGIN_OF_TEXT, END_HEADER, END_OF_TURN, IM_END, IM_START, START_HEADER};
use crate::profile::{ConversationProtocol, ModelArchitecture, ModelFamily, Quantization};
use crate::runtime::CancellationSignal;
#[cfg(all(test, not(target_family = "wasm")))]
use crate::runtime::TestControl;
#[cfg(not(target_family = "wasm"))]
use crate::runtime::{CancelOnDrop, acquire_concurrency};
use crate::types::*;
#[cfg(test)]
use crate::validation::*;
const DEFAULT_MAX_CONCURRENT_REQUESTS: usize = 1;
#[cfg(not(target_family = "wasm"))]
const STREAM_CHANNEL_CAPACITY: usize = 8;
#[derive(Clone)]
enum ModelState {
Ready(Arc<LoadedModel>),
UnsupportedMake,
}
#[derive(Clone)]
pub struct CandleModel {
state: ModelState,
}
pub struct CandleModelBuilder<'a> {
source: ModelSource<'a>,
family: Option<ModelFamily>,
generation: GenerationConfig,
max_concurrent_requests: usize,
}
enum ModelSource<'a> {
Owned(ModelArtifacts),
BorrowedGguf(GgufModelData<'a>),
}
pub type LlamaModel = CandleModel;
pub type LlamaModelBuilder<'a> = CandleModelBuilder<'a>;
impl CandleModel {
pub fn from_safetensors(data: ModelData) -> Result<Self, CandleError> {
Self::builder(data).build()
}
pub fn from_gguf(data: ModelData) -> Result<Self, CandleError> {
Self::builder_from_artifacts(ModelArtifacts::Gguf(data)).build()
}
pub fn from_gguf_bytes(data: GgufModelData<'_>) -> Result<Self, CandleError> {
Self::builder_from_gguf_bytes(data).build()
}
pub fn from_artifacts(artifacts: ModelArtifacts) -> Result<Self, CandleError> {
Self::builder_from_artifacts(artifacts).build()
}
pub fn builder(data: ModelData) -> CandleModelBuilder<'static> {
Self::builder_from_artifacts(ModelArtifacts::Safetensors(data))
}
pub fn builder_from_artifacts(artifacts: ModelArtifacts) -> CandleModelBuilder<'static> {
CandleModelBuilder {
source: ModelSource::Owned(artifacts),
family: None,
generation: GenerationConfig::default(),
max_concurrent_requests: DEFAULT_MAX_CONCURRENT_REQUESTS,
}
}
pub fn builder_from_gguf_bytes<'a>(data: GgufModelData<'a>) -> CandleModelBuilder<'a> {
CandleModelBuilder {
source: ModelSource::BorrowedGguf(data),
family: None,
generation: GenerationConfig::default(),
max_concurrent_requests: DEFAULT_MAX_CONCURRENT_REQUESTS,
}
}
#[cfg(not(target_family = "wasm"))]
pub async fn from_safetensors_async(data: ModelData) -> Result<Self, CandleError> {
Self::builder(data).build_async().await
}
#[cfg(not(target_family = "wasm"))]
pub async fn from_gguf_async(data: ModelData) -> Result<Self, CandleError> {
Self::builder_from_artifacts(ModelArtifacts::Gguf(data))
.build_async()
.await
}
#[cfg(not(target_family = "wasm"))]
pub async fn from_gguf_bytes_async(data: GgufModelData<'static>) -> Result<Self, CandleError> {
Self::builder_from_gguf_bytes(data).build_async().await
}
#[cfg(not(target_family = "wasm"))]
pub async fn from_artifacts_async(artifacts: ModelArtifacts) -> Result<Self, CandleError> {
Self::builder_from_artifacts(artifacts).build_async().await
}
pub fn conversation_protocol(&self) -> Option<ConversationProtocol> {
match &self.state {
ModelState::Ready(loaded) => Some(loaded.profile.definition.protocol),
ModelState::UnsupportedMake => None,
}
}
pub fn model_family(&self) -> Option<ModelFamily> {
self.conversation_protocol()
}
pub fn architecture(&self) -> Option<ModelArchitecture> {
match &self.state {
ModelState::Ready(loaded) => Some(loaded.profile.definition.architecture),
ModelState::UnsupportedMake => None,
}
}
pub fn quantization(&self) -> Option<Quantization> {
match &self.state {
ModelState::Ready(loaded) => loaded.profile.definition.quantization,
ModelState::UnsupportedMake => None,
}
}
}
impl<'a> CandleModelBuilder<'a> {
pub fn conversation_protocol(mut self, protocol: ConversationProtocol) -> Self {
self.family = Some(protocol);
self
}
pub fn model_family(mut self, family: ModelFamily) -> Self {
self.family = Some(family);
self
}
pub fn max_tokens(mut self, max_tokens: u64) -> Self {
self.generation.max_tokens = max_tokens;
self
}
pub fn temperature(mut self, temperature: f64) -> Self {
self.generation.temperature = temperature;
self
}
pub fn seed(mut self, seed: u64) -> Self {
self.generation.seed = seed;
self
}
pub fn top_k(mut self, top_k: Option<usize>) -> Self {
self.generation.top_k = top_k;
self
}
pub fn top_p(mut self, top_p: Option<f64>) -> Self {
self.generation.top_p = top_p;
self
}
pub fn repeat_penalty(mut self, repeat_penalty: f32) -> Self {
self.generation.repeat_penalty = repeat_penalty;
self
}
pub fn repeat_last_n(mut self, repeat_last_n: usize) -> Self {
self.generation.repeat_last_n = repeat_last_n;
self
}
pub fn max_concurrent_requests(mut self, max_concurrent_requests: usize) -> Self {
self.max_concurrent_requests = max_concurrent_requests;
self
}
pub fn build(self) -> Result<CandleModel, CandleError> {
validate_generation(&self.generation, None)?;
if self.max_concurrent_requests == 0 {
return Err(CandleError::InvalidConcurrencyLimit);
}
let loaded = match self.source {
ModelSource::Owned(artifacts) => load_model_with_family(
artifacts,
self.family,
self.generation,
self.max_concurrent_requests,
)?,
ModelSource::BorrowedGguf(data) => load_gguf_model(
data,
self.family,
self.generation,
self.max_concurrent_requests,
)?,
};
Ok(CandleModel {
state: ModelState::Ready(Arc::new(loaded)),
})
}
}
#[cfg(not(target_family = "wasm"))]
impl CandleModelBuilder<'static> {
pub async fn build_async(self) -> Result<CandleModel, CandleError> {
join_model_load(tokio::task::spawn_blocking(move || self.build())).await
}
}
#[cfg(not(target_family = "wasm"))]
async fn join_model_load(
task: tokio::task::JoinHandle<Result<CandleModel, CandleError>>,
) -> Result<CandleModel, CandleError> {
task.await
.map_err(|error| CandleError::BlockingTaskJoin(error.to_string()))?
}
#[cfg(test)]
fn render_prompt(request: &CompletionRequest) -> Result<String, CandleError> {
render_prompt_for(request, ModelFamily::Llama3)
}
#[cfg(test)]
fn render_prompt_for(
request: &CompletionRequest,
family: ModelFamily,
) -> Result<String, CandleError> {
crate::protocol::render_prompt(request, family)
}
#[cfg(not(target_family = "wasm"))]
type CandleStreamItem = Result<RawStreamingChoice<CandleCompletionResponse>, CompletionError>;
#[cfg(not(target_family = "wasm"))]
struct CandleReceiverStream {
receiver: tokio::sync::mpsc::Receiver<CandleStreamItem>,
cancellation: CancellationSignal,
}
#[cfg(not(target_family = "wasm"))]
impl Stream for CandleReceiverStream {
type Item = CandleStreamItem;
fn poll_next(
self: std::pin::Pin<&mut Self>,
context: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Self::Item>> {
self.get_mut().receiver.poll_recv(context)
}
}
#[cfg(not(target_family = "wasm"))]
impl Drop for CandleReceiverStream {
fn drop(&mut self) {
self.cancellation.cancel();
}
}
#[cfg(not(target_family = "wasm"))]
fn stream_infer(
loaded: &LoadedModel,
request: CompletionRequest,
cancellation: &CancellationSignal,
sender: &tokio::sync::mpsc::Sender<CandleStreamItem>,
) -> Result<(), CandleError> {
let response = stream_generate(loaded, request, cancellation, |choice| {
#[cfg(test)]
if let Some(control) = &loaded.test_control {
control.record_delivery_attempt();
}
sender
.blocking_send(Ok(choice))
.map_err(|_| CandleError::StreamingChannelClosed)
})?;
sender
.blocking_send(Ok(RawStreamingChoice::FinalResponse(response)))
.map_err(|_| CandleError::StreamingChannelClosed)
}
impl CompletionModel for CandleModel {
type Response = CandleCompletionResponse;
type StreamingResponse = CandleCompletionResponse;
type Client = ();
fn make(_: &Self::Client, _: impl Into<String>) -> Self {
Self {
state: ModelState::UnsupportedMake,
}
}
async fn completion(
&self,
request: CompletionRequest,
) -> Result<CompletionResponse<Self::Response>, CompletionError> {
let ModelState::Ready(loaded) = &self.state else {
return Err(CandleError::UnsupportedMake.into());
};
#[cfg(not(target_family = "wasm"))]
{
let cancellation = CancellationSignal::default();
let mut cancel_on_drop = CancelOnDrop::new(cancellation.clone());
let permit = acquire_concurrency(Arc::clone(&loaded.concurrency)).await?;
let loaded = Arc::clone(loaded);
let result = tokio::task::spawn_blocking(move || {
let result = loaded
.runtime
.device()
.with_context(|| infer(&loaded, request, &cancellation));
drop(permit);
result
})
.await
.map_err(|error| CandleError::BlockingTaskJoin(error.to_string()));
cancel_on_drop.disarm();
result?.map_err(CompletionError::from)
}
#[cfg(target_family = "wasm")]
{
infer(loaded, request, &CancellationSignal).map_err(CompletionError::from)
}
}
async fn stream(
&self,
request: CompletionRequest,
) -> Result<StreamingCompletionResponse<Self::StreamingResponse>, CompletionError> {
let ModelState::Ready(loaded) = &self.state else {
return Err(CandleError::UnsupportedMake.into());
};
#[cfg(not(target_family = "wasm"))]
{
let cancellation = CancellationSignal::default();
let mut cancel_on_drop = CancelOnDrop::new(cancellation.clone());
let permit = acquire_concurrency(Arc::clone(&loaded.concurrency)).await?;
let loaded = Arc::clone(loaded);
let (sender, receiver) = tokio::sync::mpsc::channel(STREAM_CHANNEL_CAPACITY);
let producer_sender = sender.clone();
let producer_cancellation = cancellation.clone();
let task = tokio::task::spawn_blocking(move || {
let result = loaded.runtime.device().with_context(|| {
stream_infer(&loaded, request, &producer_cancellation, &producer_sender)
});
if let Err(error) = result {
let _ = producer_sender.blocking_send(Err(error.into()));
}
drop(permit);
});
tokio::spawn(async move {
if let Err(error) = task.await {
let error = CandleError::BlockingTaskJoin(error.to_string());
let _ = sender.send(Err(error.into())).await;
}
});
let stream: StreamingResult<CandleCompletionResponse> =
Box::pin(CandleReceiverStream {
receiver,
cancellation,
});
cancel_on_drop.disarm();
Ok(StreamingCompletionResponse::stream(stream))
}
#[cfg(target_family = "wasm")]
{
let mut events = Vec::new();
let response = stream_generate(loaded, request, &CancellationSignal, |choice| {
events.push(Ok(choice));
Ok(())
})?;
events.push(Ok(RawStreamingChoice::FinalResponse(response)));
let stream: StreamingResult<CandleCompletionResponse> =
Box::pin(futures::stream::iter(events));
Ok(StreamingCompletionResponse::stream(stream))
}
}
}
#[cfg(test)]
#[allow(clippy::panic_in_result_fn)]
mod tests;