use super::gguf::Gguf; use super::{ChatTurn, ModelFamily}; use crate::settings::ReasoningMode; use std::collections::HashMap; const MAX_REASONING_PREFIX: &str = "Reasoning Effort: Absolute maximum with no shortcuts permitted.\n\ You MUST be very thorough in your thinking and comprehensively decompose the problem to resolve the root cause, rigorously stress-testing your logic against all potential paths, edge cases, and adversarial scenarios.\n\ Explicitly write out your entire deliberation process, documenting every intermediate step, considered alternative, and rejected hypothesis to ensure absolutely no assumption is left unchecked.\n\n"; pub(super) struct Tokenizer { family: ModelFamily, tokens: Vec>, token_to_id: HashMap, i32>, merge_rank: HashMap, usize>, bos: i32, eos: i32, system: i32, user: i32, assistant: i32, observation: i32, sop: i32, think_start: i32, think_end: i32, } impl Tokenizer { pub(super) fn load(model: &Gguf, family: ModelFamily) -> Result { let tokens = model.strings("tokenizer.ggml.tokens")?.to_vec(); let merges = model.strings("tokenizer.ggml.merges")?; let token_to_id = tokens .iter() .enumerate() .map(|(id, token)| (token.clone(), id as i32)) .collect::>(); let merge_rank = merges .iter() .enumerate() .map(|(rank, merge)| (merge.clone(), rank)) .collect::>(); let lookup = |token: &[u8]| token_to_id.get(token).copied().unwrap_or(-1); let required = |token: &[u8]| { token_to_id.get(token).copied().ok_or_else(|| { format!( "required tokenizer token is missing: {}", String::from_utf8_lossy(token) ) }) }; let (bos, eos, system, user, assistant, observation, sop, think_start, think_end) = match family { ModelFamily::DeepSeek => ( required("<|begin▁of▁sentence|>".as_bytes())?, required("<|end▁of▁sentence|>".as_bytes())?, -1, required("<|User|>".as_bytes())?, required("<|Assistant|>".as_bytes())?, -1, -1, required(b"")?, required(b"")?, ), ModelFamily::Glm => ( model .token_id("tokenizer.ggml.bos_token_id") .ok() .unwrap_or_else(|| lookup(b"")), model .token_id("tokenizer.ggml.eos_token_id") .ok() .unwrap_or_else(|| lookup(b"<|endoftext|>")), lookup(b"<|system|>"), lookup(b"<|user|>"), lookup(b"<|assistant|>"), lookup(b"<|observation|>"), lookup(b""), lookup(b""), lookup(b""), ), }; if [bos, user, assistant, think_start, think_end] .into_iter() .any(|token| token < 0) || (family == ModelFamily::Glm && system < 0) { return Err("tokenizer does not provide the required DS4 chat markers".into()); } Ok(Self { family, tokens, token_to_id, merge_rank, bos, eos, system, user, assistant, observation, sop, think_start, think_end, }) } pub(super) fn vocab_size(&self) -> usize { self.tokens.len() } pub(super) fn token_bytes(&self, token: i32) -> Option> { let token = self.tokens.get(usize::try_from(token).ok()?)?; if token.windows(3).any(|window| window == [0xef, 0xbd, 0x9c]) { return Some(token.clone()); } let text = std::str::from_utf8(token).ok()?; Some(text.chars().filter_map(gpt2_codepoint_to_byte).collect()) } pub(super) fn tokenize(&self, text: &str) -> Vec { let mut output = Vec::new(); if self.family == ModelFamily::Glm { self.tokenize_glm(text, &mut output); } else { self.tokenize_joyai(text, &mut output); } output } pub(super) fn encode_chat( &self, system_prompt: &str, prompt: &str, reasoning: ReasoningMode, ) -> Vec { self.encode_conversation( system_prompt, &[ChatTurn { user: true, reasoning: None, reasoning_complete: true, content: prompt.to_owned(), }], reasoning, ) } pub(super) fn encode_conversation( &self, system_prompt: &str, messages: &[ChatTurn], reasoning: ReasoningMode, ) -> Vec { let mut output = vec![self.bos]; if self.family == ModelFamily::Glm && self.sop >= 0 { output.push(self.sop); } match (self.family, reasoning) { (ModelFamily::Glm, ReasoningMode::High | ReasoningMode::Max) => { output.push(self.system); output.extend(self.tokenize(if reasoning == ReasoningMode::Max { "Reasoning Effort: Max" } else { "Reasoning Effort: High" })); } (ModelFamily::DeepSeek, ReasoningMode::Max) => { output.extend(self.tokenize(MAX_REASONING_PREFIX)); } _ => {} } if !system_prompt.is_empty() { if self.family == ModelFamily::Glm { output.push(self.system); } output.extend(self.tokenize(system_prompt)); } for message in messages { output.push(if message.user { self.user } else { self.assistant }); if !message.user { if let Some(reasoning) = &message.reasoning { output.push(self.think_start); output.extend(self.tokenize(reasoning)); if message.reasoning_complete { output.push(self.think_end); } } else if self.family == ModelFamily::Glm { output.extend([self.think_start, self.think_end]); } else { output.push(self.think_end); } } output.extend(self.tokenize(&message.content)); if !message.user && self.family == ModelFamily::DeepSeek { output.push(self.eos); } } output.push(self.assistant); if reasoning != ReasoningMode::Direct { output.push(self.think_start); } else if self.family == ModelFamily::Glm { output.extend([self.think_start, self.think_end]); } else { output.push(self.think_end); } output } pub(super) fn eos(&self) -> i32 { self.eos } pub(super) fn is_stop(&self, token: i32) -> bool { token == self.eos || (self.family == ModelFamily::Glm && [self.system, self.user, self.assistant, self.observation].contains(&token)) } pub(super) fn is_think_start(&self, token: i32) -> bool { token == self.think_start } pub(super) fn is_think_end(&self, token: i32) -> bool { token == self.think_end } fn emit_piece(&self, raw: &[u8], output: &mut Vec) { if raw.is_empty() { return; } let encoded = byte_encode(raw); let mut symbols = encoded .char_indices() .map(|(start, character)| { encoded.as_bytes()[start..start + character.len_utf8()].to_vec() }) .collect::>(); loop { let best = symbols .windows(2) .enumerate() .filter_map(|(index, pair)| { let mut key = Vec::with_capacity(pair[0].len() + pair[1].len() + 1); key.extend(&pair[0]); key.push(b' '); key.extend(&pair[1]); self.merge_rank.get(&key).map(|rank| (index, *rank)) }) .min_by_key(|(_, rank)| *rank); let Some((index, _)) = best else { break }; let right = symbols.remove(index + 1); symbols[index].extend(right); } for symbol in symbols { if let Some(token) = self.token_to_id.get(symbol.as_slice()) { output.push(*token); } else { output.extend( symbol .iter() .filter_map(|byte| self.token_to_id.get(std::slice::from_ref(byte))) .copied(), ); } } } fn tokenize_joyai(&self, text: &str, output: &mut Vec) { let bytes = text.as_bytes(); let mut position = 0; while position < bytes.len() { let start = position; let byte = bytes[position]; if byte.is_ascii_digit() { let mut digits = 0; while position < bytes.len() && bytes[position].is_ascii_digit() && digits < 3 { position += 1; digits += 1; } } else if cjk_at(text, position) { while position < bytes.len() && cjk_at(text, position) { position = next_char(text, position); } } else if ascii_punct(byte) && position + 1 < bytes.len() && bytes[position + 1].is_ascii_alphabetic() { position += 1; while position < bytes.len() && bytes[position].is_ascii_alphabetic() { position += 1; } } else if letter_like(text, position) { position = consume_letters(text, position); } else if !ascii_newline(byte) && !ascii_punct(byte) && position + 1 < bytes.len() && letter_like(text, position + 1) { position = consume_letters(text, position + 1); } else if byte == b' ' && position + 1 < bytes.len() && ascii_punct(bytes[position + 1]) { position += 1; while position < bytes.len() && ascii_punct(bytes[position]) { position += 1; } while position < bytes.len() && ascii_newline(bytes[position]) { position += 1; } } else if ascii_punct(byte) { while position < bytes.len() && ascii_punct(bytes[position]) { position += 1; } while position < bytes.len() && ascii_newline(bytes[position]) { position += 1; } } else if byte.is_ascii_whitespace() { let mut scan = position; let mut last_newline = None; while scan < bytes.len() && bytes[scan].is_ascii_whitespace() { scan += 1; if ascii_newline(bytes[scan - 1]) { last_newline = Some(scan); } } position = last_newline.unwrap_or_else(|| { if scan < bytes.len() && scan > position + 1 && (letter_like(text, scan) || ascii_punct(bytes[scan])) { scan - 1 } else { scan } }); } else { position = next_char(text, position); } if position == start { position = next_char(text, position); } self.emit_piece(&bytes[start..position], output); } } fn tokenize_glm(&self, text: &str, output: &mut Vec) { let mut position = 0; while position < text.len() { let start = position; let current = char_info(text, position); if current.ch == '\'' { let next = char_info(text, current.next); let lower = next.ch.to_ascii_lowercase(); if matches!(lower, 's' | 't' | 'm' | 'd') { position = next.next; self.emit_piece(&text.as_bytes()[start..position], output); continue; } if next.next < text.len() { let next2 = char_info(text, next.next); if matches!( (lower, next2.ch.to_ascii_lowercase()), ('r', 'e') | ('v', 'e') | ('l', 'l') ) { position = next2.next; self.emit_piece(&text.as_bytes()[start..position], output); continue; } } } let following = char_info(text, current.next); if !matches!(current.ch, '\r' | '\n') && !current.number && (current.letter || following.letter) { position = current.next; while position < text.len() && char_info(text, position).letter { position = char_info(text, position).next; } } else if current.number { let mut digits = 0; while position < text.len() && char_info(text, position).number && digits < 3 { position = char_info(text, position).next; digits += 1; } } else { let (punctuation, punct_start) = if current.ch == ' ' { (char_info(text, current.next), current.next) } else { (current, position) }; if punctuation.punctuation { position = punct_start; while position < text.len() && char_info(text, position).punctuation { position = char_info(text, position).next; } while position < text.len() && matches!(char_info(text, position).ch, '\r' | '\n') { position = char_info(text, position).next; } } else if current.whitespace { let mut scan = position; let mut last_newline = None; let mut last_start = position; let mut count = 0; while scan < text.len() && char_info(text, scan).whitespace { let item = char_info(text, scan); last_start = scan; if matches!(item.ch, '\r' | '\n') { last_newline = Some(item.next); } scan = item.next; count += 1; } position = last_newline.unwrap_or(if count > 1 && scan < text.len() { last_start } else { scan }); } else { position = current.next; } } if position == start { position = next_char(text, position); } self.emit_piece(&text.as_bytes()[start..position], output); } } } fn byte_encode(bytes: &[u8]) -> String { bytes .iter() .map(|byte| char::from_u32(gpt2_byte_to_codepoint(*byte)).unwrap()) .collect() } fn gpt2_byte_to_codepoint(byte: u8) -> u32 { if (33..=126).contains(&byte) || (161..=172).contains(&byte) || byte >= 174 { return u32::from(byte); } let mut index = 0; for candidate in 0..=255_u16 { let candidate = candidate as u8; if (33..=126).contains(&candidate) || (161..=172).contains(&candidate) || candidate >= 174 { continue; } if candidate == byte { return 256 + index; } index += 1; } u32::from(byte) } fn gpt2_codepoint_to_byte(character: char) -> Option { let codepoint = character as u32; if (33..=126).contains(&codepoint) || (161..=172).contains(&codepoint) || (174..=255).contains(&codepoint) { return Some(codepoint as u8); } let mut index = 0; for byte in 0..=255_u16 { let byte = byte as u8; if (33..=126).contains(&byte) || (161..=172).contains(&byte) || byte >= 174 { continue; } if codepoint == 256 + index { return Some(byte); } index += 1; } None } fn next_char(text: &str, position: usize) -> usize { if position >= text.len() { return text.len(); } position + text[position..].chars().next().unwrap().len_utf8() } fn cjk_at(text: &str, position: usize) -> bool { let character = text[position..].chars().next().unwrap() as u32; (0x4e00..=0x9fa5).contains(&character) || (0x3040..=0x309f).contains(&character) || (0x30a0..=0x30ff).contains(&character) } fn letter_like(text: &str, position: usize) -> bool { let byte = text.as_bytes()[position]; byte >= 128 || byte.is_ascii_alphabetic() } fn consume_letters(text: &str, mut position: usize) -> usize { while position < text.len() && letter_like(text, position) { position = next_char(text, position); } position } fn ascii_punct(byte: u8) -> bool { byte.is_ascii_punctuation() } fn ascii_newline(byte: u8) -> bool { matches!(byte, b'\n' | b'\r') } #[derive(Clone, Copy)] struct CharInfo { ch: char, next: usize, letter: bool, number: bool, whitespace: bool, punctuation: bool, } fn char_info(text: &str, position: usize) -> CharInfo { if position >= text.len() { return CharInfo { ch: '\0', next: text.len(), letter: false, number: false, whitespace: false, punctuation: false, }; } let ch = text[position..].chars().next().unwrap(); let codepoint = ch as u32; let whitespace = unicode_whitespace(codepoint); let number = unicode_number(codepoint); let punctuation = unicode_punctuation(codepoint); CharInfo { ch, next: position + ch.len_utf8(), letter: if codepoint < 128 { ch.is_ascii_alphabetic() } else { !whitespace && !number && !punctuation }, number, whitespace, punctuation, } } fn unicode_whitespace(cp: u32) -> bool { (cp < 128 && (cp as u8).is_ascii_whitespace()) || matches!( cp, 0x0085 | 0x00a0 | 0x1680 | 0x2028 | 0x2029 | 0x202f | 0x205f | 0x3000 ) || (0x2000..=0x200a).contains(&cp) } fn unicode_number(cp: u32) -> bool { (cp < 128 && (cp as u8).is_ascii_digit()) || [ 0x0660..=0x0669, 0x06f0..=0x06f9, 0x07c0..=0x07c9, 0x0966..=0x096f, 0x09e6..=0x09ef, 0x0a66..=0x0a6f, 0x0ae6..=0x0aef, 0x0b66..=0x0b6f, 0x0be6..=0x0bef, 0x0c66..=0x0c6f, 0x0ce6..=0x0cef, 0x0d66..=0x0d6f, 0x0de6..=0x0def, 0x0e50..=0x0e59, 0x0ed0..=0x0ed9, 0x0f20..=0x0f29, 0x1040..=0x1049, 0x1090..=0x1099, 0x17e0..=0x17e9, 0x1810..=0x1819, 0xff10..=0xff19, ] .iter() .any(|range| range.contains(&cp)) } fn unicode_punctuation(cp: u32) -> bool { if cp < 128 { return (cp as u8).is_ascii_punctuation(); } matches!( cp, 0x00b4 | 0x00bb | 0x00bf | 0x00d7 | 0x00f7 | 0x0387 | 0x05c3 | 0x061b | 0x061e | 0x061f | 0x066a | 0x066d | 0x06d4 ) || [ 0x00a1..=0x00a9, 0x00ab..=0x00ac, 0x00ae..=0x00b1, 0x00b6..=0x00b8, 0x02c2..=0x02df, 0x02e5..=0x02eb, 0x02ed..=0x02ff, 0x0375..=0x037e, 0x0384..=0x0385, 0x055a..=0x055f, 0x0589..=0x058a, 0x05be..=0x05c0, 0x05c6..=0x05c7, 0x0609..=0x060a, 0x060c..=0x060d, 0x2000..=0x206f, 0x20a0..=0x20cf, 0x2100..=0x214f, 0x2190..=0x23ff, 0x2460..=0x24ff, 0x2500..=0x2775, 0x2794..=0x2bff, 0x2e00..=0x2e7f, 0x3000..=0x303f, 0xfd3e..=0xfd3f, 0xfe10..=0xfe6f, 0xff01..=0xff0f, 0xff1a..=0xff20, 0xff3b..=0xff40, 0xff5b..=0xff65, 0x1f000..=0x1faff, ] .iter() .any(|range| range.contains(&cp)) }