diff options
| author | sigoden <sigoden@gmail.com> | 2023-10-30 16:32:11 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2023-10-30 16:32:11 +0800 |
| commit | 5c0383f908eaee86539103b536a95bd94d32d401 (patch) | |
| tree | c7a4b26c31a6b7b59a4f017a1db1731bbfab5ea2 /src/utils/tiktoken.rs | |
| parent | 2168610dbda294420a88a0bee4e56157c8fa3407 (diff) | |
| download | aichat-5c0383f908eaee86539103b536a95bd94d32d401.tar.gz | |
fix: dry run on role or session (#181)
Diffstat (limited to 'src/utils/tiktoken.rs')
| -rw-r--r-- | src/utils/tiktoken.rs | 198 |
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, |
