summaryrefslogtreecommitdiffstats
path: root/src/config/model_info.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-11-01 22:15:55 +0800
committerGitHub <noreply@github.com>2023-11-01 22:15:55 +0800
commitf6da06dad9b2a76016209a7d58ad923c2c72f150 (patch)
tree1c74f306c194d3289fe06ed31a39a00c27ad0ca9 /src/config/model_info.rs
parentda3c541b681b140feb52b4297167ca233535fa90 (diff)
downloadaichat-f6da06dad9b2a76016209a7d58ad923c2c72f150.tar.gz
refactor: improve code quanity (#194)
- extends ModelInfo for tokens calculating - refactor config/session.rs, improve export, render, getter/setter - modify main.rs, allow --model override session.model
Diffstat (limited to 'src/config/model_info.rs')
-rw-r--r--src/config/model_info.rs64
1 files changed, 58 insertions, 6 deletions
diff --git a/src/config/model_info.rs b/src/config/model_info.rs
index 1f8d6f0..fa51b91 100644
--- a/src/config/model_info.rs
+++ b/src/config/model_info.rs
@@ -1,27 +1,79 @@
+use super::Message;
+
+use crate::utils::count_tokens;
+
+use anyhow::{bail, Result};
+
#[derive(Debug, Clone)]
pub struct ModelInfo {
pub client: String,
pub name: String,
- pub max_tokens: Option<usize>,
pub index: usize,
+ pub max_tokens: Option<usize>,
+ pub per_message_tokens: usize,
+ pub bias_tokens: usize,
}
impl Default for ModelInfo {
fn default() -> Self {
- ModelInfo::new("", "", None, 0)
+ ModelInfo::new(0, "", "")
}
}
impl ModelInfo {
- pub fn new(client: &str, name: &str, max_tokens: Option<usize>, index: usize) -> Self {
+ pub fn new(index: usize, client: &str, name: &str) -> Self {
Self {
+ index,
client: client.into(),
name: name.into(),
- max_tokens,
- index,
+ max_tokens: None,
+ per_message_tokens: 0,
+ bias_tokens: 0,
}
}
- pub fn stringify(&self) -> String {
+
+ pub fn set_max_tokens(mut self, max_tokens: Option<usize>) -> Self {
+ match max_tokens {
+ None | Some(0) => self.max_tokens = None,
+ _ => self.max_tokens = max_tokens,
+ }
+ 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;
+ self
+ }
+
+ pub fn full_name(&self) -> String {
format!("{}:{}", self.client, self.name)
}
+
+ pub fn messages_tokens(&self, messages: &[Message]) -> usize {
+ messages.iter().map(|v| count_tokens(&v.content)).sum()
+ }
+
+ pub fn totatl_tokens(&self, messages: &[Message]) -> usize {
+ if messages.is_empty() {
+ return 0;
+ }
+ let num_messages = messages.len();
+ let message_tokens = self.messages_tokens(messages);
+ if messages[num_messages - 1].role.is_user() {
+ num_messages * self.per_message_tokens + message_tokens
+ } else {
+ (num_messages - 1) * self.per_message_tokens + message_tokens
+ }
+ }
+
+ pub fn max_tokens_limit(&self, messages: &[Message]) -> Result<()> {
+ let total_tokens = self.totatl_tokens(messages) + self.bias_tokens;
+ if let Some(max_tokens) = self.max_tokens {
+ if total_tokens >= max_tokens {
+ bail!("Exceed max tokens limit")
+ }
+ }
+ Ok(())
+ }
}