use std::collections::HashMap;
use anyhow::{Context, Result};
use regex::Regex;
use crate::gguf::{GgufFile, GgufValue};
enum Segment<'a> {
Text(&'a str),
Special(u32),
}
pub struct BpeTokenizer {
vocab: Vec<Vec<u8>>,
token_to_id: HashMap<Vec<u8>, u32>,
merge_ranks: HashMap<(Vec<u8>, Vec<u8>), usize>,
special_tokens: HashMap<String, u32>,
bos_id: Option<u32>,
eos_id: Option<u32>,
add_bos: bool,
add_eos: bool,
chat_template: Option<String>,
pretokenize_re: Regex,
digits_split_bare: bool,
byte_to_unicode: [char; 256],
unicode_to_byte: HashMap<char, u8>,
}
impl BpeTokenizer {
pub fn from_gguf(gguf: &GgufFile) -> Result<Self> {
let tokens = gguf
.get_string_array("tokenizer.ggml.tokens")
.context("missing tokenizer.ggml.tokens")?;
let vocab_size = tokens.len();
let mut vocab: Vec<Vec<u8>> = Vec::with_capacity(vocab_size);
let mut token_to_id: HashMap<Vec<u8>, u32> = HashMap::with_capacity(vocab_size);
for (id, token_str) in tokens.iter().enumerate() {
let token_bytes = unescape_token(token_str);
token_to_id.insert(token_bytes.clone(), id as u32);
vocab.push(token_bytes);
}
let mut merge_ranks: HashMap<(Vec<u8>, Vec<u8>), usize> = HashMap::new();
if let Some(merges) = gguf.get_string_array("tokenizer.ggml.merges") {
for (rank, merge_str) in merges.iter().enumerate() {
if let Some((a, b)) = merge_str.split_once(' ') {
let a_bytes = unescape_token(a);
let b_bytes = unescape_token(b);
merge_ranks.insert((a_bytes, b_bytes), rank);
}
}
}
let mut special_tokens = HashMap::new();
if let Some(GgufValue::Array(types)) = gguf.metadata.get("tokenizer.ggml.token_type") {
for (id, type_val) in types.iter().enumerate() {
if let GgufValue::I32(t) = type_val {
if (*t == 3 || *t == 4) && id < vocab_size {
let token_str = &tokens[id];
special_tokens.insert(token_str.to_string(), id as u32);
}
}
}
}
let bos_id = gguf.get_u32("tokenizer.ggml.bos_token_id");
let eos_id = gguf.get_u32("tokenizer.ggml.eos_token_id");
let add_bos = gguf
.get_bool("tokenizer.ggml.add_bos_token")
.unwrap_or(false);
let add_eos = gguf
.get_bool("tokenizer.ggml.add_eos_token")
.unwrap_or(false);
let chat_template = gguf
.get_str("tokenizer.chat_template")
.map(|raw| strip_generation_markers(raw).into_owned());
let pre_type = gguf.get_str("tokenizer.ggml.pre").unwrap_or("gpt2");
let pretokenize_re = build_pretokenize_regex(pre_type);
let digits_split_bare = pre_type == "refact";
let byte_to_unicode = build_byte_to_unicode();
let unicode_to_byte = build_unicode_to_byte();
Ok(BpeTokenizer {
vocab,
token_to_id,
merge_ranks,
special_tokens,
bos_id,
eos_id,
add_bos,
add_eos,
chat_template,
pretokenize_re,
digits_split_bare,
byte_to_unicode,
unicode_to_byte,
})
}
pub fn encode(&self, text: &str) -> Vec<u32> {
if text.is_empty() {
return vec![];
}
let segments = self.split_special_tokens(text);
let mut result = Vec::new();
for segment in &segments {
match segment {
Segment::Special(id) => result.push(*id),
Segment::Text(s) => {
for chunk in self.pretokenize(s) {
result.extend(self.bpe_encode_chunk(chunk));
}
}
}
}
result
}
pub fn encode_special(&self, text: &str, add_special: bool) -> Vec<u32> {
let bos = if add_special && self.add_bos {
self.bos_id
} else {
None
};
let eos = if add_special && self.add_eos {
self.eos_id
} else {
None
};
if bos.is_none() && eos.is_none() {
return self.encode(text);
}
let encoded = self.encode(text);
let mut result =
Vec::with_capacity(encoded.len() + bos.is_some() as usize + eos.is_some() as usize);
result.extend(bos);
result.extend_from_slice(&encoded);
result.extend(eos);
result
}
fn pretokenize<'a>(&self, s: &'a str) -> Vec<&'a str> {
let mut ranges: Vec<(usize, usize)> = self
.pretokenize_re
.find_iter(s)
.map(|m| (m.start(), m.end()))
.collect();
for i in 0..ranges.len() {
let (a, b) = ranges[i];
let is_ws_run = b - a >= 2 && s.as_bytes()[a..b].iter().all(u8::is_ascii_whitespace);
let last_is_spacetab = matches!(s.as_bytes()[b - 1], b' ' | b'\t');
let next_char = s[b..].chars().next();
let next_non_ws = next_char.is_some_and(|c| !c.is_whitespace());
let next_takes_space =
!(self.digits_split_bare && next_char.is_some_and(char::is_numeric));
if is_ws_run
&& last_is_spacetab
&& next_non_ws
&& next_takes_space
&& i + 1 < ranges.len()
{
ranges[i].1 = b - 1;
ranges[i + 1].0 = b - 1;
}
}
ranges.iter().map(|&(a, b)| &s[a..b]).collect()
}
fn split_special_tokens<'a>(&self, text: &'a str) -> Vec<Segment<'a>> {
if self.special_tokens.is_empty() {
return vec![Segment::Text(text)];
}
let mut segments = Vec::new();
let mut remaining = text;
while !remaining.is_empty() {
let mut best: Option<(usize, usize, u32)> = None; for (tok_str, &tok_id) in &self.special_tokens {
if let Some(pos) = remaining.find(tok_str.as_str()) {
let end = pos + tok_str.len();
if best.is_none()
|| pos < best.unwrap().0
|| (pos == best.unwrap().0 && end > best.unwrap().1)
{
best = Some((pos, end, tok_id));
}
}
}
match best {
Some((start, end, id)) => {
if start > 0 {
segments.push(Segment::Text(&remaining[..start]));
}
segments.push(Segment::Special(id));
remaining = &remaining[end..];
}
None => {
segments.push(Segment::Text(remaining));
break;
}
}
}
segments
}
fn bpe_encode_chunk(&self, chunk: &str) -> Vec<u32> {
if chunk.is_empty() {
return vec![];
}
let unicode_str = bytes_to_gpt2_unicode(chunk.as_bytes(), &self.byte_to_unicode);
let mut tokens: Vec<Vec<u8>> = unicode_str
.chars()
.map(|c| {
let mut buf = [0u8; 4];
let s = c.encode_utf8(&mut buf);
s.as_bytes().to_vec()
})
.collect();
loop {
if tokens.len() < 2 {
break;
}
let mut best_rank = usize::MAX;
let mut best_idx = 0;
for i in 0..tokens.len() - 1 {
let pair = (tokens[i].clone(), tokens[i + 1].clone());
if let Some(&rank) = self.merge_ranks.get(&pair)
&& rank < best_rank
{
best_rank = rank;
best_idx = i;
}
}
if best_rank == usize::MAX {
break;
}
let merged = [tokens[best_idx].as_slice(), tokens[best_idx + 1].as_slice()].concat();
tokens[best_idx] = merged;
tokens.remove(best_idx + 1);
}
tokens
.iter()
.map(|t| self.token_to_id.get(t).copied().unwrap_or(0))
.collect()
}
pub fn decode(&self, token_ids: &[u32]) -> String {
String::from_utf8_lossy(&self.decode_bytes(token_ids)).into_owned()
}
pub fn decode_bytes(&self, token_ids: &[u32]) -> Vec<u8> {
let mut raw_bytes = Vec::new();
for &id in token_ids {
if let Some(token_bytes) = self.vocab.get(id as usize) {
match std::str::from_utf8(token_bytes) {
Ok(s) => {
for ch in s.chars() {
if let Some(&b) = self.unicode_to_byte.get(&ch) {
raw_bytes.push(b);
} else {
let mut buf = [0u8; 4];
let encoded = ch.encode_utf8(&mut buf);
raw_bytes.extend_from_slice(encoded.as_bytes());
}
}
}
Err(_) => {
raw_bytes.extend_from_slice(token_bytes);
}
}
}
}
raw_bytes
}
pub fn vocab_size(&self) -> usize {
self.vocab.len()
}
pub fn bos_token(&self) -> Option<u32> {
self.bos_id
}
pub fn eos_token(&self) -> Option<u32> {
self.eos_id
}
pub fn add_bos_token(&self) -> bool {
self.add_bos
}
pub fn add_eos_token(&self) -> bool {
self.add_eos
}
pub fn chat_template(&self) -> Option<&str> {
self.chat_template.as_deref()
}
pub fn special_token_id(&self, name: &str) -> Option<u32> {
self.special_tokens.get(name).copied()
}
pub fn is_special_token(&self, id: u32) -> bool {
self.special_tokens.values().any(|&v| v == id)
}
pub fn token_output_bytes(&self, id: u32) -> Vec<u8> {
let Some(token_bytes) = self.vocab.get(id as usize) else {
return Vec::new();
};
match std::str::from_utf8(token_bytes) {
Ok(s) => {
let mut out = Vec::with_capacity(token_bytes.len());
for ch in s.chars() {
if let Some(&b) = self.unicode_to_byte.get(&ch) {
out.push(b);
} else {
let mut buf = [0u8; 4];
out.extend_from_slice(ch.encode_utf8(&mut buf).as_bytes());
}
}
out
}
Err(_) => token_bytes.clone(),
}
}
}
#[derive(Debug, Clone, serde::Serialize)]
pub struct ChatMessage {
pub role: String,
pub content: String,
}
#[derive(Debug, Clone, serde::Serialize)]
#[serde(tag = "type", rename_all = "lowercase")]
pub enum ContentItem {
Text { text: String },
Image,
}
#[derive(Debug, Clone, serde::Serialize)]
pub struct ChatMessageMultimodal {
pub role: String,
pub content: Vec<ContentItem>,
}
pub fn apply_chat_template<M: serde::Serialize>(
tokenizer: &BpeTokenizer,
messages: &[M],
add_generation_prompt: bool,
) -> Result<String> {
apply_chat_template_with_tools(tokenizer, messages, &[], add_generation_prompt)
}
pub fn apply_chat_template_with_tools<M: serde::Serialize>(
tokenizer: &BpeTokenizer,
messages: &[M],
tools: &[crate::tools::ToolDef],
add_generation_prompt: bool,
) -> Result<String> {
let template_str = tokenizer
.chat_template()
.context("model has no chat template")?;
let mut env = minijinja::Environment::new();
env.add_template("chat", template_str)
.context("invalid chat template")?;
let tmpl = env.get_template("chat").unwrap();
let bos_token = tokenizer
.bos_token()
.and_then(|id| tokenizer.vocab.get(id as usize))
.map(|b| String::from_utf8_lossy(b).into_owned())
.unwrap_or_default();
let eos_token = tokenizer
.eos_token()
.and_then(|id| tokenizer.vocab.get(id as usize))
.map(|b| String::from_utf8_lossy(b).into_owned())
.unwrap_or_default();
let ctx = minijinja::context! {
messages => messages,
tools => tools,
bos_token => bos_token,
eos_token => eos_token,
add_generation_prompt => add_generation_prompt,
};
tmpl.render(ctx).context("rendering chat template")
}
fn strip_generation_markers(template: &str) -> std::borrow::Cow<'_, str> {
use std::sync::OnceLock;
static RE: OnceLock<Regex> = OnceLock::new();
let re = RE.get_or_init(|| {
Regex::new(r"\{%-?\s*(?:end)?generation\b\s*-?%\}").expect("static regex must compile")
});
re.replace_all(template, "")
}
fn build_byte_to_unicode() -> [char; 256] {
let mut table = ['\0'; 256];
let mut n = 0u32;
for b in 0u16..256 {
let ch = match b {
0x21..=0x7E | 0xA1..=0xAC | 0xAE..=0xFF => b as u32,
_ => {
let c = 256 + n;
n += 1;
c
}
};
table[b as usize] = char::from_u32(ch).unwrap();
}
table
}
fn bytes_to_gpt2_unicode(bytes: &[u8], table: &[char; 256]) -> String {
bytes.iter().map(|&b| table[b as usize]).collect()
}
fn build_unicode_to_byte() -> HashMap<char, u8> {
let table = build_byte_to_unicode();
table
.iter()
.enumerate()
.map(|(b, &ch)| (ch, b as u8))
.collect()
}
fn build_pretokenize_regex(pre_type: &str) -> Regex {
let pattern = match pre_type {
"lfm2" | "llama3" | "llama-v3" | "llama-bpe" => concat!(
r"(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])",
r"|[^\r\n\p{L}\p{N}]?\p{L}+",
r"|\p{N}{1,3}",
r"| ?[^\s\p{L}\p{N}]+[\r\n]*",
r"|\s*[\r\n]+",
r"|\s+",
),
"qwen2" => concat!(
r"(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])",
r"|[^\r\n\p{L}\p{N}]?\p{L}+",
r"|\p{N}",
r"| ?[^\s\p{L}\p{N}]+[\r\n]*",
r"|\s*[\r\n]+",
r"|\s+",
),
"gpt2" => concat!(
r"(?:'s|'t|'re|'ve|'m|'ll|'d)",
r"| ?\p{L}+",
r"| ?\p{N}+",
r"| ?[^\s\p{L}\p{N}]+",
r"|\s+",
),
"refact" => concat!(
r"(?:'s|'t|'re|'ve|'m|'ll|'d)",
r"| ?\p{L}+",
r"|\p{N}",
r"| ?[^\s\p{L}\p{N}]+",
r"|\s+",
),
other => {
tracing::warn!(
"unknown tokenizer.ggml.pre type '{other}', defaulting to LLAMA3 pretokenizer"
);
concat!(
r"(?:'[sS]|'[tT]|'[rR][eE]|'[vV][eE]|'[mM]|'[lL][lL]|'[dD])",
r"|[^\r\n\p{L}\p{N}]?\p{L}+",
r"|\p{N}{1,3}",
r"| ?[^\s\p{L}\p{N}]+[\r\n]*",
r"|\s*[\r\n]+",
r"|\s+",
)
}
};
Regex::new(pattern).expect("invalid pretokenizer regex")
}
fn unescape_token(s: &str) -> Vec<u8> {
if s.starts_with("<0x")
&& s.ends_with('>')
&& s.len() == 6
&& let Ok(byte) = u8::from_str_radix(&s[3..5], 16)
{
return vec![byte];
}
let s = s.replace('▁', " ");
s.into_bytes()
}
#[cfg(test)]
mod tests {
use super::*;
fn make_test_tokenizer() -> BpeTokenizer {
let mut vocab: Vec<Vec<u8>> = Vec::new();
let mut token_to_id: HashMap<Vec<u8>, u32> = HashMap::new();
for b in 0u8..=255 {
vocab.push(vec![b]);
token_to_id.insert(vec![b], b as u32);
}
let merged_tokens = vec![
(256u32, b"he".to_vec()),
(257, b"ll".to_vec()),
(258, b"lo".to_vec()),
(259, b"hell".to_vec()),
(260, b"hello".to_vec()),
];
for (id, bytes) in &merged_tokens {
vocab.push(bytes.clone());
token_to_id.insert(bytes.clone(), *id);
}
let mut merge_ranks = HashMap::new();
merge_ranks.insert((b"h".to_vec(), b"e".to_vec()), 0); merge_ranks.insert((b"l".to_vec(), b"l".to_vec()), 1); merge_ranks.insert((b"l".to_vec(), b"o".to_vec()), 2); merge_ranks.insert((b"he".to_vec(), b"ll".to_vec()), 3); merge_ranks.insert((b"hell".to_vec(), b"o".to_vec()), 4);
BpeTokenizer {
vocab,
token_to_id,
merge_ranks,
special_tokens: HashMap::new(),
bos_id: None,
eos_id: None,
add_bos: false,
add_eos: false,
chat_template: None,
pretokenize_re: build_pretokenize_regex("lfm2"),
digits_split_bare: false,
byte_to_unicode: build_byte_to_unicode(),
unicode_to_byte: build_unicode_to_byte(),
}
}
#[test]
fn test_byte_to_unicode_space() {
let table = build_byte_to_unicode();
assert_eq!(table[0x20], '\u{0120}');
assert_eq!(table[b'A' as usize], 'A');
assert_eq!(table[b'z' as usize], 'z');
assert_eq!(table[b'0' as usize], '0');
assert_ne!(table[0x0A], '\n');
}
#[test]
fn encode_special_prepends_bos_when_declared() {
let mut tok = make_test_tokenizer();
tok.bos_id = Some(1000);
tok.add_bos = true;
let plain = tok.encode("hello");
let special = tok.encode_special("hello", true);
assert_eq!(special[0], 1000);
assert_eq!(&special[1..], &plain[..]);
assert_eq!(tok.encode_special("hello", false), plain);
}
#[test]
fn encode_special_without_metadata_flag_is_plain_encode() {
let mut tok = make_test_tokenizer();
tok.bos_id = Some(1000);
assert_eq!(tok.encode_special("hello", true), tok.encode("hello"));
}
#[test]
fn encode_special_flag_without_bos_id_is_noop() {
let mut tok = make_test_tokenizer();
tok.add_bos = true;
assert_eq!(tok.encode_special("hello", true), tok.encode("hello"));
}
#[test]
fn encode_special_appends_eos_and_handles_empty_text() {
let mut tok = make_test_tokenizer();
tok.eos_id = Some(1001);
tok.add_eos = true;
let special = tok.encode_special("hello", true);
assert_eq!(special.last(), Some(&1001));
assert_eq!(&special[..special.len() - 1], &tok.encode("hello")[..]);
tok.bos_id = Some(1000);
tok.add_bos = true;
assert_eq!(tok.encode_special("", true), vec![1000, 1001]);
}
#[test]
fn test_byte_unicode_roundtrip() {
let table = build_byte_to_unicode();
let reverse = build_unicode_to_byte();
for b in 0u8..=255 {
let ch = table[b as usize];
assert_eq!(reverse[&ch], b, "roundtrip failed for byte {b:#04x}");
}
}
#[test]
fn test_pretokenize_splits_words() {
let re = build_pretokenize_regex("lfm2");
let chunks: Vec<&str> = re
.find_iter("The meaning of life")
.map(|m| m.as_str())
.collect();
assert_eq!(chunks, vec!["The", " meaning", " of", " life"]);
}
#[test]
fn test_pretokenize_contractions() {
let re = build_pretokenize_regex("lfm2");
let chunks: Vec<&str> = re.find_iter("I'm don't").map(|m| m.as_str()).collect();
assert!(chunks.contains(&"'m"));
assert!(chunks.contains(&"'t"));
}
#[test]
fn test_pretokenize_numbers() {
let re = build_pretokenize_regex("lfm2");
let chunks: Vec<&str> = re.find_iter("test 12345").map(|m| m.as_str()).collect();
assert_eq!(chunks, vec!["test", " ", "123", "45"]);
}
#[test]
fn test_pretokenize_refact_splits_each_digit() {
let re = build_pretokenize_regex("refact");
let chunks: Vec<&str> = re.find_iter("12345").map(|m| m.as_str()).collect();
assert_eq!(chunks, vec!["1", "2", "3", "4", "5"]);
}
#[test]
fn test_pretokenize_refact_contractions_and_words() {
let re = build_pretokenize_regex("refact");
let chunks: Vec<&str> = re.find_iter("I'm ok").map(|m| m.as_str()).collect();
assert_eq!(chunks, vec!["I", "'m", " ok"]);
}
#[test]
fn test_refact_keeps_whitespace_run_before_digit() {
let tok = make_pretok_tokenizer("refact");
assert_eq!(tok.pretokenize(" 3"), vec![" ", "3"]);
assert_eq!(tok.pretokenize("x 9"), vec!["x", " ", "9"]);
assert_eq!(tok.pretokenize("\n 3"), vec!["\n ", "3"]);
assert_eq!(tok.pretokenize(" 3"), vec![" ", "3"]);
}
#[test]
fn test_refact_still_donates_space_before_word() {
let tok = make_pretok_tokenizer("refact");
assert_eq!(tok.pretokenize(" x"), vec![" ", " x"]);
assert_eq!(tok.pretokenize("\n return"), vec!["\n ", " return"]);
}
#[test]
fn test_gpt2_digit_absorbs_donated_space() {
let tok = make_pretok_tokenizer("gpt2");
assert_eq!(tok.pretokenize(" 3"), vec![" ", " 3"]);
}
fn make_pretok_tokenizer(pre: &str) -> BpeTokenizer {
BpeTokenizer {
vocab: Vec::new(),
token_to_id: HashMap::new(),
merge_ranks: HashMap::new(),
special_tokens: HashMap::new(),
bos_id: None,
eos_id: None,
add_bos: false,
add_eos: false,
chat_template: None,
pretokenize_re: build_pretokenize_regex(pre),
digits_split_bare: pre == "refact",
byte_to_unicode: build_byte_to_unicode(),
unicode_to_byte: build_unicode_to_byte(),
}
}
fn make_gpt2_style_tokenizer() -> BpeTokenizer {
let table = build_byte_to_unicode();
let mut vocab: Vec<Vec<u8>> = Vec::new();
let mut token_to_id: HashMap<Vec<u8>, u32> = HashMap::new();
vocab.push(b"<pad>".to_vec());
token_to_id.insert(b"<pad>".to_vec(), 0);
for b in 0u8..=127 {
let ch = table[b as usize];
let mut buf = [0u8; 4];
let s = ch.encode_utf8(&mut buf);
let bytes = s.as_bytes().to_vec();
let id = vocab.len() as u32;
vocab.push(bytes.clone());
token_to_id.insert(bytes, id);
}
let space_char = table[b' ' as usize]; let add_token =
|vocab: &mut Vec<Vec<u8>>, map: &mut HashMap<Vec<u8>, u32>, s: &str| -> u32 {
let bytes = s.as_bytes().to_vec();
let id = vocab.len() as u32;
vocab.push(bytes.clone());
map.insert(bytes, id);
id
};
let _hi_id = add_token(&mut vocab, &mut token_to_id, "Hi");
let space_world = format!("{space_char}world");
let _world_id = add_token(&mut vocab, &mut token_to_id, &space_world);
let mut merge_ranks = HashMap::new();
merge_ranks.insert((b"H".to_vec(), b"i".to_vec()), 0);
let g_bytes = {
let mut buf = [0u8; 4];
space_char.encode_utf8(&mut buf).as_bytes().to_vec()
};
merge_ranks.insert((g_bytes.clone(), b"w".to_vec()), 1);
let gw_bytes = [g_bytes.as_slice(), b"w"].concat();
merge_ranks.insert((gw_bytes.clone(), b"o".to_vec()), 2);
let gwor = [gw_bytes.as_slice(), b"o"].concat();
merge_ranks.insert((gwor.clone(), b"r".to_vec()), 3);
let gworl = [gwor.as_slice(), b"r"].concat();
merge_ranks.insert((gworl.clone(), b"l".to_vec()), 4);
let gworld = [gworl.as_slice(), b"l"].concat();
merge_ranks.insert((gworld.clone(), b"d".to_vec()), 5);
BpeTokenizer {
vocab,
token_to_id,
merge_ranks,
special_tokens: HashMap::new(),
bos_id: None,
eos_id: None,
add_bos: false,
add_eos: false,
chat_template: None,
pretokenize_re: build_pretokenize_regex("lfm2"),
digits_split_bare: false,
byte_to_unicode: build_byte_to_unicode(),
unicode_to_byte: build_unicode_to_byte(),
}
}
#[test]
fn test_gpt2_encode_space_prefix() {
let tok = make_gpt2_style_tokenizer();
let ids = tok.encode("Hi world");
let hi_id = *tok.token_to_id.get(b"Hi".as_slice()).unwrap();
let table = build_byte_to_unicode();
let space_world = format!("{}world", table[b' ' as usize]);
let world_id = *tok.token_to_id.get(space_world.as_bytes()).unwrap();
assert_eq!(ids, vec![hi_id, world_id]);
}
#[test]
fn test_gpt2_decode_reverses_encode() {
let tok = make_gpt2_style_tokenizer();
let ids = tok.encode("Hi world");
let decoded = tok.decode(&ids);
assert_eq!(decoded, "Hi world");
}
#[test]
fn test_encode_single_bytes() {
let tok = make_test_tokenizer();
let ids = tok.encode("ab");
assert_eq!(ids, vec![b'a' as u32, b'b' as u32]);
}
#[test]
fn test_encode_with_merges() {
let tok = make_test_tokenizer();
let ids = tok.encode("hello");
assert_eq!(ids, vec![260]); }
#[test]
fn test_encode_partial_merges() {
let tok = make_test_tokenizer();
let ids = tok.encode("hell");
assert_eq!(ids, vec![259]); }
#[test]
fn test_decode_roundtrip() {
let tok = make_test_tokenizer();
let text = "hello";
let ids = tok.encode(text);
let decoded = tok.decode(&ids);
assert_eq!(decoded, text);
}
#[test]
fn test_decode_single_bytes() {
let tok = make_test_tokenizer();
let decoded = tok.decode(&[72, 105]); assert_eq!(decoded, "Hi");
}
#[test]
fn test_encode_empty() {
let tok = make_test_tokenizer();
assert_eq!(tok.encode(""), Vec::<u32>::new());
}
#[test]
fn test_unescape_byte_token() {
assert_eq!(unescape_token("<0x0A>"), vec![0x0A]); assert_eq!(unescape_token("<0xFF>"), vec![0xFF]);
assert_eq!(unescape_token("<0x00>"), vec![0x00]);
}
#[test]
fn test_unescape_space_marker() {
assert_eq!(unescape_token("▁hello"), b" hello");
assert_eq!(unescape_token("▁"), b" ");
}
#[test]
fn test_unescape_regular() {
assert_eq!(unescape_token("hello"), b"hello");
}
#[test]
fn test_chat_template_rendering() {
let mut tok = make_test_tokenizer();
tok.chat_template = Some(
"{% for msg in messages %}{{ msg.role }}: {{ msg.content }}\n{% endfor %}{% if add_generation_prompt %}assistant: {% endif %}"
.to_string(),
);
let messages = vec![ChatMessage {
role: "user".to_string(),
content: "Hello!".to_string(),
}];
let result = apply_chat_template(&tok, &messages, true).unwrap();
assert_eq!(result, "user: Hello!\nassistant: ");
}
#[test]
fn strip_generation_markers_removes_all_variants() {
let cases = [
("{% generation %}", ""),
("{%- generation -%}", ""),
("{%- generation %}", ""),
("{% generation -%}", ""),
("{% endgeneration %}", ""),
("{%- endgeneration -%}", ""),
("a{% generation %}b{% endgeneration %}c", "abc"),
(
"before{%- generation -%}inside{%- endgeneration -%}after",
"beforeinsideafter",
),
("no markers here", "no markers here"),
];
for (input, want) in cases {
let got = strip_generation_markers(input);
assert_eq!(&*got, want, "input: {input:?}");
}
}
#[test]
fn chat_template_with_generation_block_renders() {
let mut tok = make_test_tokenizer();
let raw = "{% for msg in messages %}\
{%- if msg.role == 'assistant' -%}\
{%- generation -%}{{ msg.role }}: {{ msg.content }}\n{%- endgeneration -%}\
{%- else -%}{{ msg.role }}: {{ msg.content }}\n{%- endif -%}\
{% endfor %}";
tok.chat_template = Some(strip_generation_markers(raw).into_owned());
let messages = vec![
ChatMessage {
role: "user".to_string(),
content: "Hi".to_string(),
},
ChatMessage {
role: "assistant".to_string(),
content: "Hello!".to_string(),
},
];
let got = apply_chat_template(&tok, &messages, false).unwrap();
assert!(
got.contains("user: Hi") && got.contains("assistant: Hello!"),
"expected user+assistant content; got {got:?}"
);
}
#[test]
fn chat_template_multimodal_emits_image_marker() {
let mut tok = make_test_tokenizer();
tok.chat_template = Some(
"{% for msg in messages %}\
{%- if msg.content is string -%}\
{{ msg.role }}: {{ msg.content }}\n\
{%- else -%}\
{{ msg.role }}: \
{%- for item in msg.content -%}\
{%- if item.type == 'image' -%}<image>\
{%- elif item.type == 'text' -%}{{ item.text }}\
{%- endif -%}\
{%- endfor -%}\n\
{%- endif -%}\
{% endfor %}"
.to_string(),
);
let text_only = vec![ChatMessage {
role: "user".to_string(),
content: "Hi".to_string(),
}];
let rendered = apply_chat_template(&tok, &text_only, false).unwrap();
assert!(
rendered.contains("user: Hi"),
"string-content render lost the body: {rendered:?}"
);
assert!(
!rendered.contains("<image>"),
"string-content render shouldn't emit <image>: {rendered:?}"
);
let multimodal = vec![ChatMessageMultimodal {
role: "user".to_string(),
content: vec![
ContentItem::Image,
ContentItem::Text {
text: " Describe.".to_string(),
},
],
}];
let rendered = apply_chat_template(&tok, &multimodal, false).unwrap();
assert!(
rendered.contains("<image> Describe."),
"list-content render didn't combine items in order: {rendered:?}"
);
assert_eq!(
rendered.matches("<image>").count(),
1,
"expected exactly one <image> marker; got {rendered:?}"
);
}
}