From 553d0fe55b3e55969e6ca48554e1f63c8e48f97e Mon Sep 17 00:00:00 2001 From: sigoden Date: Thu, 9 Mar 2023 21:18:28 +0800 Subject: refactor: optimize counting tokens (#53) --- src/config/message.rs | 30 +++++++++++++++++++++++------- 1 file changed, 23 insertions(+), 7 deletions(-) (limited to 'src/config/message.rs') diff --git a/src/config/message.rs b/src/config/message.rs index 0a06a73..514a2fa 100644 --- a/src/config/message.rs +++ b/src/config/message.rs @@ -1,6 +1,6 @@ use serde::{Deserialize, Serialize}; -pub const MESSAGE_EXTRA_TOKENS: usize = 6; +use crate::utils::count_tokens; #[derive(Debug, Clone, Deserialize, Serialize)] pub struct Message { @@ -25,10 +25,26 @@ pub enum MessageRole { User, } -#[test] -fn test_serde() { - assert_eq!( - serde_json::to_string(&Message::new("Hello World")).unwrap(), - "{\"role\":\"user\",\"content\":\"Hello World\"}" - ) +pub fn num_tokens_from_messages(messages: &[Message]) -> usize { + let mut num_tokens = 0; + for message in messages.iter() { + num_tokens += 4; + num_tokens += count_tokens(&message.content); + num_tokens += 1; // role always take 1 token + } + num_tokens += 2; + num_tokens +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_serde() { + assert_eq!( + serde_json::to_string(&Message::new("Hello World")).unwrap(), + "{\"role\":\"user\",\"content\":\"Hello World\"}" + ) + } } -- cgit v1.2.3