From f9c40e52dabda7b037805c0635b84ccb6d75f5a8 Mon Sep 17 00:00:00 2001 From: sigoden Date: Fri, 3 Nov 2023 06:52:57 +0800 Subject: refactor: improve code quanity (#203) - update field name of ModelInfo - rename ModelInfo to Model --- src/client/model.rs | 80 +++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 80 insertions(+) create mode 100644 src/client/model.rs (limited to 'src/client/model.rs') diff --git a/src/client/model.rs b/src/client/model.rs new file mode 100644 index 0000000..d00ad46 --- /dev/null +++ b/src/client/model.rs @@ -0,0 +1,80 @@ +use super::message::Message; + +use crate::utils::count_tokens; + +use anyhow::{bail, Result}; + +pub type TokensCountFactors = (usize, usize); // (per-messages, bias) + +#[derive(Debug, Clone)] +pub struct Model { + pub client_index: usize, + pub client_name: String, + pub llm_name: String, + pub max_tokens: Option, + pub tokens_count_factors: TokensCountFactors, +} + +impl Default for Model { + fn default() -> Self { + Model::new(0, "", "") + } +} + +impl Model { + pub fn new(client_index: usize, client_name: &str, name: &str) -> Self { + Self { + client_index, + client_name: client_name.into(), + llm_name: name.into(), + max_tokens: None, + tokens_count_factors: Default::default(), + } + } + + pub fn id(&self) -> String { + format!("{}:{}", self.client_name, self.llm_name) + } + + pub fn set_max_tokens(mut self, max_tokens: Option) -> Self { + match max_tokens { + None | Some(0) => self.max_tokens = None, + _ => self.max_tokens = max_tokens, + } + self + } + + pub fn set_tokens_count_factors(mut self, tokens_count_factors: TokensCountFactors) -> Self { + self.tokens_count_factors = tokens_count_factors; + self + } + + pub fn messages_tokens(&self, messages: &[Message]) -> usize { + messages.iter().map(|v| count_tokens(&v.content)).sum() + } + + pub fn total_tokens(&self, messages: &[Message]) -> usize { + if messages.is_empty() { + return 0; + } + let num_messages = messages.len(); + let message_tokens = self.messages_tokens(messages); + let (per_messages, _) = self.tokens_count_factors; + if messages[num_messages - 1].role.is_user() { + num_messages * per_messages + message_tokens + } else { + (num_messages - 1) * per_messages + message_tokens + } + } + + pub fn max_tokens_limit(&self, messages: &[Message]) -> Result<()> { + let (_, bias) = self.tokens_count_factors; + let total_tokens = self.total_tokens(messages) + bias; + if let Some(max_tokens) = self.max_tokens { + if total_tokens >= max_tokens { + bail!("Exceed max tokens limit") + } + } + Ok(()) + } +} -- cgit v1.2.3