Initial commit

This commit is contained in:
Slavi Pantaleev
2024-09-12 13:44:06 +03:00
commit 946aa9d9e9
220 changed files with 26033 additions and 0 deletions

View File

@@ -0,0 +1,17 @@
#[derive(Debug, Clone, PartialEq)]
pub enum Author {
Prompt,
Assistant,
User,
}
#[derive(Debug, Clone)]
pub struct Message {
pub author: Author,
pub message_text: String,
}
#[derive(Debug)]
pub struct Conversation {
pub messages: Vec<Message>,
}

View File

@@ -0,0 +1,10 @@
mod entity;
mod tokenization;
mod utils;
#[cfg(test)]
mod tests;
pub use entity::*;
pub use tokenization::shorten_messages_list_to_context_size;
pub use utils::*;

View File

@@ -0,0 +1,84 @@
use crate::utils::status::create_error_message_text;
use crate::utils::text_to_speech::create_transcribed_message_text;
use super::*;
#[test]
fn test_messages_by_the_bot_are_identified_correctly() {
let bot_user_id = "@bot:example.com";
let matrix_message = super::super::matrix::MatrixMessage {
sender_id: bot_user_id.to_owned(),
message_type: super::super::matrix::MatrixMessageType::Text,
message_text: "Hello!".to_owned(),
};
let llm_message = convert_matrix_message_to_llm_message(&matrix_message, bot_user_id).unwrap();
assert_eq!(llm_message.author, Author::Assistant);
assert_eq!(llm_message.message_text, "Hello!");
}
#[test]
fn test_notice_messages_by_bot_with_speech_to_text_prefix_are_cleaned_up_and_considered_sent_by_user(
) {
let bot_user_id = "@bot:example.com";
let source_message_text = "Hello!";
let message_text = create_transcribed_message_text(source_message_text);
assert_ne!(source_message_text, message_text);
let matrix_message = super::super::matrix::MatrixMessage {
sender_id: bot_user_id.to_owned(),
message_type: super::super::matrix::MatrixMessageType::Notice,
message_text,
};
let llm_message = convert_matrix_message_to_llm_message(&matrix_message, bot_user_id).unwrap();
assert_eq!(llm_message.author, Author::User);
assert_eq!(llm_message.message_text, source_message_text);
}
#[test]
fn test_notice_error_messages_by_bot_are_ignored() {
let bot_user_id = "@bot:example.com";
let source_message_text = "Some error happened";
let message_text = create_error_message_text(source_message_text);
assert_ne!(source_message_text, message_text);
let matrix_message = super::super::matrix::MatrixMessage {
sender_id: bot_user_id.to_owned(),
message_type: super::super::matrix::MatrixMessageType::Notice,
message_text,
};
let llm_message = convert_matrix_message_to_llm_message(&matrix_message, bot_user_id);
assert!(llm_message.is_none());
}
#[test]
fn test_other_notice_messages_by_the_bot_are_ignored() {
// Also see `test_notice_error_messages_by_bot_are_ignored()`.
// That one passes accidentally, because we ignore all messages by the bot that are notices
// (except for speech-to-text-created transcriptions - see `test_notice_messages_by_bot_with_speech_to_text_prefix_are_cleaned_up_and_considered_sent_by_user()`).
// This test is to make sure that we don't accidentally start accepting other notice messages.
let bot_user_id = "@bot:example.com";
let message_text = "Something something";
let matrix_message = super::super::matrix::MatrixMessage {
sender_id: bot_user_id.to_owned(),
message_type: super::super::matrix::MatrixMessageType::Notice,
message_text: message_text.to_owned(),
};
let llm_message = convert_matrix_message_to_llm_message(&matrix_message, bot_user_id);
assert!(llm_message.is_none());
}

View File

@@ -0,0 +1,287 @@
use tiktoken_rs::get_bpe_from_tokenizer;
use tiktoken_rs::tokenizer;
use tiktoken_rs::CoreBPE;
use super::{Author, Message};
fn get_bpe_for_model(model: &str) -> CoreBPE {
let tokenizer = tokenizer::get_tokenizer(model)
.or_else(|| tokenizer::get_tokenizer("gpt-4"))
.unwrap();
get_bpe_from_tokenizer(tokenizer).unwrap()
}
pub fn shorten_messages_list_to_context_size(
model: &str,
prompt_message: &Option<Message>,
mut messages: Vec<Message>,
max_response_tokens: u32,
max_context_tokens: u32,
) -> Vec<Message> {
// Loading the tokenization data is an expensive process, so
// se construct the BPE instance once and then use it for all messages.
let bpe = get_bpe_for_model(model);
// We want to retain the prompt in all cases, so we always count it first.
// We also always reserve enough tokens for the maximum response we expect.
let mut current_context_length: u32 = if let Some(prompt_message) = prompt_message {
calculate_token_size_for_message(&bpe, model, prompt_message) + max_response_tokens
} else {
0
};
messages.reverse();
let mut messages_to_keep: Vec<Message> = Vec::new();
for message in messages {
let tokens_for_message = calculate_token_size_for_message(&bpe, model, &message);
if current_context_length + tokens_for_message > max_context_tokens {
break;
}
current_context_length += tokens_for_message;
messages_to_keep.push(message);
}
messages_to_keep.reverse();
messages_to_keep
}
/// Calculate the token size of a message for a given model, with a preloaded CoreBPE object.
/// Related to `calculate_token_size_for_model_message`.
fn calculate_token_size_for_message(bpe: &CoreBPE, model: &str, message: &Message) -> u32 {
let (tokens_per_message, tokens_per_name) = if model.starts_with("gpt-3.5") {
(
4, // every message follows <im_start>{role/name}\n{content}<im_end>\n
-1, // if there's a name, the role is omitted
)
} else {
(3, 1)
};
let role_length = match message.author {
Author::Assistant => bpe.encode_with_special_tokens("assistant").len() as i32,
Author::User => bpe.encode_with_special_tokens("user").len() as i32,
Author::Prompt => bpe.encode_with_special_tokens("system").len() as i32,
};
let text_length = bpe.encode_with_special_tokens(&message.message_text).len() as i32;
(text_length + role_length + tokens_per_message + tokens_per_name) as u32
}
pub mod test {
#[test]
fn message_size_counting_works() {
let model = "gpt-4";
let bpe = super::get_bpe_for_model(model);
let message = super::Message {
author: super::Author::User,
message_text: "Hello there!".to_owned(),
};
let tokens = super::calculate_token_size_for_message(&bpe, model, &message);
assert_eq!(8, tokens);
}
#[test]
fn shortening_works_with_english() {
let model = "gpt-4";
let bpe = super::get_bpe_for_model(model);
let max_response_tokens: u32 = 5;
let prompt = super::Message {
author: super::Author::Prompt,
message_text: "You are a bot!".to_owned(),
};
let prompt_length = 10;
assert_eq!(
prompt_length,
super::calculate_token_size_for_message(&bpe, model, &prompt)
);
let mut conversation_messages = Vec::new();
let first = super::Message {
author: super::Author::User,
message_text: "Hello there!".to_owned(),
};
let first_length = 8;
assert_eq!(
first_length,
super::calculate_token_size_for_message(&bpe, model, &first)
);
conversation_messages.push(first);
let second = super::Message {
author: super::Author::Assistant,
message_text: "Hello!".to_owned(),
};
let second_length = 7;
assert_eq!(
second_length,
super::calculate_token_size_for_message(&bpe, model, &second)
);
conversation_messages.push(second);
let third = super::Message {
author: super::Author::User,
message_text: "This is the 3rd message in this conversation. It shall be preserved."
.to_owned(),
};
let third_length = 21;
assert_eq!(
third_length,
super::calculate_token_size_for_message(&bpe, model, &third)
);
conversation_messages.push(third.clone());
let forth = super::Message {
author: super::Author::Assistant,
message_text: "This is yet another message that shall be preserved.".to_owned(),
};
let forth_length = 15;
assert_eq!(
forth_length,
super::calculate_token_size_for_message(&bpe, model, &forth)
);
conversation_messages.push(forth.clone());
assert_eq!(4, conversation_messages.len());
let new_conversation_messages = super::shorten_messages_list_to_context_size(
model,
&Some(prompt),
conversation_messages,
max_response_tokens,
prompt_length + max_response_tokens + forth_length + third_length,
);
assert_eq!(2, new_conversation_messages.len());
assert_eq!(
new_conversation_messages.first().unwrap().message_text,
third.message_text
);
assert_eq!(
new_conversation_messages.last().unwrap().message_text,
forth.message_text
);
}
#[test]
fn shortening_works_with_japanese() {
let model = "gpt-4";
let bpe = super::get_bpe_for_model(model);
let max_response_tokens: u32 = 5;
let prompt = super::Message {
author: super::Author::User,
message_text: "あなたはボットです。".to_owned(),
};
let prompt_length = 14;
assert_eq!(
prompt_length,
super::calculate_token_size_for_message(&bpe, model, &prompt)
);
let mut conversation_messages = Vec::new();
let first = super::Message {
author: super::Author::User,
message_text: "こんにちは!".to_owned(),
};
let first_length = 7;
assert_eq!(
first_length,
super::calculate_token_size_for_message(&bpe, model, &first)
);
conversation_messages.push(first);
let second = super::Message {
author: super::Author::Assistant,
message_text: "こんにちは。今日は元気ですか。".to_owned(),
};
let second_length = 15;
assert_eq!(
second_length,
super::calculate_token_size_for_message(&bpe, model, &second)
);
conversation_messages.push(second);
let third = super::Message {
author: super::Author::User,
message_text: "これは第3のメッセージなので、保存されます。".to_owned(),
};
let third_length = 22;
assert_eq!(
third_length,
super::calculate_token_size_for_message(&bpe, model, &third)
);
conversation_messages.push(third.clone());
let forth = super::Message {
author: super::Author::Assistant,
message_text: "これはもう一つの保存されますメッセージです。".to_owned(),
};
let forth_length = 21;
assert_eq!(
forth_length,
super::calculate_token_size_for_message(&bpe, model, &forth)
);
conversation_messages.push(forth.clone());
assert_eq!(4, conversation_messages.len());
let new_conversation_messages = super::shorten_messages_list_to_context_size(
model,
&Some(prompt),
conversation_messages,
max_response_tokens,
prompt_length + max_response_tokens + forth_length + third_length,
);
assert_eq!(2, new_conversation_messages.len());
assert_eq!(
new_conversation_messages.first().unwrap().message_text,
third.message_text
);
assert_eq!(
new_conversation_messages.last().unwrap().message_text,
forth.message_text
);
}
}

View File

@@ -0,0 +1,51 @@
use super::{Author, Message};
use crate::conversation::matrix::{MatrixMessage, MatrixMessageType};
use crate::utils::text_to_speech as text_to_speech_utils;
pub fn convert_matrix_message_to_llm_message(
matrix_message: &MatrixMessage,
bot_user_id: &str,
) -> Option<Message> {
if matrix_message.sender_id == bot_user_id {
return convert_bot_message(matrix_message);
}
convert_user_message(matrix_message)
}
fn convert_bot_message(matrix_message: &MatrixMessage) -> Option<Message> {
match matrix_message.message_type {
MatrixMessageType::Text => convert_bot_text_message(&matrix_message.message_text),
MatrixMessageType::Notice => convert_bot_notice_message(&matrix_message.message_text),
}
}
fn convert_bot_text_message(text: &str) -> Option<Message> {
Some(Message {
author: Author::Assistant,
message_text: text.to_owned(),
})
}
fn convert_bot_notice_message(text: &str) -> Option<Message> {
// Notice messages sent by the bot are usually transcriptions of previous messages sent by the user.
// Such transcriptions are prefixed with an emoji and blockquoted.
// If we find a notice that doesn't match this pattern, we skip it.
if let Some(text) = text_to_speech_utils::parse_transcribed_message_text(text) {
// This is a transcription message. We remove the prefix and consider it as a message sent by the user.
return Some(Message {
author: Author::User,
message_text: text.to_owned(),
});
}
None
}
fn convert_user_message(matrix_message: &MatrixMessage) -> Option<Message> {
Some(Message {
author: Author::User,
message_text: matrix_message.message_text.clone(),
})
}