rig_experimental/providers/candle/
completion.rs1use rig::OneOrMany;
2use rig::client::{
3 AsAudioGeneration, AsEmbeddings, AsTranscription, CompletionClient, ProviderClient,
4};
5use rig::message::{AssistantContent, Message, Text, UserContent};
6use serde::Deserialize;
7use serde::de::Deserializer;
8use std::collections::HashSet;
9use std::marker::PhantomData;
10
11use anyhow::Result;
12
13use candle_transformers::models::mistral::{Config, Model as Mistral};
14
15use candle_core::{DType, Device, Tensor};
16use candle_nn::VarBuilder;
17use candle_transformers::generation::LogitsProcessor;
18use hf_hub::{Repo, RepoType, api::sync::ApiBuilder};
19use tokenizers::Tokenizer;
20
21pub struct TokenOutputStream {
22 tokenizer: tokenizers::Tokenizer,
23 tokens: Vec<u32>,
24 prev_index: usize,
25 current_index: usize,
26}
27
28impl TokenOutputStream {
29 pub fn new(tokenizer: tokenizers::Tokenizer) -> Self {
30 Self {
31 tokenizer,
32 tokens: Vec::new(),
33 prev_index: 0,
34 current_index: 0,
35 }
36 }
37
38 pub fn into_inner(self) -> tokenizers::Tokenizer {
39 self.tokenizer
40 }
41
42 fn decode(&self, tokens: &[u32]) -> candle_core::Result<String> {
43 match self.tokenizer.decode(tokens, true) {
44 Ok(str) => Ok(str),
45 Err(err) => candle_core::bail!("cannot decode: {err}"),
46 }
47 }
48
49 pub fn next_token(&mut self, token: u32) -> candle_core::Result<Option<String>> {
51 let prev_text = if self.tokens.is_empty() {
52 String::new()
53 } else {
54 let tokens = &self.tokens[self.prev_index..self.current_index];
55 self.decode(tokens)?
56 };
57 self.tokens.push(token);
58 let text = self.decode(&self.tokens[self.prev_index..])?;
59 if text.len() > prev_text.len() && text.chars().last().unwrap().is_alphanumeric() {
60 let text = text.split_at(prev_text.len());
61 self.prev_index = self.current_index;
62 self.current_index = self.tokens.len();
63 Ok(Some(text.1.to_string()))
64 } else {
65 Ok(None)
66 }
67 }
68
69 pub fn decode_rest(&self) -> Result<Option<String>> {
70 let prev_text = if self.tokens.is_empty() {
71 String::new()
72 } else {
73 let tokens = &self.tokens[self.prev_index..self.current_index];
74 self.decode(tokens)?
75 };
76 let text = self.decode(&self.tokens[self.prev_index..])?;
77 if text.len() > prev_text.len() {
78 let text = text.split_at(prev_text.len());
79 Ok(Some(text.1.to_string()))
80 } else {
81 Ok(None)
82 }
83 }
84
85 pub fn decode_all(&self) -> candle_core::Result<String> {
86 self.decode(&self.tokens)
87 }
88
89 pub fn get_token(&self, token_s: &str) -> Option<u32> {
90 self.tokenizer.get_vocab(true).get(token_s).copied()
91 }
92
93 pub fn tokenizer(&self) -> &tokenizers::Tokenizer {
94 &self.tokenizer
95 }
96
97 pub fn clear(&mut self) {
98 self.tokens.clear();
99 self.prev_index = 0;
100 self.current_index = 0;
101 }
102}
103
104impl<T> From<(T, Device, Tokenizer)> for CompletionModel<T>
105where
106 T: CandleModel + Clone + std::fmt::Debug + Sync + Send + 'static,
107{
108 fn from((model, device, tokenizer): (T, Device, Tokenizer)) -> Self {
109 Self {
110 model,
111 device,
112 tokenizer,
113 }
114 }
115}
116
117impl<T> From<CompletionModel<T>> for TextGeneration<T>
118where
119 T: CandleModel + Clone + std::fmt::Debug + Sync + Send + 'static,
120{
121 fn from(e: CompletionModel<T>) -> Self {
122 Self::new(
123 e.model,
124 e.tokenizer,
125 299792458, Some(0.), None, 1.1, 64, &e.device,
131 )
132 }
133}
134
135impl<T> From<&CompletionModel<T>> for TextGeneration<T>
136where
137 T: CandleModel + Clone + std::fmt::Debug + Sync + Send + 'static,
138{
139 fn from(e: &CompletionModel<T>) -> Self {
140 Self::new(
141 e.model.clone(),
142 e.tokenizer.clone(),
143 299792458, Some(0.), None, 1.1, 64, &e.device,
149 )
150 }
151}
152
153struct TextGeneration<T> {
154 model: T,
155 device: Device,
156 tokenizer: TokenOutputStream,
157 logits_processor: LogitsProcessor,
158 repeat_penalty: f32,
159 repeat_last_n: usize,
160}
161
162impl<T> TextGeneration<T>
163where
164 T: CandleModel + Clone + std::fmt::Debug + Sync + Send + 'static,
165{
166 #[allow(clippy::too_many_arguments)]
167 fn new(
168 model: T,
169 tokenizer: Tokenizer,
170 seed: u64,
171 _temp: Option<f64>,
172 _top_p: Option<f64>,
173 repeat_penalty: f32,
174 repeat_last_n: usize,
175 device: &Device,
176 ) -> Self {
177 let logits_processor = LogitsProcessor::new(seed, Some(0.0), None);
178
179 Self {
180 model,
181 tokenizer: TokenOutputStream::new(tokenizer),
182 logits_processor,
183 repeat_penalty,
184 repeat_last_n,
185 device: device.clone(),
186 }
187 }
188
189 fn run(mut self, prompt: String, sample_len: usize) -> CompletionResponse {
190 self.tokenizer.clear();
191 let mut tokens = self
192 .tokenizer
193 .tokenizer()
194 .encode(prompt, true)
195 .unwrap()
196 .get_ids()
197 .to_vec();
198
199 let eos_token = match self.tokenizer.get_token("</s>") {
200 Some(token) => token,
201 _ => panic!("cannot find the </s> token"),
202 };
203
204 let mut string = String::new();
205
206 let mut token_usage = 0;
207
208 for index in 0..sample_len {
209 let context_size = if index > 0 { 1 } else { tokens.len() };
210 let start_pos = tokens.len().saturating_sub(context_size);
211 let ctxt = &tokens[start_pos..];
212 let input = Tensor::new(ctxt, &self.device)
213 .unwrap()
214 .unsqueeze(0)
215 .unwrap();
216 let logits = self.model.forward(&input, start_pos).unwrap();
217 let logits = logits
218 .squeeze(0)
219 .unwrap()
220 .squeeze(0)
221 .unwrap()
222 .to_dtype(DType::F32)
223 .unwrap();
224 let logits = if self.repeat_penalty == 1. {
225 logits
226 } else {
227 let start_at = tokens.len().saturating_sub(self.repeat_last_n);
228 candle_transformers::utils::apply_repeat_penalty(
229 &logits,
230 self.repeat_penalty,
231 &tokens[start_at..],
232 )
233 .unwrap()
234 };
235
236 let next_token = self.logits_processor.sample(&logits).unwrap();
237 tokens.push(next_token);
238
239 if next_token == eos_token {
240 token_usage = index + 1;
241 break;
242 }
243
244 if let Some(t) = self.tokenizer.next_token(next_token).unwrap() {
245 println!("Found token: {t}");
246 string.push_str(&t);
247 }
248 }
249
250 CompletionResponse {
251 response: string,
252 token_usage,
253 }
254 }
255}
256
257pub struct CompletionResponse {
258 response: String,
259 pub token_usage: usize,
260}
261
262impl TryFrom<CompletionResponse> for rig::completion::CompletionResponse<CompletionResponse> {
263 type Error = rig::completion::CompletionError;
264
265 fn try_from(raw_response: CompletionResponse) -> std::result::Result<Self, Self::Error> {
266 let text = raw_response.response.clone();
267 Ok(Self {
268 choice: OneOrMany::one(AssistantContent::Text(Text { text })),
269 raw_response,
270 })
271 }
272}
273
274#[derive(Debug, Deserialize)]
275struct Weightmaps {
276 #[serde(deserialize_with = "deserialize_weight_map")]
277 weight_map: HashSet<String>,
278}
279
280fn deserialize_weight_map<'de, D>(deserializer: D) -> anyhow::Result<HashSet<String>, D::Error>
282where
283 D: Deserializer<'de>,
284{
285 let map = serde_json::Value::deserialize(deserializer)?;
286 match map {
287 serde_json::Value::Object(obj) => Ok(obj
288 .values()
289 .filter_map(|v| v.as_str().map(ToString::to_string))
290 .collect::<HashSet<String>>()),
291 _ => Err(serde::de::Error::custom(
292 "Expected an object for weight_map",
293 )),
294 }
295}
296
297pub fn hub_load_safetensors(
298 repo: &hf_hub::api::sync::ApiRepo,
299 json_file: &str,
300) -> Result<Vec<std::path::PathBuf>> {
301 let json_file = repo.get(json_file).map_err(candle_core::Error::wrap)?;
302 let json_file = std::fs::File::open(json_file)?;
303 let json: Weightmaps = serde_json::from_reader(&json_file).map_err(candle_core::Error::wrap)?;
304
305 let pathbufs: Vec<std::path::PathBuf> = json
306 .weight_map
307 .iter()
308 .map(|f| repo.get(f).unwrap())
309 .collect();
310
311 Ok(pathbufs)
312}
313
314#[derive(Debug, Clone)]
315pub struct Client<T> {
316 api_key: Option<String>,
317 model_ty: PhantomData<T>,
318}
319
320impl<T> Client<T> {
321 pub fn new(api_key: &str) -> Self {
322 Self {
323 api_key: Some(api_key.to_string()),
324 model_ty: PhantomData,
325 }
326 }
327
328 pub fn no_api_key() -> Self {
329 Self {
330 api_key: None,
331 model_ty: PhantomData,
332 }
333 }
334}
335
336impl<T> ProviderClient for Client<T>
337where
338 T: CandleModel + Clone + std::fmt::Debug + Send + Sync + 'static,
339{
340 fn from_env() -> Self {
341 let api_key = std::env::var("HUGGINGFACE_API_KEY").ok();
342
343 Self {
344 api_key,
345 model_ty: PhantomData,
346 }
347 }
348}
349
350#[derive(Clone)]
351pub struct CompletionModel<T> {
352 model: T,
353 device: Device,
354 tokenizer: Tokenizer,
355}
356
357impl<T> rig::completion::CompletionModel for CompletionModel<T>
358where
359 T: CandleModel + Clone + Send + Sync + std::fmt::Debug + 'static,
360{
361 type Response = CompletionResponse;
362 type StreamingResponse = String;
363
364 async fn completion(
365 &self,
366 request: rig::completion::CompletionRequest,
367 ) -> std::result::Result<
368 rig::completion::CompletionResponse<Self::Response>,
369 rig::completion::CompletionError,
370 > {
371 let max_tokens = if let Some(max_tokens) = request.max_tokens {
372 max_tokens as usize
373 } else {
374 1024
375 };
376 println!("Loading text generator...");
377 let text_generation = TextGeneration::from(self);
378 let prompt = convert_messages_to_mistral_compat(request.preamble, request.chat_history);
379
380 println!("Running text generator...");
381 let response = text_generation.run(prompt, max_tokens);
382
383 response.try_into()
384 }
385
386 async fn stream(
387 &self,
388 _request: rig::completion::CompletionRequest,
389 ) -> std::result::Result<
390 rig::streaming::StreamingCompletionResponse<Self::StreamingResponse>,
391 rig::completion::CompletionError,
392 > {
393 todo!()
394 }
395}
396
397impl<T> AsEmbeddings for Client<T>
398where
399 T: CandleModel + std::fmt::Debug + Clone + Send + Sync,
400{
401 fn as_embeddings(&self) -> Option<Box<dyn rig::client::embeddings::EmbeddingsClientDyn>> {
402 None
403 }
404}
405
406impl<T> AsTranscription for Client<T>
407where
408 T: CandleModel + std::fmt::Debug + Clone + Send + Sync,
409{
410 fn as_transcription(
411 &self,
412 ) -> Option<Box<dyn rig::client::transcription::TranscriptionClientDyn>> {
413 None
414 }
415}
416
417impl<T> AsAudioGeneration for Client<T>
418where
419 T: CandleModel + std::fmt::Debug + Clone + Send + Sync,
420{
421 fn as_audio_generation(
422 &self,
423 ) -> Option<Box<dyn rig::client::audio_generation::AudioGenerationClientDyn>> {
424 None
425 }
426}
427
428#[cfg(feature = "image")]
429impl<T> AsImageGeneration for Client<T>
430where
431 T: CandleModel + std::fmt::Debug + Clone + Send + Sync,
432{
433 fn as_image_generation(
434 &self,
435 ) -> Option<Box<dyn rig::client::image_generation::ImageGenerationClientDyn>> {
436 None
437 }
438}
439
440impl<T> CompletionClient for Client<T>
441where
442 T: CandleModel + std::fmt::Debug + Clone + Send + Sync + 'static,
443 T::Config: Clone + std::fmt::Debug,
444{
445 type CompletionModel = CompletionModel<T>;
446 fn completion_model(&self, model: &str) -> Self::CompletionModel {
447 let api = ApiBuilder::new()
448 .with_token(self.api_key.clone())
449 .build()
450 .expect("to successfully build the HuggingFace API client");
451
452 let repo = api.repo(Repo::with_revision(
453 model.to_string(),
454 RepoType::Model,
455 "main".to_string(),
456 ));
457
458 let tokenizer = {
459 let tokenizer_filename = repo.get("tokenizer.json").unwrap();
460 Tokenizer::from_file(tokenizer_filename).unwrap()
461 };
462
463 let device = Device::Cpu;
464 let filenames = hub_load_safetensors(&repo, "model.safetensors.index.json").unwrap();
465
466 let model = {
467 let dtype = DType::F32;
468 let vb =
469 unsafe { VarBuilder::from_mmaped_safetensors(&filenames, dtype, &device).unwrap() };
470 T::new(vb)
471 };
472
473 CompletionModel::from((model, device, tokenizer))
474 }
475}
476
477trait CandleModel {
478 type Config: Clone + std::fmt::Debug;
479 fn new(vb: VarBuilder<'_>) -> Self;
480
481 fn forward(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> candle_core::Result<Tensor>;
482}
483
484impl CandleModel for Mistral {
485 type Config = Config;
486
487 fn new(vb: VarBuilder<'_>) -> Self {
488 let config = Config::config_7b_v0_1(false);
489 Self::new(&config, vb).unwrap()
490 }
491
492 fn forward(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> candle_core::Result<Tensor> {
493 self.forward(input_ids, seqlen_offset)
494 }
495}
496
497fn convert_messages_to_mistral_compat(
498 premable: Option<String>,
499 messages: OneOrMany<Message>,
500) -> String {
501 let mut str = premable.unwrap_or_default();
502 let messages = messages
503 .into_iter()
504 .map(convert_message_to_mistral)
505 .collect::<Vec<String>>()
506 .join("\n");
507 str.push('\n');
508 str.push_str(&messages);
509 str.push_str("\n<|assistant|>");
510
511 str
512}
513
514fn convert_message_to_mistral(message: Message) -> String {
515 match message {
516 Message::User { content } => content
517 .into_iter()
518 .map(|x| match x {
519 UserContent::Text(Text { text }) => format!("<|user|>{text}"),
520 _ => unimplemented!(
521 "Only text messages are supported for local Candle models currently!"
522 ),
523 })
524 .collect::<Vec<String>>()
525 .join("\n"),
526 Message::Assistant { content } => content
527 .into_iter()
528 .map(|x| match x {
529 AssistantContent::Text(Text { text }) => format!("<|assistant|>\n{text}"),
530 _ => unimplemented!(
531 "Only text messages are supported for local Candle models currently!"
532 ),
533 })
534 .collect::<Vec<String>>()
535 .join("\n"),
536 }
537}