use crate::{chat::ChatMessage, error::Result};
use super::{CompletionOptions, Llama};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TextGeneration {
pub text: String,
pub prompt_tokens: u32,
pub completion_tokens: u32,
pub total_tokens: u32,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ChatGeneration {
pub message: ChatMessage,
pub prompt_tokens: u32,
pub completion_tokens: u32,
pub total_tokens: u32,
}
#[derive(Debug, Clone, PartialEq)]
pub struct EmbeddingBatch {
pub vectors: Vec<Vec<f32>>,
pub prompt_tokens: u32,
pub total_tokens: u32,
}
impl Llama {
pub fn generate_text(&mut self, prompt: &str, max_tokens: usize) -> Result<TextGeneration> {
let prompt_tokens = self.model().tokenize(prompt, true, true)?.len() as u32;
let completion =
self.create_completion_with_options(prompt, CompletionOptions::new(max_tokens))?;
let completion_tokens = completion.n_tokens as u32;
Ok(TextGeneration {
text: completion.text,
prompt_tokens,
completion_tokens,
total_tokens: prompt_tokens + completion_tokens,
})
}
pub fn generate_chat(
&mut self,
messages: &[ChatMessage],
max_tokens: usize,
) -> Result<ChatGeneration> {
let prompt_tokens = messages
.iter()
.map(|message| {
self.model()
.tokenize(&message.content, true, true)
.map_or(0, |tokens| tokens.len() as u32)
})
.sum::<u32>();
let message = self.create_chat_completion(messages, max_tokens)?;
let completion_tokens = self.model().tokenize(&message.content, false, true)?.len() as u32;
Ok(ChatGeneration {
message,
prompt_tokens,
completion_tokens,
total_tokens: prompt_tokens + completion_tokens,
})
}
pub fn embed_texts(&mut self, texts: &[String], normalize: bool) -> Result<EmbeddingBatch> {
let mut prompt_tokens = 0_u32;
let mut vectors = Vec::with_capacity(texts.len());
for text in texts {
prompt_tokens += self.model().tokenize(text, true, false)?.len() as u32;
vectors.push(self.embed(text, normalize)?);
}
Ok(EmbeddingBatch {
vectors,
prompt_tokens,
total_tokens: prompt_tokens,
})
}
}