diff options
| author | sigoden <sigoden@gmail.com> | 2023-11-03 12:30:40 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2023-11-03 12:30:40 +0800 |
| commit | c20cd20d535dbf6194cd10d44c91eb87fc45b343 (patch) | |
| tree | 6b062176d68ab6fa6962ae036b61c5396bfee08a /src | |
| parent | b34e40e25f0b10fccc9de64b9f06b9290be88b22 (diff) | |
| download | aichat-c20cd20d535dbf6194cd10d44c91eb87fc45b343.tar.gz | |
fix: wrap and tokenize algorithm (#205)
* fix: wrap and tokenize algorithm
* update tests
* remove unnecessary tokenize from tiktoken
Diffstat (limited to 'src')
| -rw-r--r-- | src/render/markdown.rs | 20 | ||||
| -rw-r--r-- | src/utils/mod.rs | 30 | ||||
| -rw-r--r-- | src/utils/tiktoken.rs | 111 |
3 files changed, 44 insertions, 117 deletions
diff --git a/src/render/markdown.rs b/src/render/markdown.rs index 834c47b..f491fc1 100644 --- a/src/render/markdown.rs +++ b/src/render/markdown.rs @@ -82,13 +82,13 @@ impl MarkdownRender { .join("\n") } - pub fn render_with_indent(&mut self, text: &str, padding: usize) -> String { - let text = format!("{}{}", " ".repeat(padding), text); + pub fn render_with_indent(&mut self, text: &str, indent: usize) -> String { + let text = format!("{}{}", " ".repeat(indent), text); let output = self.render(&text); if output.starts_with('\n') { output } else { - output.chars().skip(padding).collect() + output.chars().skip(indent).collect() } } @@ -186,7 +186,7 @@ impl MarkdownRender { if is_code && !self.options.wrap_code { return line; } - textwrap::wrap(&line, width as usize).join("\n") + wrap(&line, width as usize) } else { line } @@ -203,6 +203,12 @@ impl MarkdownRender { } } +fn wrap(text: &str, width: usize) -> String { + let indent: usize = text.chars().take_while(|c| *c == ' ').count(); + let wrap_options = textwrap::Options::new(width).initial_indent(&text[0..indent]); + textwrap::wrap(&text[indent..], wrap_options).join("\n") +} + #[derive(Debug, Clone, Default)] pub struct RenderOptions { pub theme: Option<Theme>, @@ -381,5 +387,11 @@ std::error::Error>> { let expect = "To unzip a file in Rust, you can use the\n`zip` crate. Here's an example code"; assert_eq!(output, expect); + + let input = "Unzip a file"; + let output = render.render_with_indent(input, 76); + let expect = "\nUnzip a file"; + + assert_eq!(output, expect); } } 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(¤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<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("δΈη"), ["δΈ", "η"]); + } +} 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<Mutex<CoreBPE>>, text: &str) -> Result<Vec<us 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(); @@ -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<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(); @@ -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<usize> { let allowed_special = self .special_tokens_encoder |
