summaryrefslogtreecommitdiffstats
path: root/src/utils/mod.rs
diff options
context:
space:
mode:
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("δΈ–η•Œ"), ["δΈ–", "η•Œ"]);
+ }
+}