649 lines
21 KiB
Rust
649 lines
21 KiB
Rust
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<Vec<u8>>,
|
||
token_to_id: HashMap<Vec<u8>, i32>,
|
||
merge_rank: HashMap<Vec<u8>, 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<Self, String> {
|
||
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::<HashMap<_, _>>();
|
||
let merge_rank = merges
|
||
.iter()
|
||
.enumerate()
|
||
.map(|(rank, merge)| (merge.clone(), rank))
|
||
.collect::<HashMap<_, _>>();
|
||
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"<think>")?,
|
||
required(b"</think>")?,
|
||
),
|
||
ModelFamily::Glm => (
|
||
model
|
||
.token_id("tokenizer.ggml.bos_token_id")
|
||
.ok()
|
||
.unwrap_or_else(|| lookup(b"<sop>")),
|
||
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"<sop>"),
|
||
lookup(b"<think>"),
|
||
lookup(b"</think>"),
|
||
),
|
||
};
|
||
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<Vec<u8>> {
|
||
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<i32> {
|
||
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<i32> {
|
||
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<i32> {
|
||
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<i32>) {
|
||
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::<Vec<_>>();
|
||
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<i32>) {
|
||
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<i32>) {
|
||
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<u8> {
|
||
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))
|
||
}
|