summaryrefslogtreecommitdiffstats
path: root/src/utils/mod.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-11-03 12:30:40 +0800
committerGitHub <noreply@github.com>2023-11-03 12:30:40 +0800
commitc20cd20d535dbf6194cd10d44c91eb87fc45b343 (patch)
tree6b062176d68ab6fa6962ae036b61c5396bfee08a /src/utils/mod.rs
parentb34e40e25f0b10fccc9de64b9f06b9290be88b22 (diff)
downloadaichat-c20cd20d535dbf6194cd10d44c91eb87fc45b343.tar.gz
fix: wrap and tokenize algorithm (#205)
* fix: wrap and tokenize algorithm * update tests * remove unnecessary tokenize from tiktoken
Diffstat (limited to 'src/utils/mod.rs')
-rw-r--r--src/utils/mod.rs30
1 files changed, 28 insertions, 2 deletions
diff --git a/src/utils/mod.rs b/src/utils/mod.rs
index 1deec3e..c9fa608 100644
--- a/src/utils/mod.rs
+++ b/src/utils/mod.rs
@@ -23,8 +23,23 @@ pub fn get_env_name(key: &str) -> String {
/// Split text to tokens
pub fn tokenize(text: &str) -> Vec<String> {
- let tokens = cl100k_base_singleton().lock().tokenize(text);
- tokens.into_iter().map(|(_, text)| text).collect()
+ let tokens = cl100k_base_singleton()
+ .lock()
+ .encode_with_special_tokens(text);
+ let token_bytes: Vec<Vec<u8>> = tokens
+ .into_iter()
+ .map(|v| cl100k_base_singleton().lock().decode_bytes(vec![v]))
+ .collect();
+ let mut output = vec![];
+ let mut current_bytes = vec![];
+ for bytes in token_bytes {
+ current_bytes.extend(bytes);
+ if let Ok(v) = std::str::from_utf8(&current_bytes) {
+ output.push(v.to_string());
+ current_bytes.clear();
+ }
+ }
+ output
}
/// Count how many tokens a piece of text needs to consume
@@ -60,3 +75,14 @@ pub fn init_tokio_runtime() -> anyhow::Result<tokio::runtime::Runtime> {
.build()
.with_context(|| "Failed to init tokio")
}
+
+#[cfg(test)]
+mod tests {
+ use super::*;
+
+ #[test]
+ fn test_tokenize() {
+ assert_eq!(tokenize("😊 hello world"), ["😊", " hello", " world"]);
+ assert_eq!(tokenize("δΈ–η•Œ"), ["δΈ–", "η•Œ"]);
+ }
+}