summaryrefslogtreecommitdiffstats
path: root/src/utils/tiktoken.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src/utils/tiktoken.rs')
-rw-r--r--src/utils/tiktoken.rs198
1 files changed, 153 insertions, 45 deletions
diff --git a/src/utils/tiktoken.rs b/src/utils/tiktoken.rs
index ce47c17..3f01d9e 100644
--- a/src/utils/tiktoken.rs
+++ b/src/utils/tiktoken.rs
@@ -1,6 +1,6 @@
//! Use tiktoken for count tokens
//!
-//! Copy from [https://github.com/dust-tt/dust/tree/main/core/src/providers/tiktoken](https://github.com/dust-tt/dust/tree/main/core/src/providers/tiktoken)
+//! Fork from https://github.com/dust-tt/dust/tree/main/core/src/providers/tiktoken
#![allow(unused)]
@@ -12,23 +12,7 @@ use parking_lot::Mutex;
use rustc_hash::FxHashMap as HashMap;
use std::collections::HashSet;
use std::sync::Arc;
-
-/// Count how many tokens a piece of text needs to consume
-pub fn count_tokens(text: &str) -> usize {
- text_to_tokens(text).len()
-}
-
-/// Convert a plain text to tokens
-pub fn text_to_tokens(text: &str) -> Vec<usize> {
- cl100k_base_singleton()
- .lock()
- .encode_with_special_tokens(text)
-}
-
-/// Convert tokens to plan text
-pub fn tokens_to_text(tokens: &[usize]) -> Result<String> {
- cl100k_base_singleton().lock().decode(tokens)
-}
+use tokio::task;
pub fn cl100k_base() -> Result<CoreBPE> {
let cl100k_base = include_str!("../../assets/cl100k_base.tiktoken");
@@ -43,11 +27,11 @@ pub fn cl100k_base() -> Result<CoreBPE> {
}
let mut special_tokens = HashMap::default();
- special_tokens.insert(String::from("<|endoftext|>"), 100_257);
- special_tokens.insert(String::from("<|fim_prefix|>"), 100_258);
- special_tokens.insert(String::from("<|fim_middle|>"), 100_259);
- special_tokens.insert(String::from("<|fim_suffix|>"), 100_260);
- special_tokens.insert(String::from("<|endofprompt|>"), 100_276);
+ special_tokens.insert(String::from("<|endoftext|>"), 100257);
+ special_tokens.insert(String::from("<|fim_prefix|>"), 100258);
+ special_tokens.insert(String::from("<|fim_middle|>"), 100259);
+ special_tokens.insert(String::from("<|fim_suffix|>"), 100260);
+ special_tokens.insert(String::from("<|endofprompt|>"), 100276);
CoreBPE::new(
encoder,
@@ -63,6 +47,22 @@ pub fn cl100k_base_singleton() -> Arc<Mutex<CoreBPE>> {
CL100K_BASE.clone()
}
+pub async fn decode_async(bpe: Arc<Mutex<CoreBPE>>, tokens: Vec<usize>) -> Result<String> {
+ task::spawn_blocking(move || bpe.lock().decode(tokens)).await?
+}
+
+pub async fn encode_async(bpe: Arc<Mutex<CoreBPE>>, text: &str) -> Result<Vec<usize>> {
+ let text = text.to_string();
+ let r = task::spawn_blocking(move || bpe.lock().encode_with_special_tokens(&text)).await?;
+ Ok(r)
+}
+
+pub async fn tokenize_async(bpe: Arc<Mutex<CoreBPE>>, text: &str) -> Result<Vec<(usize, String)>> {
+ let text = text.to_string();
+ let r = task::spawn_blocking(move || bpe.lock().tokenize(&text)).await?;
+ Ok(r)
+}
+
fn _byte_pair_merge(piece: &[u8], ranks: &HashMap<Vec<u8>, usize>) -> Vec<std::ops::Range<usize>> {
let mut parts: Vec<_> = (0..piece.len()).map(|i| i..i + 1).collect();
@@ -156,11 +156,11 @@ pub struct CoreBPE {
}
impl CoreBPE {
- const fn _get_regex(&self) -> &Regex {
+ fn _get_regex(&self) -> &Regex {
&self.regex
}
- const fn _get_special_regex(&self) -> &Regex {
+ fn _get_special_regex(&self) -> &Regex {
&self.special_regex
}
@@ -192,6 +192,102 @@ impl CoreBPE {
ret
}
+ fn _tokenize(&self, text: &str) -> Vec<(usize, String)> {
+ let regex = self._get_regex();
+ let mut results = vec![];
+
+ for mat in regex.find_iter(text) {
+ let string = mat.unwrap().as_str();
+ let piece = string.as_bytes();
+ if let Some(token) = self.encoder.get(piece) {
+ results.push((*token, string.to_string()));
+ continue;
+ }
+
+ results.extend(Self::_tokenize_byte_pair_encode(piece, &self.encoder));
+ }
+ results
+ }
+
+ // Copy of _encode_native but returns both the tokens and the associated string in a tuple
+ // As needed in tokenize function
+ fn _tokenize_with_spe_regex(
+ &self,
+ text: &str,
+ allowed_special: &HashSet<&str>,
+ ) -> Vec<(usize, String)> {
+ let special_regex = self._get_special_regex();
+ let regex = self._get_regex();
+ let mut ret = vec![];
+
+ let mut start = 0;
+ loop {
+ let mut next_special;
+ let mut start_find = start;
+ loop {
+ // Find the next allowed special token, if any
+ next_special = special_regex.find_from_pos(text, start_find).unwrap();
+ match next_special {
+ Some(m) => {
+ if allowed_special.contains(&text[m.start()..m.end()]) {
+ break;
+ }
+ start_find = m.start() + 1;
+ }
+ None => break,
+ }
+ }
+ let end = next_special.map_or(text.len(), |m| m.start());
+
+ // Okay, here we go, compare this logic to _encode_ordinary_native
+ for mat in regex.find_iter(&text[start..end]) {
+ let string = mat.unwrap().as_str();
+ let piece = string.as_bytes();
+ if let Some(token) = self.encoder.get(piece) {
+ ret.push((*token, string.to_string()));
+ continue;
+ }
+ ret.extend(Self::_tokenize_byte_pair_encode(piece, &self.encoder));
+ }
+
+ match next_special {
+ // And here we push the special token
+ Some(m) => {
+ let piece = m.as_str();
+ let token = self.special_tokens_encoder[piece];
+ ret.push((token, piece.to_string()));
+ start = m.end();
+ }
+ None => break,
+ }
+ }
+ ret
+ }
+
+ /**
+ * Implemented to match the logic in _encode_ordinary_native
+ * Used in tokenize function
+ */
+ pub fn _tokenize_byte_pair_encode(
+ piece: &[u8],
+ ranks: &HashMap<Vec<u8>, usize>,
+ ) -> Vec<(usize, String)> {
+ if piece.len() == 1 {
+ let string = String::from_utf8_lossy(piece);
+ return vec![(ranks[piece], string.to_string())];
+ }
+
+ _byte_pair_merge(piece, ranks)
+ .iter()
+ .map(|p| {
+ (
+ ranks[&piece[p.start..p.end]],
+ String::from_utf8_lossy(&piece[p.start..p.end]).to_string(),
+ )
+ })
+ .collect()
+ }
+
fn _encode_native(&self, text: &str, allowed_special: &HashSet<&str>) -> (Vec<usize>, usize) {
let special_regex = self._get_special_regex();
let regex = self._get_regex();
@@ -262,12 +358,15 @@ impl CoreBPE {
// Here is a quick and dirty fix:
{
let token_is_all_space = |token| {
- self.decoder.get(token).map_or(false, |token_bytes| {
- token_bytes
- .iter()
- .rev()
- .all(|&b| [b' ', b'\n', b'\t'].contains(&b))
- })
+ self.decoder
+ .get(token)
+ .map(|token_bytes| {
+ token_bytes
+ .iter()
+ .rev()
+ .all(|&b| [b' ', b'\n', b'\t'].contains(&b))
+ })
+ .unwrap_or(false)
};
if last_piece_token_len > 0
&& token_is_all_space(&tokens[tokens.len() - last_piece_token_len])
@@ -383,7 +482,7 @@ impl CoreBPE {
if unstable_bytes.len() > 1 {
let last_decoded = bstr::decode_last_utf8(unstable_bytes.as_slice());
if unstable_bytes.len() - last_decoded.1 > 0
- && last_decoded.0.map_or(false, char::is_whitespace)
+ && last_decoded.0.map_or(false, |c| c.is_whitespace())
{
let mut reencoded = byte_pair_encode(
&unstable_bytes[..unstable_bytes.len() - last_decoded.1],
@@ -410,11 +509,11 @@ impl CoreBPE {
let regex = Regex::new(pattern)?;
let special_regex = {
- let parts = special_tokens_encoder
+ let _parts = special_tokens_encoder
.keys()
.map(|s| fancy_regex::escape(s))
.collect::<Vec<_>>();
- Regex::new(&parts.join("|"))?
+ Regex::new(&_parts.join("|"))?
};
let decoder: HashMap<usize, Vec<u8>> =
@@ -431,7 +530,7 @@ impl CoreBPE {
let mut sorted_token_bytes: Vec<Vec<u8>> = encoder.keys().cloned().collect();
sorted_token_bytes.sort();
- Ok(Self {
+ Ok(CoreBPE {
encoder,
special_tokens_encoder,
decoder,
@@ -450,15 +549,24 @@ impl CoreBPE {
self._encode_ordinary_native(text)
}
- pub fn encode(&self, text: &str, allowed_special: &HashSet<&str>) -> Vec<usize> {
- self._encode_native(text, allowed_special).0
+ pub fn encode(&self, text: &str, allowed_special: HashSet<&str>) -> Vec<usize> {
+ self._encode_native(text, &allowed_special).0
+ }
+
+ pub fn tokenize(&self, text: &str) -> Vec<(usize, String)> {
+ let allowed_special = self
+ .special_tokens_encoder
+ .keys()
+ .map(|s| s.as_str())
+ .collect();
+ self._tokenize_with_spe_regex(text, &allowed_special)
}
pub fn encode_with_special_tokens(&self, text: &str) -> Vec<usize> {
let allowed_special = self
.special_tokens_encoder
.keys()
- .map(std::string::String::as_str)
+ .map(|s| s.as_str())
.collect();
self._encode_native(text, &allowed_special).0
}
@@ -492,9 +600,9 @@ impl CoreBPE {
fn encode_with_unstable(
&self,
text: &str,
- allowed_special: &HashSet<&str>,
+ allowed_special: HashSet<&str>,
) -> (Vec<usize>, HashSet<Vec<usize>>) {
- self._encode_unstable_native(text, allowed_special)
+ self._encode_unstable_native(text, &allowed_special)
}
#[allow(dead_code)]
@@ -522,12 +630,12 @@ impl CoreBPE {
// Decoding
// ====================
- pub fn decode_bytes(&self, tokens: &[usize]) -> Vec<u8> {
- self._decode_native(tokens)
+ pub fn decode_bytes(&self, tokens: Vec<usize>) -> Vec<u8> {
+ self._decode_native(&tokens)
}
- pub fn decode(&self, tokens: &[usize]) -> Result<String> {
- match String::from_utf8(self._decode_native(tokens)) {
+ pub fn decode(&self, tokens: Vec<usize>) -> Result<String> {
+ match String::from_utf8(self._decode_native(&tokens)) {
Ok(text) => Ok(text),
Err(e) => Err(anyhow!("Unable to decode into a valid UTF-8 string: {}", e)),
}
@@ -573,7 +681,7 @@ mod tests {
fn cl100k_base_test() {
let bpe = cl100k_base().unwrap();
let tokens = bpe.encode_with_special_tokens("This is a test with a lot of spaces");
- let decoded = bpe.decode(&tokens).unwrap();
+ let decoded = bpe.decode(tokens.clone()).unwrap();
assert_eq!(decoded, "This is a test with a lot of spaces");
assert_eq!(
tokens,