1use std::sync::Arc;
65
66#[cfg(not(target_family = "wasm"))]
67use futures::Stream;
68use rig_core::completion::{
69 CompletionError, CompletionModel, CompletionRequest, CompletionResponse,
70};
71#[cfg(test)]
72use rig_core::message::{Message, UserContent};
73use rig_core::streaming::{RawStreamingChoice, StreamingCompletionResponse, StreamingResult};
74#[cfg(test)]
75use tokenizers::Tokenizer;
76
77use crate::artifacts::{GgufModelData, ModelArtifacts, ModelData};
78use crate::generation::{GenerationConfig, infer, stream_generate, validate_generation};
79#[cfg(test)]
80use crate::generation::{
81 IncrementalTextDecoder, effective_generation, effective_output_limit, max_tokens_to_usize,
82 next_cache_position, recent_tokens, sampling,
83};
84#[cfg(test)]
85use crate::loader::*;
86use crate::loader::{LoadedModel, load_gguf_model, load_model_with_family};
87#[cfg(test)]
88use crate::profile::{ArtifactFormat, LoaderBackend, definition_for};
89#[cfg(test)]
90use crate::profile::{BEGIN_OF_TEXT, END_HEADER, END_OF_TURN, IM_END, IM_START, START_HEADER};
91use crate::profile::{ConversationProtocol, ModelArchitecture, ModelFamily, Quantization};
92use crate::runtime::CancellationSignal;
93#[cfg(all(test, not(target_family = "wasm")))]
94use crate::runtime::TestControl;
95#[cfg(not(target_family = "wasm"))]
96use crate::runtime::{CancelOnDrop, acquire_concurrency};
97use crate::types::*;
98#[cfg(test)]
99use crate::validation::*;
100
101const DEFAULT_MAX_CONCURRENT_REQUESTS: usize = 1;
102#[cfg(not(target_family = "wasm"))]
103const STREAM_CHANNEL_CAPACITY: usize = 8;
104
105#[derive(Clone)]
106enum ModelState {
107 Ready(Arc<LoadedModel>),
108 UnsupportedMake,
109}
110
111#[derive(Clone)]
113pub struct CandleModel {
114 state: ModelState,
115}
116
117pub struct CandleModelBuilder<'a> {
119 source: ModelSource<'a>,
120 family: Option<ModelFamily>,
121 generation: GenerationConfig,
122 max_concurrent_requests: usize,
123}
124
125enum ModelSource<'a> {
126 Owned(ModelArtifacts),
127 BorrowedGguf(GgufModelData<'a>),
128}
129
130pub type LlamaModel = CandleModel;
135
136pub type LlamaModelBuilder<'a> = CandleModelBuilder<'a>;
138
139impl CandleModel {
140 pub fn from_safetensors(data: ModelData) -> Result<Self, CandleError> {
142 Self::builder(data).build()
143 }
144
145 pub fn from_gguf(data: ModelData) -> Result<Self, CandleError> {
147 Self::builder_from_artifacts(ModelArtifacts::Gguf(data)).build()
148 }
149
150 pub fn from_gguf_bytes(data: GgufModelData<'_>) -> Result<Self, CandleError> {
155 Self::builder_from_gguf_bytes(data).build()
156 }
157
158 pub fn from_artifacts(artifacts: ModelArtifacts) -> Result<Self, CandleError> {
160 Self::builder_from_artifacts(artifacts).build()
161 }
162
163 pub fn builder(data: ModelData) -> CandleModelBuilder<'static> {
165 Self::builder_from_artifacts(ModelArtifacts::Safetensors(data))
166 }
167
168 pub fn builder_from_artifacts(artifacts: ModelArtifacts) -> CandleModelBuilder<'static> {
170 CandleModelBuilder {
171 source: ModelSource::Owned(artifacts),
172 family: None,
173 generation: GenerationConfig::default(),
174 max_concurrent_requests: DEFAULT_MAX_CONCURRENT_REQUESTS,
175 }
176 }
177
178 pub fn builder_from_gguf_bytes<'a>(data: GgufModelData<'a>) -> CandleModelBuilder<'a> {
184 CandleModelBuilder {
185 source: ModelSource::BorrowedGguf(data),
186 family: None,
187 generation: GenerationConfig::default(),
188 max_concurrent_requests: DEFAULT_MAX_CONCURRENT_REQUESTS,
189 }
190 }
191
192 #[cfg(not(target_family = "wasm"))]
194 pub async fn from_safetensors_async(data: ModelData) -> Result<Self, CandleError> {
195 Self::builder(data).build_async().await
196 }
197
198 #[cfg(not(target_family = "wasm"))]
200 pub async fn from_gguf_async(data: ModelData) -> Result<Self, CandleError> {
201 Self::builder_from_artifacts(ModelArtifacts::Gguf(data))
202 .build_async()
203 .await
204 }
205
206 #[cfg(not(target_family = "wasm"))]
211 pub async fn from_gguf_bytes_async(data: GgufModelData<'static>) -> Result<Self, CandleError> {
212 Self::builder_from_gguf_bytes(data).build_async().await
213 }
214
215 #[cfg(not(target_family = "wasm"))]
217 pub async fn from_artifacts_async(artifacts: ModelArtifacts) -> Result<Self, CandleError> {
218 Self::builder_from_artifacts(artifacts).build_async().await
219 }
220
221 pub fn conversation_protocol(&self) -> Option<ConversationProtocol> {
223 match &self.state {
224 ModelState::Ready(loaded) => Some(loaded.profile.definition.protocol),
225 ModelState::UnsupportedMake => None,
226 }
227 }
228
229 pub fn model_family(&self) -> Option<ModelFamily> {
231 self.conversation_protocol()
232 }
233
234 pub fn architecture(&self) -> Option<ModelArchitecture> {
236 match &self.state {
237 ModelState::Ready(loaded) => Some(loaded.profile.definition.architecture),
238 ModelState::UnsupportedMake => None,
239 }
240 }
241
242 pub fn quantization(&self) -> Option<Quantization> {
244 match &self.state {
245 ModelState::Ready(loaded) => loaded.profile.definition.quantization,
246 ModelState::UnsupportedMake => None,
247 }
248 }
249}
250
251impl<'a> CandleModelBuilder<'a> {
252 pub fn conversation_protocol(mut self, protocol: ConversationProtocol) -> Self {
254 self.family = Some(protocol);
255 self
256 }
257
258 pub fn model_family(mut self, family: ModelFamily) -> Self {
260 self.family = Some(family);
261 self
262 }
263 pub fn max_tokens(mut self, max_tokens: u64) -> Self {
265 self.generation.max_tokens = max_tokens;
266 self
267 }
268
269 pub fn temperature(mut self, temperature: f64) -> Self {
271 self.generation.temperature = temperature;
272 self
273 }
274
275 pub fn seed(mut self, seed: u64) -> Self {
277 self.generation.seed = seed;
278 self
279 }
280
281 pub fn top_k(mut self, top_k: Option<usize>) -> Self {
283 self.generation.top_k = top_k;
284 self
285 }
286
287 pub fn top_p(mut self, top_p: Option<f64>) -> Self {
289 self.generation.top_p = top_p;
290 self
291 }
292
293 pub fn repeat_penalty(mut self, repeat_penalty: f32) -> Self {
295 self.generation.repeat_penalty = repeat_penalty;
296 self
297 }
298
299 pub fn repeat_last_n(mut self, repeat_last_n: usize) -> Self {
301 self.generation.repeat_last_n = repeat_last_n;
302 self
303 }
304
305 pub fn max_concurrent_requests(mut self, max_concurrent_requests: usize) -> Self {
310 self.max_concurrent_requests = max_concurrent_requests;
311 self
312 }
313
314 pub fn build(self) -> Result<CandleModel, CandleError> {
316 validate_generation(&self.generation, None)?;
317 if self.max_concurrent_requests == 0 {
318 return Err(CandleError::InvalidConcurrencyLimit);
319 }
320 let loaded = match self.source {
321 ModelSource::Owned(artifacts) => load_model_with_family(
322 artifacts,
323 self.family,
324 self.generation,
325 self.max_concurrent_requests,
326 )?,
327 ModelSource::BorrowedGguf(data) => load_gguf_model(
328 data,
329 self.family,
330 self.generation,
331 self.max_concurrent_requests,
332 )?,
333 };
334 Ok(CandleModel {
335 state: ModelState::Ready(Arc::new(loaded)),
336 })
337 }
338}
339
340#[cfg(not(target_family = "wasm"))]
341impl CandleModelBuilder<'static> {
342 pub async fn build_async(self) -> Result<CandleModel, CandleError> {
347 join_model_load(tokio::task::spawn_blocking(move || self.build())).await
348 }
349}
350
351#[cfg(not(target_family = "wasm"))]
352async fn join_model_load(
353 task: tokio::task::JoinHandle<Result<CandleModel, CandleError>>,
354) -> Result<CandleModel, CandleError> {
355 task.await
356 .map_err(|error| CandleError::BlockingTaskJoin(error.to_string()))?
357}
358
359#[cfg(test)]
360fn render_prompt(request: &CompletionRequest) -> Result<String, CandleError> {
361 render_prompt_for(request, ModelFamily::Llama3)
362}
363
364#[cfg(test)]
365fn render_prompt_for(
366 request: &CompletionRequest,
367 family: ModelFamily,
368) -> Result<String, CandleError> {
369 crate::protocol::render_prompt(request, family)
370}
371
372#[cfg(not(target_family = "wasm"))]
373type CandleStreamItem = Result<RawStreamingChoice<CandleCompletionResponse>, CompletionError>;
374
375#[cfg(not(target_family = "wasm"))]
376struct CandleReceiverStream {
377 receiver: tokio::sync::mpsc::Receiver<CandleStreamItem>,
378 cancellation: CancellationSignal,
379}
380
381#[cfg(not(target_family = "wasm"))]
382impl Stream for CandleReceiverStream {
383 type Item = CandleStreamItem;
384
385 fn poll_next(
386 self: std::pin::Pin<&mut Self>,
387 context: &mut std::task::Context<'_>,
388 ) -> std::task::Poll<Option<Self::Item>> {
389 self.get_mut().receiver.poll_recv(context)
390 }
391}
392
393#[cfg(not(target_family = "wasm"))]
394impl Drop for CandleReceiverStream {
395 fn drop(&mut self) {
396 self.cancellation.cancel();
397 }
398}
399
400#[cfg(not(target_family = "wasm"))]
401fn stream_infer(
402 loaded: &LoadedModel,
403 request: CompletionRequest,
404 cancellation: &CancellationSignal,
405 sender: &tokio::sync::mpsc::Sender<CandleStreamItem>,
406) -> Result<(), CandleError> {
407 let response = stream_generate(loaded, request, cancellation, |choice| {
408 #[cfg(test)]
409 if let Some(control) = &loaded.test_control {
410 control.record_delivery_attempt();
411 }
412 sender
413 .blocking_send(Ok(choice))
414 .map_err(|_| CandleError::StreamingChannelClosed)
415 })?;
416 sender
417 .blocking_send(Ok(RawStreamingChoice::FinalResponse(response)))
418 .map_err(|_| CandleError::StreamingChannelClosed)
419}
420
421impl CompletionModel for CandleModel {
422 type Response = CandleCompletionResponse;
423 type StreamingResponse = CandleCompletionResponse;
424 type Client = ();
425
426 fn make(_: &Self::Client, _: impl Into<String>) -> Self {
427 Self {
428 state: ModelState::UnsupportedMake,
429 }
430 }
431
432 async fn completion(
433 &self,
434 request: CompletionRequest,
435 ) -> Result<CompletionResponse<Self::Response>, CompletionError> {
436 let ModelState::Ready(loaded) = &self.state else {
437 return Err(CandleError::UnsupportedMake.into());
438 };
439
440 #[cfg(not(target_family = "wasm"))]
441 {
442 let cancellation = CancellationSignal::default();
443 let mut cancel_on_drop = CancelOnDrop::new(cancellation.clone());
444 let permit = acquire_concurrency(Arc::clone(&loaded.concurrency)).await?;
445 let loaded = Arc::clone(loaded);
446 let result = tokio::task::spawn_blocking(move || {
447 let result = loaded
448 .runtime
449 .device()
450 .with_context(|| infer(&loaded, request, &cancellation));
451 drop(permit);
452 result
453 })
454 .await
455 .map_err(|error| CandleError::BlockingTaskJoin(error.to_string()));
456 cancel_on_drop.disarm();
457 result?.map_err(CompletionError::from)
458 }
459
460 #[cfg(target_family = "wasm")]
461 {
462 infer(loaded, request, &CancellationSignal).map_err(CompletionError::from)
463 }
464 }
465
466 async fn stream(
467 &self,
468 request: CompletionRequest,
469 ) -> Result<StreamingCompletionResponse<Self::StreamingResponse>, CompletionError> {
470 let ModelState::Ready(loaded) = &self.state else {
471 return Err(CandleError::UnsupportedMake.into());
472 };
473
474 #[cfg(not(target_family = "wasm"))]
475 {
476 let cancellation = CancellationSignal::default();
477 let mut cancel_on_drop = CancelOnDrop::new(cancellation.clone());
478 let permit = acquire_concurrency(Arc::clone(&loaded.concurrency)).await?;
479 let loaded = Arc::clone(loaded);
480 let (sender, receiver) = tokio::sync::mpsc::channel(STREAM_CHANNEL_CAPACITY);
481 let producer_sender = sender.clone();
482 let producer_cancellation = cancellation.clone();
483 let task = tokio::task::spawn_blocking(move || {
484 let result = loaded.runtime.device().with_context(|| {
485 stream_infer(&loaded, request, &producer_cancellation, &producer_sender)
486 });
487 if let Err(error) = result {
488 let _ = producer_sender.blocking_send(Err(error.into()));
489 }
490 drop(permit);
491 });
492 tokio::spawn(async move {
493 if let Err(error) = task.await {
494 let error = CandleError::BlockingTaskJoin(error.to_string());
495 let _ = sender.send(Err(error.into())).await;
496 }
497 });
498 let stream: StreamingResult<CandleCompletionResponse> =
499 Box::pin(CandleReceiverStream {
500 receiver,
501 cancellation,
502 });
503 cancel_on_drop.disarm();
504 Ok(StreamingCompletionResponse::stream(stream))
505 }
506
507 #[cfg(target_family = "wasm")]
508 {
509 let mut events = Vec::new();
510 let response = stream_generate(loaded, request, &CancellationSignal, |choice| {
511 events.push(Ok(choice));
512 Ok(())
513 })?;
514 events.push(Ok(RawStreamingChoice::FinalResponse(response)));
515 let stream: StreamingResult<CandleCompletionResponse> =
516 Box::pin(futures::stream::iter(events));
517 Ok(StreamingCompletionResponse::stream(stream))
518 }
519 }
520}
521
522#[cfg(test)]
523#[allow(clippy::panic_in_result_fn)]
524mod tests;