summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2023-11-03 12:30:40 +0800
committerGitHub <noreply@github.com>2023-11-03 12:30:40 +0800
commitc20cd20d535dbf6194cd10d44c91eb87fc45b343 (patch)
tree6b062176d68ab6fa6962ae036b61c5396bfee08a /src
parentb34e40e25f0b10fccc9de64b9f06b9290be88b22 (diff)
downloadaichat-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.rs20
-rw-r--r--src/utils/mod.rs30
-rw-r--r--src/utils/tiktoken.rs111
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(&current_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