diff options
| author | sigoden <sigoden@gmail.com> | 2023-03-09 21:18:28 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2023-03-09 21:18:28 +0800 |
| commit | 553d0fe55b3e55969e6ca48554e1f63c8e48f97e (patch) | |
| tree | b8f40a3b6aacd5207207c41849fb6bbf5ccc5a2f /src/config/role.rs | |
| parent | 9767c07eeeb20cb2f8386faf66cb88452987988b (diff) | |
| download | aichat-553d0fe55b3e55969e6ca48554e1f63c8e48f97e.tar.gz | |
refactor: optimize counting tokens (#53)
Diffstat (limited to 'src/config/role.rs')
| -rw-r--r-- | src/config/role.rs | 25 |
1 files changed, 3 insertions, 22 deletions
diff --git a/src/config/role.rs b/src/config/role.rs index d0ed623..5155ea1 100644 --- a/src/config/role.rs +++ b/src/config/role.rs @@ -1,12 +1,9 @@ -use super::message::{Message, MessageRole, MESSAGE_EXTRA_TOKENS}; - -use crate::utils::count_tokens; +use super::message::{Message, MessageRole}; use serde::{Deserialize, Serialize}; const TEMP_NAME: &str = "P"; const INPUT_PLACEHOLDER: &str = "__INPUT__"; -const INPUT_PLACEHOLDER_TOKENS: usize = 3; #[derive(Debug, Clone, Deserialize, Serialize)] pub struct Role { @@ -19,37 +16,21 @@ pub struct Role { pub prompt: String, /// What sampling temperature to use, between 0 and 2 pub temperature: Option<f64>, - /// Number of tokens - /// - /// System prompt consume extra 6 tokens - #[serde(skip_deserializing)] - pub tokens: usize, } impl Role { pub fn new(prompt: &str, temperature: Option<f64>) -> Self { - let mut value = Self { + Self { name: TEMP_NAME.into(), prompt: prompt.into(), temperature, - tokens: 0, - }; - value.tokens = value.consume_tokens(); - value + } } pub fn is_temp(&self) -> bool { self.name == TEMP_NAME } - pub fn consume_tokens(&self) -> usize { - if self.embeded() { - count_tokens(&self.prompt) + MESSAGE_EXTRA_TOKENS - INPUT_PLACEHOLDER_TOKENS - } else { - count_tokens(&self.prompt) + 2 * MESSAGE_EXTRA_TOKENS - } - } - pub fn embeded(&self) -> bool { self.prompt.contains(INPUT_PLACEHOLDER) } |
