summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/client/azure_openai.rs8
-rw-r--r--src/client/localai.rs8
-rw-r--r--src/client/mod.rs6
-rw-r--r--src/client/openai.rs17
-rw-r--r--src/config/mod.rs2
-rw-r--r--src/config/model_info.rs21
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<String>,
@@ -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<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")