From 5c0383f908eaee86539103b536a95bd94d32d401 Mon Sep 17 00:00:00 2001 From: sigoden Date: Mon, 30 Oct 2023 16:32:11 +0800 Subject: fix: dry run on role or session (#181) --- src/utils/tiktoken.rs | 198 ++++++++++++++++++++++++++++++++++++++------------ 1 file changed, 153 insertions(+), 45 deletions(-) (limited to 'src/utils/tiktoken.rs') 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 { - cl100k_base_singleton() - .lock() - .encode_with_special_tokens(text) -} - -/// Convert tokens to plan text -pub fn tokens_to_text(tokens: &[usize]) -> Result { - cl100k_base_singleton().lock().decode(tokens) -} +use tokio::task; pub fn cl100k_base() -> Result { let cl100k_base = include_str!("../../assets/cl100k_base.tiktoken"); @@ -43,11 +27,11 @@ pub fn cl100k_base() -> Result { } 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> { CL100K_BASE.clone() } +pub async fn decode_async(bpe: Arc>, tokens: Vec) -> Result { + task::spawn_blocking(move || bpe.lock().decode(tokens)).await? +} + +pub async fn encode_async(bpe: Arc>, text: &str) -> Result> { + 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>, text: &str) -> Result> { + 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, usize>) -> Vec> { 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, 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) { 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::>(); - Regex::new(&parts.join("|"))? + Regex::new(&_parts.join("|"))? }; let decoder: HashMap> = @@ -431,7 +530,7 @@ impl CoreBPE { let mut sorted_token_bytes: Vec> = 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 { - self._encode_native(text, allowed_special).0 + pub fn encode(&self, text: &str, allowed_special: HashSet<&str>) -> Vec { + 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 { 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, HashSet>) { - 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 { - self._decode_native(tokens) + pub fn decode_bytes(&self, tokens: Vec) -> Vec { + self._decode_native(&tokens) } - pub fn decode(&self, tokens: &[usize]) -> Result { - match String::from_utf8(self._decode_native(tokens)) { + pub fn decode(&self, tokens: Vec) -> Result { + 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, -- cgit v1.2.3