summaryrefslogtreecommitdiffstats
path: root/src/config/role.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-03-09 21:18:28 +0800
committerGitHub <noreply@github.com>2023-03-09 21:18:28 +0800
commit553d0fe55b3e55969e6ca48554e1f63c8e48f97e (patch)
treeb8f40a3b6aacd5207207c41849fb6bbf5ccc5a2f /src/config/role.rs
parent9767c07eeeb20cb2f8386faf66cb88452987988b (diff)
downloadaichat-553d0fe55b3e55969e6ca48554e1f63c8e48f97e.tar.gz
refactor: optimize counting tokens (#53)
Diffstat (limited to 'src/config/role.rs')
-rw-r--r--src/config/role.rs25
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)
}