From c20cd20d535dbf6194cd10d44c91eb87fc45b343 Mon Sep 17 00:00:00 2001 From: sigoden Date: Fri, 3 Nov 2023 12:30:40 +0800 Subject: fix: wrap and tokenize algorithm (#205) * fix: wrap and tokenize algorithm * update tests * remove unnecessary tokenize from tiktoken --- src/utils/mod.rs | 30 +++++++++++++- src/utils/tiktoken.rs | 111 -------------------------------------------------- 2 files changed, 28 insertions(+), 113 deletions(-) (limited to 'src/utils') 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 { - 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> = 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(¤t_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 { .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("δΈ–η•Œ"), ["δΈ–", "η•Œ"]); + } +} diff --git a/src/utils/tiktoken.rs b/src/utils/tiktoken.rs index 3f01d9e..8b7e712 100644 --- a/src/utils/tiktoken.rs +++ b/src/utils/tiktoken.rs @@ -57,12 +57,6 @@ pub async fn encode_async(bpe: Arc>, text: &str) -> Result>, 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(); @@ -192,102 +186,6 @@ 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(); @@ -553,15 +451,6 @@ impl CoreBPE { 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 -- cgit v1.2.3