Skip to main content

rig_experimental/providers/candle/
completion.rs

1use 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    // https://github.com/huggingface/text-generation-inference/blob/5ba53d44a18983a4de32d122f4cb46f4a17d9ef6/server/text_generation_server/models/model.py#L68
50    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, // seed RNG
126            Some(0.),  // temperature
127            None,      // top_p - Nucleus sampling probability stuff
128            1.1,       // repeat penalty
129            64,        // context size to consider for the repeat penalty
130            &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, // seed RNG
144            Some(0.),  // temperature
145            None,      // top_p - Nucleus sampling probability stuff
146            1.1,       // repeat penalty
147            64,        // context size to consider for the repeat penalty
148            &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
280// Custom deserializer for the weight_map to directly extract values into a HashSet
281fn 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}