Implement DS4 model intake
This commit is contained in:
600
src/engine/tokenizer.rs
Normal file
600
src/engine/tokenizer.rs
Normal file
@@ -0,0 +1,600 @@
|
||||
use super::ModelFamily;
|
||||
use super::gguf::Gguf;
|
||||
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> {
|
||||
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));
|
||||
}
|
||||
output.push(self.user);
|
||||
output.extend(self.tokenize(prompt));
|
||||
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))
|
||||
}
|
||||
|
||||
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))
|
||||
}
|
||||
Reference in New Issue
Block a user