diff options
| author | sigoden <sigoden@gmail.com> | 2023-11-02 07:08:54 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2023-11-02 07:08:54 +0800 |
| commit | 5c7bfd92ff3e557477969be9db0638ef0d3d3659 (patch) | |
| tree | dc626111b7e7f9aa1a3c4b90aab22468d61e226f /src/config | |
| parent | 0238c8734ecadfc6172fdded0fbcf83949b198fa (diff) | |
| download | aichat-5c7bfd92ff3e557477969be9db0638ef0d3d3659.tar.gz | |
refactor: set tokens count factors (#195)
Diffstat (limited to 'src/config')
| -rw-r--r-- | src/config/mod.rs | 2 | ||||
| -rw-r--r-- | src/config/model_info.rs | 21 |
2 files changed, 12 insertions, 11 deletions
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<usize>, - 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") |
