summaryrefslogtreecommitdiffstats
path: root/src/client/model.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-04-11 05:03:58 +0800
committerGitHub <noreply@github.com>2024-04-11 05:03:58 +0800
commit3b5843fe2e6b4e3725bbc1b269c23115c953b2b0 (patch)
treea1d2c36dbccf0c82079f774d1efe08c993ee31f2 /src/client/model.rs
parent01ebc8734838e13fb533022c1d24e7003ed42935 (diff)
downloadaichat-3b5843fe2e6b4e3725bbc1b269c23115c953b2b0.tar.gz
refactor: all clients use openai token counter (#402)
Diffstat (limited to 'src/client/model.rs')
-rw-r--r--src/client/model.rs18
1 files changed, 5 insertions, 13 deletions
diff --git a/src/client/model.rs b/src/client/model.rs
index ce181a5..03877b9 100644
--- a/src/client/model.rs
+++ b/src/client/model.rs
@@ -5,7 +5,8 @@ use crate::utils::count_tokens;
use anyhow::{bail, Result};
use serde::{Deserialize, Deserializer};
-pub type TokensCountFactors = (usize, usize); // (per-messages, bias)
+const PER_MESSAGES_TOKENS: usize = 5;
+const BASIS_TOKENS: usize = 2;
#[derive(Debug, Clone)]
pub struct Model {
@@ -13,7 +14,6 @@ pub struct Model {
pub name: String,
pub max_input_tokens: Option<usize>,
pub extra_fields: Option<serde_json::Map<String, serde_json::Value>>,
- pub tokens_count_factors: TokensCountFactors,
pub capabilities: ModelCapabilities,
}
@@ -30,7 +30,6 @@ impl Model {
name: name.into(),
extra_fields: None,
max_input_tokens: None,
- tokens_count_factors: Default::default(),
capabilities: ModelCapabilities::Text,
}
}
@@ -91,11 +90,6 @@ impl Model {
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()
@@ -114,17 +108,15 @@ impl Model {
}
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
+ num_messages * PER_MESSAGES_TOKENS + message_tokens
} else {
- (num_messages - 1) * per_messages + message_tokens
+ (num_messages - 1) * PER_MESSAGES_TOKENS + message_tokens
}
}
pub fn max_input_tokens_limit(&self, messages: &[Message]) -> Result<()> {
- let (_, bias) = self.tokens_count_factors;
- let total_tokens = self.total_tokens(messages) + bias;
+ let total_tokens = self.total_tokens(messages) + BASIS_TOKENS;
if let Some(max_input_tokens) = self.max_input_tokens {
if total_tokens >= max_input_tokens {
bail!("Exceed max input tokens limit")