diff options
| author | sigoden <sigoden@gmail.com> | 2023-11-01 22:15:55 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2023-11-01 22:15:55 +0800 |
| commit | f6da06dad9b2a76016209a7d58ad923c2c72f150 (patch) | |
| tree | 1c74f306c194d3289fe06ed31a39a00c27ad0ca9 /src/config/model_info.rs | |
| parent | da3c541b681b140feb52b4297167ca233535fa90 (diff) | |
| download | aichat-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.rs | 64 |
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(()) + } } |
