From 5c7bfd92ff3e557477969be9db0638ef0d3d3659 Mon Sep 17 00:00:00 2001 From: sigoden Date: Thu, 2 Nov 2023 07:08:54 +0800 Subject: refactor: set tokens count factors (#195) --- src/client/azure_openai.rs | 8 ++++++-- src/client/localai.rs | 8 ++++++-- src/client/mod.rs | 6 +++++- src/client/openai.rs | 17 +++++++++++------ src/config/mod.rs | 2 +- src/config/model_info.rs | 21 +++++++++++---------- 6 files changed, 40 insertions(+), 22 deletions(-) diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs index 3a48145..5603666 100644 --- a/src/client/azure_openai.rs +++ b/src/client/azure_openai.rs @@ -1,4 +1,4 @@ -use super::openai::{openai_build_body, openai_tokens_formula}; +use super::openai::{openai_build_body, OPENAI_TOKENS_COUNT_FACTORS}; use super::{AzureOpenAIClient, ExtraConfig, ModelInfo, PromptKind, PromptType, SendData}; use anyhow::{anyhow, Result}; @@ -46,7 +46,11 @@ impl AzureOpenAIClient { local_config .models .iter() - .map(|v| openai_tokens_formula(ModelInfo::new(index, client, &v.name).set_max_tokens(v.max_tokens))) + .map(|v| { + ModelInfo::new(index, client, &v.name) + .set_max_tokens(v.max_tokens) + .set_tokens_count_factors(OPENAI_TOKENS_COUNT_FACTORS) + }) .collect() } diff --git a/src/client/localai.rs b/src/client/localai.rs index 026ae27..d438388 100644 --- a/src/client/localai.rs +++ b/src/client/localai.rs @@ -1,4 +1,4 @@ -use super::openai::{openai_build_body, openai_tokens_formula}; +use super::openai::{openai_build_body, OPENAI_TOKENS_COUNT_FACTORS}; use super::{ExtraConfig, LocalAIClient, ModelInfo, PromptKind, PromptType, SendData}; use anyhow::Result; @@ -45,7 +45,11 @@ impl LocalAIClient { local_config .models .iter() - .map(|v| openai_tokens_formula(ModelInfo::new(index, client, &v.name).set_max_tokens(v.max_tokens))) + .map(|v| { + ModelInfo::new(index, client, &v.name) + .set_max_tokens(v.max_tokens) + .set_tokens_count_factors(OPENAI_TOKENS_COUNT_FACTORS) + }) .collect() } diff --git a/src/client/mod.rs b/src/client/mod.rs index dba049d..5fa6146 100644 --- a/src/client/mod.rs +++ b/src/client/mod.rs @@ -3,7 +3,11 @@ mod common; pub use common::*; -use crate::{config::ModelInfo, repl::ReplyStreamHandler, utils::PromptKind}; +use crate::{ + config::{ModelInfo, TokensCountFactors}, + repl::ReplyStreamHandler, + utils::PromptKind, +}; register_client!( (openai, "openai", OpenAI, OpenAIConfig, OpenAIClient), diff --git a/src/client/openai.rs b/src/client/openai.rs index 387baf0..a82ee9a 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -1,4 +1,7 @@ -use super::{ExtraConfig, ModelInfo, OpenAIClient, PromptKind, PromptType, SendData, ReplyStreamHandler}; +use super::{ + ExtraConfig, ModelInfo, OpenAIClient, PromptKind, PromptType, ReplyStreamHandler, SendData, + TokensCountFactors, +}; use anyhow::{anyhow, bail, Result}; use async_trait::async_trait; @@ -18,6 +21,8 @@ const MODELS: [(&str, usize); 4] = [ ("gpt-4-32k", 32768), ]; +pub const OPENAI_TOKENS_COUNT_FACTORS: TokensCountFactors = (5, 2); + #[derive(Debug, Clone, Deserialize, Default)] pub struct OpenAIConfig { pub name: Option, @@ -38,7 +43,11 @@ impl OpenAIClient { let client = Self::name(local_config); MODELS .into_iter() - .map(|(name, max_tokens)| openai_tokens_formula(ModelInfo::new(index, client, name).set_max_tokens(Some(max_tokens)))) + .map(|(name, max_tokens)| { + ModelInfo::new(index, client, name) + .set_max_tokens(Some(max_tokens)) + .set_tokens_count_factors(OPENAI_TOKENS_COUNT_FACTORS) + }) .collect() } @@ -135,7 +144,3 @@ pub fn openai_build_body(data: SendData, model: String) -> Value { } body } - -pub fn openai_tokens_formula(model: ModelInfo) -> ModelInfo { - model.set_tokens_formula(5, 2) -} diff --git a/src/config/mod.rs b/src/config/mod.rs index 958400c..ea45897 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -4,7 +4,7 @@ mod role; mod session; pub use self::message::Message; -pub use self::model_info::ModelInfo; +pub use self::model_info::{ModelInfo, TokensCountFactors}; use self::role::Role; use self::session::{Session, TEMP_SESSION_NAME}; diff --git a/src/config/model_info.rs b/src/config/model_info.rs index c747d82..793c014 100644 --- a/src/config/model_info.rs +++ b/src/config/model_info.rs @@ -4,14 +4,15 @@ use crate::utils::count_tokens; use anyhow::{bail, Result}; +pub type TokensCountFactors = (usize, usize); // (per-messages, bias) + #[derive(Debug, Clone)] pub struct ModelInfo { pub client: String, pub name: String, pub index: usize, pub max_tokens: Option, - pub per_message_tokens: usize, - pub bias_tokens: usize, + pub tokens_count_factors: TokensCountFactors, } impl Default for ModelInfo { @@ -27,8 +28,7 @@ impl ModelInfo { client: client.into(), name: name.into(), max_tokens: None, - per_message_tokens: 0, - bias_tokens: 0, + tokens_count_factors: Default::default(), } } @@ -40,9 +40,8 @@ impl ModelInfo { self } - pub fn set_tokens_formula(mut self, per_message_token: usize, bias_tokens: usize) -> Self { - self.per_message_tokens = per_message_token; - self.bias_tokens = bias_tokens; + pub fn set_tokens_count_factors(mut self, tokens_count_factors: TokensCountFactors) -> Self { + self.tokens_count_factors = tokens_count_factors; self } @@ -60,15 +59,17 @@ impl ModelInfo { } 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 * self.per_message_tokens + message_tokens + num_messages * per_messages + message_tokens } else { - (num_messages - 1) * self.per_message_tokens + message_tokens + (num_messages - 1) * per_messages + message_tokens } } pub fn max_tokens_limit(&self, messages: &[Message]) -> Result<()> { - let total_tokens = self.total_tokens(messages) + self.bias_tokens; + 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") -- cgit v1.2.3