use crate::error::InferenceError;
use crate::tokenizer::common::{
JsonValue, ThreadSafeLruCache, TokenizedInput, Tokenizer, invert_vocab, json_object_to_vocab,
json_path, known_special_id, pad_ids, parse_added_tokens, parse_json,
parse_post_processor_flags, parse_rendered_added_tokens, push_eos_preserving_limit,
vocab_txt_to_map,
};
use std::cmp::Ordering;
use std::collections::{BinaryHeap, HashMap};
use std::fs;
use std::path::Path;
use std::sync::Arc;
use tracing::warn;
pub(crate) const DEFAULT_BPE_CACHE_CAPACITY: usize = 8_192;
pub(crate) const DEFAULT_BPE_MAX_SEQ_LEN: usize = 4_096;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum PreTokenizeMode {
Gpt4Regex,
}
#[derive(Debug, Clone)]
pub struct BpeTokenizer {
inner: Arc<BpeInner>,
}
#[derive(Debug)]
struct BpeInner {
vocab: HashMap<String, u32>,
id_to_token: Vec<String>,
merges: HashMap<String, HashMap<String, usize>>,
byte_encoder: Vec<char>,
special_tokens: HashMap<String, u32>,
special_tokens_sorted: Vec<String>,
added_render: HashMap<u32, String>,
pad_id: u32,
unk_id: Option<u32>,
bos_id: Option<u32>,
eos_id: Option<u32>,
add_bos: bool,
add_eos: bool,
pre_tokenize_mode: PreTokenizeMode,
max_seq_len: usize,
cache: ThreadSafeLruCache<String, Vec<u32>>,
}
#[derive(Debug, Clone)]
struct BpeNode {
token: String,
prev: Option<usize>,
next: Option<usize>,
alive: bool,
version: u64,
}
#[derive(Debug, Clone, Eq, PartialEq)]
struct Candidate {
rank: usize,
left: usize,
right: usize,
left_version: u64,
right_version: u64,
}
impl Ord for Candidate {
fn cmp(&self, other: &Self) -> Ordering {
other
.rank
.cmp(&self.rank)
.then_with(|| other.left.cmp(&self.left))
.then_with(|| other.right.cmp(&self.right))
}
}
impl PartialOrd for Candidate {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
#[derive(Debug, Default)]
struct TokenizeScratch {
ids: Vec<u32>,
word_ids: Vec<u32>,
}
impl BpeTokenizer {
pub fn from_files(vocab_path: &Path, merges_path: &Path) -> Result<Self, InferenceError> {
let vocab_text = fs::read_to_string(vocab_path).map_err(|e| {
InferenceError::Tokenizer(format!("failed to read {}: {e}", vocab_path.display()))
})?;
let merges_text = fs::read_to_string(merges_path).map_err(|e| {
InferenceError::Tokenizer(format!("failed to read {}: {e}", merges_path.display()))
})?;
let vocab_json = parse_json(&vocab_text)?;
let vocab = json_object_to_vocab(&vocab_json)?;
Self::from_vocab_and_merges(vocab, parse_merges_txt(&merges_text))
}
pub fn from_vocab_txt_and_merges(
vocab_txt_path: &Path,
merges_path: &Path,
) -> Result<Self, InferenceError> {
let vocab_text = fs::read_to_string(vocab_txt_path).map_err(|e| {
InferenceError::Tokenizer(format!("failed to read {}: {e}", vocab_txt_path.display()))
})?;
let merges_text = fs::read_to_string(merges_path).map_err(|e| {
InferenceError::Tokenizer(format!("failed to read {}: {e}", merges_path.display()))
})?;
Self::from_vocab_and_merges(
vocab_txt_to_map(&vocab_text),
parse_merges_txt(&merges_text),
)
}
pub fn from_tokenizer_json(path: &Path) -> Result<Self, InferenceError> {
let text = fs::read_to_string(path).map_err(|e| {
InferenceError::Tokenizer(format!("failed to read {}: {e}", path.display()))
})?;
Self::from_tokenizer_json_str(&text)
}
pub fn from_tokenizer_json_str(text: &str) -> Result<Self, InferenceError> {
let root = parse_json(text)?;
let model_type = json_path(&root, &["model", "type"])
.and_then(JsonValue::as_str)
.unwrap_or("");
if model_type != "BPE" {
return Err(InferenceError::Tokenizer(format!(
"expected BPE tokenizer.json model type, found {model_type:?}"
)));
}
let vocab =
json_object_to_vocab(json_path(&root, &["model", "vocab"]).ok_or_else(|| {
InferenceError::Tokenizer("tokenizer.json missing model.vocab".into())
})?)?;
let merges_value = json_path(&root, &["model", "merges"]).ok_or_else(|| {
InferenceError::Tokenizer("tokenizer.json missing model.merges".into())
})?;
let merges = parse_merges_json(merges_value)?;
let added = parse_added_tokens(&root);
let rendered_added = parse_rendered_added_tokens(&root);
let mut tokenizer = Self::from_vocab_and_merges_with_config(
vocab,
merges,
added,
rendered_added,
DEFAULT_BPE_CACHE_CAPACITY,
DEFAULT_BPE_MAX_SEQ_LEN,
)?;
if let Some(unk_token) =
json_path(&root, &["model", "unk_token"]).and_then(JsonValue::as_str)
{
tokenizer = tokenizer.with_unk_token(unk_token);
}
let pp = parse_post_processor_flags(&root);
if pp.add_eos {
tokenizer = tokenizer.with_add_eos();
if let Some(eos_id) = pp.eos_id {
tokenizer = tokenizer.with_eos_id(eos_id);
}
}
validate_bpe_pretokenizer(&root)?;
if detect_gpt4_regex_pretokenizer(&root) {
tokenizer = tokenizer.with_pre_tokenize_mode(PreTokenizeMode::Gpt4Regex);
}
Ok(tokenizer)
}
pub fn from_vocab_and_merges(
vocab: HashMap<String, u32>,
merges: Vec<(String, String)>,
) -> Result<Self, InferenceError> {
Self::from_vocab_and_merges_with_config(
vocab,
merges,
HashMap::new(),
HashMap::new(),
DEFAULT_BPE_CACHE_CAPACITY,
DEFAULT_BPE_MAX_SEQ_LEN,
)
}
pub(crate) fn from_vocab_and_merges_with_config(
vocab: HashMap<String, u32>,
merges: Vec<(String, String)>,
added_tokens: HashMap<String, u32>,
rendered_added: HashMap<u32, String>,
cache_capacity: usize,
max_seq_len: usize,
) -> Result<Self, InferenceError> {
let id_to_token = invert_vocab(&vocab)?;
let added_render: HashMap<u32, String> = rendered_added;
let mut merge_ranks: HashMap<String, HashMap<String, usize>> = HashMap::new();
for (rank, (left, right)) in merges.into_iter().enumerate() {
merge_ranks
.entry(left)
.or_default()
.entry(right)
.or_insert(rank);
}
let mut special_tokens = added_tokens;
special_tokens.retain(|name, _| !name.is_empty());
for name in [
"<|endoftext|>",
"<|im_start|>",
"<|im_end|>",
"<pad>",
"<|pad|>",
"<bos>",
"<eos>",
"<unk>",
] {
if let Some(&id) = vocab.get(name) {
special_tokens.entry(name.to_string()).or_insert(id);
}
}
let pad_id = known_special_id(&vocab, &["<|pad|>", "<pad>", "[PAD]"])
.or_else(|| known_special_id(&special_tokens, &["<|pad|>", "<pad>"]))
.or_else(|| known_special_id(&vocab, &["<|endoftext|>", "</s>"]))
.unwrap_or(0);
let unk_id = known_special_id(&vocab, &["<unk>", "[UNK]"]);
let bos_id = known_special_id(&vocab, &["<bos>", "<s>"]);
let eos_id = known_special_id(&vocab, &["<eos>", "</s>", "<|endoftext|>"]);
let mut special_tokens_sorted: Vec<String> = special_tokens.keys().cloned().collect();
special_tokens_sorted.sort_by(|a, b| b.len().cmp(&a.len()).then_with(|| a.cmp(b)));
let inner = BpeInner {
vocab,
id_to_token,
merges: merge_ranks,
byte_encoder: bytes_to_unicode(),
special_tokens,
special_tokens_sorted,
added_render,
pad_id,
unk_id,
bos_id,
eos_id,
add_bos: false,
add_eos: false,
pre_tokenize_mode: PreTokenizeMode::Gpt4Regex,
max_seq_len,
cache: ThreadSafeLruCache::new(cache_capacity),
};
Ok(Self {
inner: Arc::new(inner),
})
}
pub fn with_max_seq_len(&self, max_seq_len: usize) -> Self {
let inner = BpeInner {
vocab: self.inner.vocab.clone(),
id_to_token: self.inner.id_to_token.clone(),
merges: self.inner.merges.clone(),
byte_encoder: self.inner.byte_encoder.clone(),
special_tokens: self.inner.special_tokens.clone(),
special_tokens_sorted: self.inner.special_tokens_sorted.clone(),
added_render: self.inner.added_render.clone(),
pad_id: self.inner.pad_id,
unk_id: self.inner.unk_id,
bos_id: self.inner.bos_id,
eos_id: self.inner.eos_id,
add_bos: self.inner.add_bos,
add_eos: self.inner.add_eos,
pre_tokenize_mode: self.inner.pre_tokenize_mode,
max_seq_len,
cache: self.inner.cache.clone(),
};
Self {
inner: Arc::new(inner),
}
}
pub fn with_unk_token(&self, token: &str) -> Self {
let unk_id = self.inner.vocab.get(token).copied();
let inner = BpeInner {
vocab: self.inner.vocab.clone(),
id_to_token: self.inner.id_to_token.clone(),
merges: self.inner.merges.clone(),
byte_encoder: self.inner.byte_encoder.clone(),
special_tokens: self.inner.special_tokens.clone(),
special_tokens_sorted: self.inner.special_tokens_sorted.clone(),
added_render: self.inner.added_render.clone(),
pad_id: self.inner.pad_id,
unk_id,
bos_id: self.inner.bos_id,
eos_id: self.inner.eos_id,
add_bos: self.inner.add_bos,
add_eos: self.inner.add_eos,
pre_tokenize_mode: self.inner.pre_tokenize_mode,
max_seq_len: self.inner.max_seq_len,
cache: self.inner.cache.clone(),
};
Self {
inner: Arc::new(inner),
}
}
pub fn with_add_eos(self) -> Self {
let inner = BpeInner {
vocab: self.inner.vocab.clone(),
id_to_token: self.inner.id_to_token.clone(),
merges: self.inner.merges.clone(),
byte_encoder: self.inner.byte_encoder.clone(),
special_tokens: self.inner.special_tokens.clone(),
special_tokens_sorted: self.inner.special_tokens_sorted.clone(),
added_render: self.inner.added_render.clone(),
pad_id: self.inner.pad_id,
unk_id: self.inner.unk_id,
bos_id: self.inner.bos_id,
eos_id: self.inner.eos_id,
add_bos: self.inner.add_bos,
add_eos: true,
pre_tokenize_mode: self.inner.pre_tokenize_mode,
max_seq_len: self.inner.max_seq_len,
cache: self.inner.cache.clone(),
};
Self {
inner: Arc::new(inner),
}
}
fn with_eos_id(self, eos_id: u32) -> Self {
let inner = BpeInner {
vocab: self.inner.vocab.clone(),
id_to_token: self.inner.id_to_token.clone(),
merges: self.inner.merges.clone(),
byte_encoder: self.inner.byte_encoder.clone(),
special_tokens: self.inner.special_tokens.clone(),
special_tokens_sorted: self.inner.special_tokens_sorted.clone(),
added_render: self.inner.added_render.clone(),
pad_id: self.inner.pad_id,
unk_id: self.inner.unk_id,
bos_id: self.inner.bos_id,
eos_id: Some(eos_id),
add_bos: self.inner.add_bos,
add_eos: self.inner.add_eos,
pre_tokenize_mode: self.inner.pre_tokenize_mode,
max_seq_len: self.inner.max_seq_len,
cache: self.inner.cache.clone(),
};
Self {
inner: Arc::new(inner),
}
}
fn with_pre_tokenize_mode(self, mode: PreTokenizeMode) -> Self {
let inner = BpeInner {
vocab: self.inner.vocab.clone(),
id_to_token: self.inner.id_to_token.clone(),
merges: self.inner.merges.clone(),
byte_encoder: self.inner.byte_encoder.clone(),
special_tokens: self.inner.special_tokens.clone(),
special_tokens_sorted: self.inner.special_tokens_sorted.clone(),
added_render: self.inner.added_render.clone(),
pad_id: self.inner.pad_id,
unk_id: self.inner.unk_id,
bos_id: self.inner.bos_id,
eos_id: self.inner.eos_id,
add_bos: self.inner.add_bos,
add_eos: self.inner.add_eos,
pre_tokenize_mode: mode,
max_seq_len: self.inner.max_seq_len,
cache: self.inner.cache.clone(),
};
Self {
inner: Arc::new(inner),
}
}
pub fn max_seq_len(&self) -> usize {
self.inner.max_seq_len
}
pub fn vocab_size(&self) -> usize {
self.inner.id_to_token.len()
}
pub fn token_for_id(&self, id: u32) -> Option<&str> {
self.inner
.id_to_token
.get(id as usize)
.map(String::as_str)
.or_else(|| self.inner.added_render.get(&id).map(String::as_str))
}
pub fn special_token_id(&self, name: &str) -> Option<u32> {
self.inner.special_tokens.get(name).copied()
}
pub(crate) fn token_ids_for_content(&self, content: &str) -> Vec<u32> {
let mut ids: Vec<u32> = Vec::new();
for id in self
.inner
.special_tokens
.get(content)
.into_iter()
.chain(self.inner.vocab.get(content))
.copied()
.chain(
self.inner
.added_render
.iter()
.filter(|(_, token)| token.as_str() == content)
.map(|(&id, _)| id),
)
{
if !ids.contains(&id) {
ids.push(id);
}
}
ids
}
pub fn vocab_bytes(&self, vocab_size: usize) -> Result<Vec<Vec<u8>>, InferenceError> {
let required_vocab_size = self
.inner
.special_tokens
.values()
.chain(self.inner.added_render.keys())
.copied()
.max()
.map_or(0usize, |id| id as usize + 1);
let required_vocab_size = self.inner.id_to_token.len().max(required_vocab_size);
if vocab_size < required_vocab_size {
return Err(InferenceError::Tokenizer(format!(
"model vocabulary size {vocab_size} cannot represent tokenizer token ID {}",
required_vocab_size - 1
)));
}
let byte_decoder = byte_decoder();
let mut vocab_bytes = Vec::with_capacity(vocab_size);
for id in 0..vocab_size {
let mut bytes = Vec::new();
self.append_token_bytes(id as u32, &byte_decoder, &mut bytes);
vocab_bytes.push(bytes);
}
Ok(vocab_bytes)
}
pub fn token_bytes_for_id(&self, id: u32) -> Option<Vec<u8>> {
let byte_decoder = byte_decoder();
let mut bytes = Vec::new();
self.append_token_bytes(id, &byte_decoder, &mut bytes)
.then_some(bytes)
}
pub(crate) fn append_token_bytes(
&self,
id: u32,
byte_decoder: &HashMap<char, u8>,
out: &mut Vec<u8>,
) -> bool {
if let Some(content) = self.inner.added_render.get(&id) {
out.extend_from_slice(content.as_bytes());
return true;
}
let Some(token_str) = self.inner.id_to_token.get(id as usize) else {
return false;
};
append_byte_decoded_token_bytes(token_str, byte_decoder, out);
true
}
fn tokenize_to_ids(&self, text: &str) -> Vec<u32> {
let mut scratch = TokenizeScratch::default();
self.tokenize_to_ids_into(text, &mut scratch)
}
#[cfg(all(feature = "metal-gpu", feature = "serve"))]
pub(crate) fn tokenize_fragments_with_inserted_ids(
&self,
before: &str,
inserted_ids: &[u32],
after: &str,
) -> Vec<u32> {
let mut scratch = TokenizeScratch::default();
scratch.ids.clear();
if self.inner.add_bos
&& let Some(bos_id) = self.inner.bos_id
{
scratch.ids.push(bos_id);
}
self.tokenize_text_into(before, &mut scratch);
scratch.ids.extend_from_slice(inserted_ids);
self.tokenize_text_into(after, &mut scratch);
if self.inner.add_eos
&& let Some(eos_id) = self.inner.eos_id
{
scratch.ids.push(eos_id);
}
scratch.ids
}
fn tokenize_to_ids_into(&self, text: &str, scratch: &mut TokenizeScratch) -> Vec<u32> {
scratch.ids.clear();
if self.inner.add_bos
&& let Some(bos_id) = self.inner.bos_id
{
scratch.ids.push(bos_id);
}
self.tokenize_text_into(text, scratch);
if self.inner.add_eos {
if let Some(eos_id) = self.inner.eos_id {
if scratch.ids.len() >= self.inner.max_seq_len {
warn!(
original_len = scratch.ids.len().saturating_add(1),
max_seq_len = self.inner.max_seq_len,
"truncating BPE tokenized input to preserve EOS within max_seq_len"
);
}
push_eos_preserving_limit(&mut scratch.ids, eos_id, self.inner.max_seq_len);
} else if scratch.ids.len() > self.inner.max_seq_len {
warn!(
original_len = scratch.ids.len(),
max_seq_len = self.inner.max_seq_len,
"truncating BPE tokenized input to max_seq_len"
);
scratch.ids.truncate(self.inner.max_seq_len);
}
} else if scratch.ids.len() > self.inner.max_seq_len {
warn!(
original_len = scratch.ids.len(),
max_seq_len = self.inner.max_seq_len,
"truncating BPE tokenized input to max_seq_len"
);
scratch.ids.truncate(self.inner.max_seq_len);
}
scratch.ids.clone()
}
fn tokenize_text_into(&self, text: &str, scratch: &mut TokenizeScratch) {
let mut segment_start = 0usize;
let mut pos = 0usize;
while pos < text.len() {
if let Some((special_end, special_id)) = self.match_special(text, pos) {
if segment_start < pos {
self.tokenize_regular_segment_into(&text[segment_start..pos], scratch);
}
scratch.ids.push(special_id);
pos = special_end;
segment_start = pos;
continue;
}
let Some(ch) = text[pos..].chars().next() else {
break;
};
pos += ch.len_utf8();
}
if segment_start < text.len() {
self.tokenize_regular_segment_into(&text[segment_start..], scratch);
}
}
fn tokenize_regular_segment_into(&self, text: &str, scratch: &mut TokenizeScratch) {
let pieces = match self.inner.pre_tokenize_mode {
PreTokenizeMode::Gpt4Regex => gpt4_regex_pretokenize(text),
};
for piece in pieces {
if let Some(cached) = self.inner.cache.get(&piece) {
scratch.ids.extend(cached.iter().copied());
continue;
}
scratch.word_ids.clear();
self.encode_piece_to_ids(&piece, &mut scratch.word_ids);
self.inner.cache.insert(piece, scratch.word_ids.clone());
scratch.ids.extend(scratch.word_ids.iter().copied());
}
}
fn match_special(&self, text: &str, pos: usize) -> Option<(usize, u32)> {
let tail = &text[pos..];
for token in &self.inner.special_tokens_sorted {
if tail.starts_with(token)
&& let Some(&id) = self.inner.special_tokens.get(token)
{
return Some((pos + token.len(), id));
}
}
None
}
fn encode_piece_to_ids(&self, piece: &str, out: &mut Vec<u32>) {
let encoded = self.byte_encode(piece);
let merged = self.bpe_merge(&encoded);
for token in merged {
if let Some(&id) = self.inner.vocab.get(token.as_str()) {
out.push(id);
continue;
}
let fallback_start = out.len();
let mut recovered = true;
for ch in token.chars() {
let one = ch.to_string();
if let Some(&id) = self.inner.vocab.get(one.as_str()) {
out.push(id);
} else {
recovered = false;
break;
}
}
if recovered {
continue;
}
out.truncate(fallback_start);
if let Some(unk_id) = self.inner.unk_id {
out.push(unk_id);
}
}
}
fn byte_encode(&self, text: &str) -> String {
let mut out = String::with_capacity(text.len());
for &byte in text.as_bytes() {
out.push(self.inner.byte_encoder[byte as usize]);
}
out
}
fn merge_rank(&self, left: &str, right: &str) -> Option<usize> {
self.inner
.merges
.get(left)
.and_then(|inner| inner.get(right))
.copied()
}
fn push_candidate(&self, nodes: &[BpeNode], heap: &mut BinaryHeap<Candidate>, left: usize) {
let Some(right) = nodes[left].next else {
return;
};
if !nodes[left].alive || !nodes[right].alive {
return;
}
let Some(rank) = self.merge_rank(nodes[left].token.as_str(), nodes[right].token.as_str())
else {
return;
};
heap.push(Candidate {
rank,
left,
right,
left_version: nodes[left].version,
right_version: nodes[right].version,
});
}
fn bpe_merge(&self, encoded: &str) -> Vec<String> {
let mut nodes: Vec<BpeNode> = encoded
.chars()
.enumerate()
.map(|(idx, ch)| BpeNode {
token: ch.to_string(),
prev: idx.checked_sub(1),
next: None,
alive: true,
version: 0,
})
.collect();
if nodes.is_empty() {
return Vec::new();
}
for idx in 0..nodes.len().saturating_sub(1) {
nodes[idx].next = Some(idx + 1);
}
let mut heap = BinaryHeap::new();
for idx in 0..nodes.len().saturating_sub(1) {
self.push_candidate(&nodes, &mut heap, idx);
}
while let Some(candidate) = heap.pop() {
if candidate.left >= nodes.len() || candidate.right >= nodes.len() {
continue;
}
let left = &nodes[candidate.left];
let right = &nodes[candidate.right];
if !left.alive
|| !right.alive
|| left.next != Some(candidate.right)
|| left.version != candidate.left_version
|| right.version != candidate.right_version
{
continue;
}
let right_next = nodes[candidate.right].next;
let right_token = nodes[candidate.right].token.clone();
nodes[candidate.left].token.push_str(right_token.as_str());
nodes[candidate.left].version = nodes[candidate.left].version.wrapping_add(1);
nodes[candidate.left].next = right_next;
if let Some(next) = right_next {
nodes[next].prev = Some(candidate.left);
}
nodes[candidate.right].alive = false;
nodes[candidate.right].version = nodes[candidate.right].version.wrapping_add(1);
if let Some(prev) = nodes[candidate.left].prev {
self.push_candidate(&nodes, &mut heap, prev);
}
self.push_candidate(&nodes, &mut heap, candidate.left);
}
let mut first = 0usize;
while first < nodes.len() && !nodes[first].alive {
first += 1;
}
if first >= nodes.len() {
return Vec::new();
}
let mut out = Vec::new();
let mut current = Some(first);
while let Some(idx) = current {
if nodes[idx].alive {
out.push(nodes[idx].token.clone());
}
current = nodes[idx].next;
}
out
}
}
impl Tokenizer for BpeTokenizer {
fn tokenize(&self, text: &str) -> TokenizedInput {
let ids = self.tokenize_to_ids(text);
pad_ids(ids, self.inner.max_seq_len, self.inner.pad_id)
}
fn tokenize_batch(&self, texts: &[&str]) -> Vec<TokenizedInput> {
if texts.is_empty() {
return Vec::new();
}
let mut scratch = TokenizeScratch::default();
let mut max_len = 0usize;
let mut all = Vec::with_capacity(texts.len());
for text in texts {
let ids = self.tokenize_to_ids_into(text, &mut scratch);
max_len = max_len.max(ids.len());
all.push(ids);
}
all.into_iter()
.map(|ids| pad_ids(ids, max_len, self.inner.pad_id))
.collect()
}
fn decode(&self, ids: &[u32]) -> Option<String> {
let byte_decoder = byte_decoder();
let mut bytes = Vec::new();
for &id in ids {
self.append_token_bytes(id, &byte_decoder, &mut bytes);
}
Some(String::from_utf8_lossy(&bytes).into_owned())
}
fn vocab_size(&self) -> usize {
self.inner.id_to_token.len()
}
fn max_seq_len(&self) -> usize {
self.inner.max_seq_len
}
}
fn parse_merges_txt(text: &str) -> Vec<(String, String)> {
let mut merges = Vec::new();
for line in text.lines() {
let line = line.trim();
if line.is_empty() || line.starts_with('#') {
continue;
}
let mut parts = line.split_whitespace();
let Some(left) = parts.next() else { continue };
let Some(right) = parts.next() else { continue };
merges.push((left.to_string(), right.to_string()));
}
merges
}
pub(crate) fn parse_merges_json(
value: &JsonValue,
) -> Result<Vec<(String, String)>, InferenceError> {
let array = value.as_array().ok_or_else(|| {
InferenceError::Tokenizer("expected tokenizer.json model.merges array".into())
})?;
let mut merges = Vec::with_capacity(array.len());
for item in array {
match item {
JsonValue::String(value) => {
let mut parts = value.split_whitespace();
let left = parts.next().ok_or_else(|| {
InferenceError::Tokenizer(format!("invalid merge entry {value:?}"))
})?;
let right = parts.next().ok_or_else(|| {
InferenceError::Tokenizer(format!("invalid merge entry {value:?}"))
})?;
merges.push((left.to_string(), right.to_string()));
}
JsonValue::Array(items) if items.len() == 2 => {
let left = items[0].as_str().ok_or_else(|| {
InferenceError::Tokenizer("invalid merge left operand".into())
})?;
let right = items[1].as_str().ok_or_else(|| {
InferenceError::Tokenizer("invalid merge right operand".into())
})?;
merges.push((left.to_string(), right.to_string()));
}
_ => {
return Err(InferenceError::Tokenizer(
"unsupported tokenizer.json merge entry".into(),
));
}
}
}
Ok(merges)
}
fn bytes_to_unicode() -> Vec<char> {
let mut bs = Vec::new();
bs.extend(33u16..=126);
bs.extend(161u16..=172);
bs.extend(174u16..=255);
let mut cs = bs.clone();
let mut n = 0u16;
for b in 0u16..=255u16 {
if !bs.contains(&b) {
bs.push(b);
cs.push(256 + n);
n += 1;
}
}
let mut table = vec!['\0'; 256];
for (b, c) in bs.into_iter().zip(cs) {
table[b as usize] =
char::from_u32(c as u32).expect("invariant: byte-to-unicode codepoint is valid");
}
table
}
pub fn byte_decode_token_bytes(token_str: &str) -> Vec<u8> {
let decoder = byte_decoder();
let mut bytes = Vec::new();
append_byte_decoded_token_bytes(token_str, &decoder, &mut bytes);
bytes
}
fn byte_decoder() -> HashMap<char, u8> {
bytes_to_unicode()
.into_iter()
.enumerate()
.map(|(byte, ch)| (ch, byte as u8))
.collect()
}
fn append_byte_decoded_token_bytes(
token_str: &str,
byte_decoder: &HashMap<char, u8>,
out: &mut Vec<u8>,
) {
out.extend(
token_str
.chars()
.filter_map(|ch| byte_decoder.get(&ch).copied()),
);
}
pub fn byte_decode_token(token_str: &str) -> String {
String::from_utf8_lossy(&byte_decode_token_bytes(token_str)).to_string()
}
fn detect_gpt4_regex_pretokenizer(root: &JsonValue) -> bool {
let Some(pt) = root.get("pre_tokenizer") else {
return false;
};
has_regex_split(pt)
}
fn has_regex_split(pt: &JsonValue) -> bool {
let pt_type = pt.get("type").and_then(JsonValue::as_str).unwrap_or("");
match pt_type {
"Split" => pt.get("pattern").and_then(|p| p.get("Regex")).is_some(),
"Sequence" => pt
.get("pretokenizers")
.and_then(JsonValue::as_array)
.is_some_and(|arr| arr.iter().any(has_regex_split)),
_ => false,
}
}
fn validate_bpe_pretokenizer(root: &JsonValue) -> Result<(), InferenceError> {
let Some(pt) = root.get("pre_tokenizer") else {
return Ok(());
};
if is_supported_bpe_pretokenizer(pt) {
return Ok(());
}
Err(InferenceError::Tokenizer(format!(
"tokenizer.json declares an unsupported BPE pre_tokenizer \
configuration ({}); refusing to silently fall back to the GPT-2 \
regex default, which could split text the declared pre_tokenizer \
would not",
describe_pretokenizer(pt)
)))
}
fn is_regex_split_node(pt: &JsonValue) -> bool {
let pt_type = pt.get("type").and_then(JsonValue::as_str).unwrap_or("");
pt_type == "Split" && pt.get("pattern").and_then(|p| p.get("Regex")).is_some()
}
fn is_supported_bpe_pretokenizer(pt: &JsonValue) -> bool {
if is_regex_split_node(pt) {
return true;
}
let pt_type = pt.get("type").and_then(JsonValue::as_str).unwrap_or("");
match pt_type {
"ByteLevel" => pt
.get("use_regex")
.and_then(JsonValue::as_bool)
.unwrap_or(true),
"Sequence" => pt
.get("pretokenizers")
.and_then(JsonValue::as_array)
.is_some_and(|arr| {
let mut seen_regex_split = false;
for child in arr {
if is_regex_split_node(child) {
seen_regex_split = true;
continue;
}
let child_type = child.get("type").and_then(JsonValue::as_str).unwrap_or("");
let child_ok = if child_type == "ByteLevel" {
let use_regex = child
.get("use_regex")
.and_then(JsonValue::as_bool)
.unwrap_or(true);
use_regex || seen_regex_split
} else {
false
};
if !child_ok {
return false;
}
}
true
}),
_ => false,
}
}
fn describe_pretokenizer(pt: &JsonValue) -> String {
let pt_type = pt
.get("type")
.and_then(JsonValue::as_str)
.unwrap_or("<untyped>");
if pt_type == "ByteLevel" {
let use_regex = pt
.get("use_regex")
.and_then(JsonValue::as_bool)
.unwrap_or(true);
return format!("ByteLevel{{use_regex:{use_regex}}}");
}
pt_type.to_string()
}
fn gpt4_regex_pretokenize(text: &str) -> Vec<String> {
let chars: Vec<char> = text.chars().collect();
let mut pieces = Vec::new();
let mut pos = 0;
while pos < chars.len() {
if let Some(end) = try_contraction(&chars, pos) {
pieces.push(chars[pos..end].iter().collect());
pos = end;
} else if let Some(end) = try_prefix_letters(&chars, pos) {
pieces.push(chars[pos..end].iter().collect());
pos = end;
} else if chars[pos].is_numeric() {
pieces.push(chars[pos].to_string());
pos += 1;
} else if let Some(end) = try_punctuation_run(&chars, pos) {
pieces.push(chars[pos..end].iter().collect());
pos = end;
} else if let Some(end) = try_newline_run(&chars, pos) {
pieces.push(chars[pos..end].iter().collect());
pos = end;
} else if let Some(end) = try_trailing_ws(&chars, pos) {
pieces.push(chars[pos..end].iter().collect());
pos = end;
} else if chars[pos].is_whitespace() {
let start = pos;
while pos < chars.len() && chars[pos].is_whitespace() {
pos += 1;
}
pieces.push(chars[start..pos].iter().collect());
} else {
pieces.push(chars[pos].to_string());
pos += 1;
}
}
pieces
}
fn eq_ci(a: char, lower: char) -> bool {
a.to_ascii_lowercase() == lower
}
fn try_contraction(chars: &[char], pos: usize) -> Option<usize> {
if chars.get(pos).copied() != Some('\'') {
return None;
}
let rest = &chars[pos + 1..];
if rest.len() >= 2 && eq_ci(rest[0], 'l') && eq_ci(rest[1], 'l') {
return Some(pos + 3);
}
if rest.len() >= 2 && eq_ci(rest[0], 'r') && eq_ci(rest[1], 'e') {
return Some(pos + 3);
}
if rest.len() >= 2 && eq_ci(rest[0], 'v') && eq_ci(rest[1], 'e') {
return Some(pos + 3);
}
if !rest.is_empty() {
let c = rest[0].to_ascii_lowercase();
if matches!(c, 's' | 't' | 'm' | 'd') {
return Some(pos + 2);
}
}
None
}
fn try_prefix_letters(chars: &[char], pos: usize) -> Option<usize> {
let mut i = pos;
if i < chars.len()
&& !chars[i].is_alphabetic()
&& !chars[i].is_numeric()
&& chars[i] != '\r'
&& chars[i] != '\n'
{
i += 1;
}
let start = i;
while i < chars.len() && chars[i].is_alphabetic() {
i += 1;
}
if i > start { Some(i) } else { None }
}
fn try_punctuation_run(chars: &[char], pos: usize) -> Option<usize> {
let mut i = pos;
if i < chars.len() && chars[i] == ' ' {
i += 1;
}
let start = i;
while i < chars.len()
&& !chars[i].is_whitespace()
&& !chars[i].is_alphabetic()
&& !chars[i].is_numeric()
{
i += 1;
}
if i == start {
return None;
}
while i < chars.len() && (chars[i] == '\r' || chars[i] == '\n') {
i += 1;
}
Some(i)
}
fn try_newline_run(chars: &[char], pos: usize) -> Option<usize> {
let mut i = pos;
while i < chars.len() && chars[i].is_whitespace() && chars[i] != '\r' && chars[i] != '\n' {
i += 1;
}
let nl_start = i;
while i < chars.len() && (chars[i] == '\r' || chars[i] == '\n') {
i += 1;
}
if i > nl_start { Some(i) } else { None }
}
fn try_trailing_ws(chars: &[char], pos: usize) -> Option<usize> {
if !chars.get(pos).is_some_and(|c| c.is_whitespace()) {
return None;
}
let mut i = pos;
while i < chars.len() && chars[i].is_whitespace() {
i += 1;
}
if i == chars.len() {
Some(i)
} else if i > pos + 1 {
Some(i - 1)
} else {
None
}
}
#[cfg(test)]
mod tests {
use super::*;
fn synthetic_bpe() -> BpeTokenizer {
let mut vocab = HashMap::new();
vocab.insert("h".to_string(), 0);
vocab.insert("e".to_string(), 1);
vocab.insert("l".to_string(), 2);
vocab.insert("o".to_string(), 3);
vocab.insert("Ġ".to_string(), 4);
vocab.insert("w".to_string(), 5);
vocab.insert("r".to_string(), 6);
vocab.insert("d".to_string(), 7);
vocab.insert("he".to_string(), 8);
vocab.insert("hel".to_string(), 9);
vocab.insert("hell".to_string(), 10);
vocab.insert("hello".to_string(), 11);
vocab.insert("Ġw".to_string(), 12);
vocab.insert("Ġwo".to_string(), 13);
vocab.insert("Ġwor".to_string(), 14);
vocab.insert("Ġworl".to_string(), 15);
vocab.insert("Ġworld".to_string(), 16);
vocab.insert("<|endoftext|>".to_string(), 17);
let merges = vec![
("h".to_string(), "e".to_string()),
("he".to_string(), "l".to_string()),
("hel".to_string(), "l".to_string()),
("hell".to_string(), "o".to_string()),
("Ġ".to_string(), "w".to_string()),
("Ġw".to_string(), "o".to_string()),
("Ġwo".to_string(), "r".to_string()),
("Ġwor".to_string(), "l".to_string()),
("Ġworl".to_string(), "d".to_string()),
];
BpeTokenizer::from_vocab_and_merges(vocab, merges).unwrap()
}
#[test]
fn test_gpt4_regex_pretokenize_matches_hf_bytelevel_use_regex_true() {
assert_eq!(
gpt4_regex_pretokenize("hello\nworld"),
vec!["hello", "\n", "world"]
);
assert_eq!(gpt4_regex_pretokenize("a'st"), vec!["a", "'s", "t"]);
}
#[test]
fn test_gpt2_raw_vocab_defaults_to_gpt4_regex_pretokenize() {
let mut vocab = HashMap::new();
for (s, i) in [("a", 0u32), ("'", 1), ("s", 2), ("t", 3), ("st", 4)] {
vocab.insert(s.to_string(), i);
}
let merges = vec![("s".to_string(), "t".to_string())];
let tokenizer = BpeTokenizer::from_vocab_and_merges(vocab, merges).unwrap();
assert_eq!(tokenizer.tokenize_to_ids("a'st"), vec![0, 1, 2, 3]);
}
fn cross_space_merge_json(pre_tokenizer: &str) -> String {
format!(
r#"{{
"model": {{
"type": "BPE",
"vocab": {{"a": 0, "Ġb": 1, "aĠb": 2, "Ġ": 3, "b": 4}},
"merges": ["Ġ b", "a Ġb"]
}},
"pre_tokenizer": {pre_tokenizer}
}}"#
)
}
#[test]
fn test_bytelevel_use_regex_false_rejected_without_supporting_split() {
let json = cross_space_merge_json(r#"{"type": "ByteLevel", "use_regex": false}"#);
let err = BpeTokenizer::from_tokenizer_json_str(&json)
.expect_err("ByteLevel(use_regex:false) without a supporting Split must fail closed");
let msg = err.to_string();
assert!(
msg.contains("unsupported BPE pre_tokenizer"),
"unexpected error message: {msg}"
);
}
#[test]
fn test_unknown_pretokenizer_type_rejected() {
let json = cross_space_merge_json(r#"{"type": "Metaspace", "replacement": "_"}"#);
let err = BpeTokenizer::from_tokenizer_json_str(&json)
.expect_err("unknown pre_tokenizer type must fail closed");
assert!(
err.to_string().contains("unsupported BPE pre_tokenizer"),
"unexpected error message: {err}"
);
}
#[test]
fn test_bare_bytelevel_use_regex_true_or_absent_accepted() {
for pre_tokenizer in [
r#"{"type": "ByteLevel"}"#,
r#"{"type": "ByteLevel", "use_regex": true}"#,
] {
let json = cross_space_merge_json(pre_tokenizer);
let tokenizer = BpeTokenizer::from_tokenizer_json_str(&json)
.unwrap_or_else(|e| panic!("bare ByteLevel must be accepted: {e}"));
assert_eq!(tokenizer.tokenize_to_ids("a b"), vec![0, 1]);
}
}
#[test]
fn test_sequence_split_then_bytelevel_use_regex_false_accepted() {
let pre_tokenizer = r#"{
"type": "Sequence",
"pretokenizers": [
{"type": "Split", "pattern": {"Regex": "\\s+"}, "behavior": "Isolated"},
{"type": "ByteLevel", "use_regex": false}
]
}"#;
let json = cross_space_merge_json(pre_tokenizer);
let tokenizer = BpeTokenizer::from_tokenizer_json_str(&json)
.unwrap_or_else(|e| panic!("Sequence[Split, ByteLevel(false)] must be accepted: {e}"));
assert_eq!(tokenizer.tokenize_to_ids("a b"), vec![0, 1]);
}
#[test]
fn test_sequence_split_then_metaspace_rejected() {
let pre_tokenizer = r#"{
"type": "Sequence",
"pretokenizers": [
{"type": "Split", "pattern": {"Regex": "\\s+"}, "behavior": "Isolated"},
{"type": "Metaspace", "replacement": "_"}
]
}"#;
let json = cross_space_merge_json(pre_tokenizer);
let err = BpeTokenizer::from_tokenizer_json_str(&json).expect_err(
"Sequence[Split, Metaspace] must fail closed even though Split is regex-based",
);
assert!(
err.to_string().contains("unsupported BPE pre_tokenizer"),
"unexpected error message: {err}"
);
}
#[test]
fn test_sequence_split_then_whitespace_rejected() {
let pre_tokenizer = r#"{
"type": "Sequence",
"pretokenizers": [
{"type": "Split", "pattern": {"Regex": "\\s+"}, "behavior": "Isolated"},
{"type": "Whitespace"}
]
}"#;
let json = cross_space_merge_json(pre_tokenizer);
let err = BpeTokenizer::from_tokenizer_json_str(&json).expect_err(
"Sequence[Split, Whitespace] must fail closed even though Split is regex-based",
);
assert!(
err.to_string().contains("unsupported BPE pre_tokenizer"),
"unexpected error message: {err}"
);
}
#[test]
fn test_bpe_merge_hello_world() {
let tokenizer = synthetic_bpe();
let ids = tokenizer.tokenize_to_ids("hello world");
assert_eq!(ids, vec![11, 16]);
}
#[test]
fn test_empty_content_special_token_does_not_hang() {
let mut vocab = HashMap::new();
for (s, i) in [("a", 0u32), ("b", 1), ("c", 2)] {
vocab.insert(s.to_string(), i);
}
let mut added = HashMap::new();
added.insert(String::new(), 0u32);
let tokenizer = BpeTokenizer::from_vocab_and_merges_with_config(
vocab,
Vec::new(),
added,
HashMap::new(),
DEFAULT_BPE_CACHE_CAPACITY,
DEFAULT_BPE_MAX_SEQ_LEN,
)
.expect("construct tokenizer with empty special token");
let (tx, rx) = std::sync::mpsc::channel();
std::thread::spawn(move || {
let _ = tx.send(tokenizer.tokenize("abc").real_length);
});
match rx.recv_timeout(std::time::Duration::from_secs(5)) {
Ok(len) => assert!(
len < 1000,
"empty special token produced runaway output ({len}); zero-length specials must be dropped"
),
Err(_) => {
panic!("tokenize() hung on a zero-length special token (infinite-loop regression)")
}
}
}
#[test]
fn test_bpe_special_token_passthrough() {
let tokenizer = synthetic_bpe();
let ids = tokenizer.tokenize_to_ids("hello<|endoftext|>world");
assert_eq!(ids, vec![11, 17, 5, 3, 6, 2, 7]);
}
#[test]
fn test_bpe_truncation_preserves_eos() {
let tokenizer = synthetic_bpe().with_add_eos().with_max_seq_len(2);
let ids = tokenizer.tokenize_to_ids("hello world");
assert_eq!(ids, vec![11, 17]);
}
#[test]
fn test_bpe_tokenize_batch_pads_to_batch_max() {
let tokenizer = synthetic_bpe();
let batch = tokenizer.tokenize_batch(&["hello", "hello world"]);
assert_eq!(batch[0].input_ids.len(), 2);
assert_eq!(batch[1].input_ids.len(), 2);
assert_eq!(batch[0].attention_mask, vec![1, 0]);
assert_eq!(batch[1].attention_mask, vec![1, 1]);
}
#[test]
fn test_bpe_decode_roundtrip() {
let tokenizer = synthetic_bpe();
let ids = tokenizer.tokenize_to_ids("hello world");
assert_eq!(ids, vec![11, 16]);
assert_eq!(tokenizer.decode(&ids), Some("hello world".to_string()));
assert_eq!(tokenizer.decode(&[]), Some(String::new()));
}
#[test]
fn test_grammar_vocab_preserves_split_utf8_bytes() {
use crate::grammar::{GrammarEngine, GrammarSpec};
let byte_encoder = bytes_to_unicode();
let mut vocab = HashMap::new();
vocab.insert(byte_encoder[0xE5].to_string(), 0);
vocab.insert(byte_encoder[0xA5].to_string(), 1);
vocab.insert(byte_encoder[0xBD].to_string(), 2);
let tokenizer = BpeTokenizer::from_vocab_and_merges(vocab, Vec::new()).unwrap();
let vocab_bytes = tokenizer.vocab_bytes(3).unwrap();
assert_eq!(vocab_bytes, vec![vec![0xE5], vec![0xA5], vec![0xBD]]);
let spec = GrammarSpec::Gbnf("root ::= \"好\"\n".to_string());
let engine = GrammarEngine::new(&spec, vocab_bytes).unwrap();
let mut state = engine.initial_state();
for token_id in 0..3u32 {
let mut logits = vec![0.0; 3];
engine
.mask_logits(&mut state, &mut logits)
.expect("matching vocab length");
assert!(
logits[token_id as usize].is_finite(),
"split UTF-8 token {token_id} must remain grammar-legal"
);
assert!(engine.advance(&mut state, token_id));
}
}
#[test]
fn test_grammar_vocab_masks_rendered_added_tokens_and_controls() {
use crate::grammar::{GrammarEngine, GrammarSpec};
let json = r#"{
"model":{"type":"BPE","vocab":{"a":0,"b":1},"merges":[]},
"added_tokens":[
{"id":99,"content":"<control>","special":true},
{"id":100,"content":"<extra>","special":false}
]
}"#;
let tokenizer = BpeTokenizer::from_tokenizer_json_str(json).unwrap();
let vocab_bytes = tokenizer.vocab_bytes(101).unwrap();
assert_eq!(vocab_bytes.len(), 101);
assert!(vocab_bytes[99].is_empty());
assert_eq!(vocab_bytes[100], b"<extra>");
let spec = GrammarSpec::Gbnf("root ::= \"a\"\n".to_string());
let engine = GrammarEngine::new(&spec, vocab_bytes).unwrap();
let mut state = engine.initial_state();
let mut logits = vec![0.0; 101];
engine
.mask_logits(&mut state, &mut logits)
.expect("matching vocab length");
assert_eq!(logits[99], f32::NEG_INFINITY);
assert_eq!(logits[100], f32::NEG_INFINITY);
}
fn non_ascii_added_tokenizer() -> BpeTokenizer {
let json = r#"{
"model":{"type":"BPE","vocab":{"a":0},"merges":[]},
"added_tokens":[
{"id":100,"content":"好","special":false},
{"id":101,"content":"café","special":false}
]
}"#;
BpeTokenizer::from_tokenizer_json_str(json).unwrap()
}
#[test]
fn grammar_vocab_fail_closes_qwen_reserved_tail() {
use crate::grammar::{GrammarEngine, GrammarSpec};
use crate::model::qwen35_config::Qwen35Config;
let json = r#"{
"model":{"type":"BPE","vocab":{"a":0},"merges":[]},
"added_tokens":[
{"id":248069,"content":"<control>","special":true}
]
}"#;
let tokenizer = BpeTokenizer::from_tokenizer_json_str(json).unwrap();
let cfg = Qwen35Config::qwen35_0_8b();
let vocab_bytes = tokenizer.vocab_bytes(cfg.vocab_size).unwrap();
let table_len = vocab_bytes.len();
let spec = GrammarSpec::Gbnf("root ::= \"a\"\n".to_string());
let engine = GrammarEngine::new(&spec, vocab_bytes).unwrap();
let mut state = engine.initial_state();
let mut logits = vec![0.0; cfg.vocab_size];
engine
.mask_logits(&mut state, &mut logits)
.expect("matching vocab length");
assert!(
logits[248_070..]
.iter()
.all(|logit| *logit == f32::NEG_INFINITY),
"reserved logits above the tokenizer maximum must be masked"
);
assert_eq!(table_len, cfg.vocab_size);
}
#[test]
fn grammar_vocab_uses_literal_non_ascii_added_token_bytes() {
use crate::grammar::{GrammarEngine, GrammarSpec};
let tokenizer = non_ascii_added_tokenizer();
let err = tokenizer.vocab_bytes(101).unwrap_err();
assert!(err.to_string().contains("token ID 101"));
let vocab_bytes = tokenizer.vocab_bytes(102).unwrap();
assert_eq!(vocab_bytes[100], "好".as_bytes());
assert_eq!(vocab_bytes[101], "café".as_bytes());
let spec = GrammarSpec::Gbnf("root ::= \"好\" | \"café\"\n".to_string());
let engine = GrammarEngine::new(&spec, vocab_bytes).unwrap();
let mut state = engine.initial_state();
let mut logits = vec![0.0; 102];
engine
.mask_logits(&mut state, &mut logits)
.expect("matching vocab length");
assert!(logits[100].is_finite());
assert!(logits[101].is_finite());
}
#[test]
fn decode_and_incremental_detokenize_render_non_ascii_added_tokens() {
use crate::tokenizer::detokenize::IncrementalDetokenizer;
let tokenizer = non_ascii_added_tokenizer();
assert_eq!(
tokenizer.token_bytes_for_id(100).as_deref(),
Some("好".as_bytes())
);
assert_eq!(
tokenizer.token_bytes_for_id(101).as_deref(),
Some("café".as_bytes())
);
assert_eq!(tokenizer.decode(&[100, 101]), Some("好café".to_string()));
let mut detok = IncrementalDetokenizer::new();
let mut text = String::new();
for id in [100, 101] {
text.push_str(&detok.push(&tokenizer, id));
}
text.push_str(&detok.finish());
assert_eq!(text, "好café");
}
#[test]
fn test_decode_renders_nonspecial_added_tokens() {
let mut vocab = HashMap::new();
for (s, i) in [("a", 0u32), ("b", 1), ("c", 2)] {
vocab.insert(s.to_string(), i);
}
let mut rendered = HashMap::new();
rendered.insert(100u32, "</think>".to_string());
rendered.insert(101u32, "<think>".to_string());
rendered.insert(102u32, "<tool_call>".to_string());
let tokenizer = BpeTokenizer::from_vocab_and_merges_with_config(
vocab,
Vec::new(),
HashMap::new(),
rendered,
DEFAULT_BPE_CACHE_CAPACITY,
DEFAULT_BPE_MAX_SEQ_LEN,
)
.expect("construct tokenizer with rendered added tokens");
assert_eq!(tokenizer.special_token_id("</think>"), None);
assert_eq!(tokenizer.token_ids_for_content("</think>"), vec![100]);
assert_eq!(tokenizer.decode(&[101]), Some("<think>".to_string()));
assert_eq!(tokenizer.decode(&[100]), Some("</think>".to_string()));
assert_eq!(tokenizer.decode(&[102]), Some("<tool_call>".to_string()));
assert_eq!(
tokenizer.decode(&[0, 1, 100, 2]),
Some("ab</think>c".to_string())
);
assert_eq!(tokenizer.decode(&[200]), Some(String::new()));
use crate::tokenizer::detokenize::IncrementalDetokenizer;
let mut detok = IncrementalDetokenizer::new();
let mut out = String::new();
for id in [0u32, 100, 1] {
out.push_str(&detok.push(&tokenizer, id));
}
out.push_str(&detok.finish());
assert_eq!(out, "a</think>b");
}
#[test]
fn token_ids_for_content_resolves_special_vocab_and_added_render_tiers() {
let mut vocab = HashMap::new();
vocab.insert("a".to_string(), 0u32);
vocab.insert("</think>".to_string(), 7u32);
let mut specials = HashMap::new();
specials.insert("<|im_end|>".to_string(), 20u32);
let mut rendered = HashMap::new();
rendered.insert(30u32, "<tool_call>".to_string());
let tokenizer = BpeTokenizer::from_vocab_and_merges_with_config(
vocab,
Vec::new(),
specials,
rendered,
DEFAULT_BPE_CACHE_CAPACITY,
DEFAULT_BPE_MAX_SEQ_LEN,
)
.expect("construct three-tier tokenizer");
assert_eq!(tokenizer.token_ids_for_content("<|im_end|>"), vec![20]);
assert_eq!(tokenizer.special_token_id("</think>"), None);
assert_eq!(tokenizer.token_ids_for_content("</think>"), vec![7]);
assert_eq!(tokenizer.special_token_id("<tool_call>"), None);
assert_eq!(tokenizer.token_ids_for_content("<tool_call>"), vec![30]);
assert!(tokenizer.token_ids_for_content("<think>").is_empty());
}
#[test]
fn token_ids_for_content_reports_every_id_sharing_one_spelling() {
let mut vocab = HashMap::new();
vocab.insert("a".to_string(), 0u32);
vocab.insert("</think>".to_string(), 7u32);
let mut rendered = HashMap::new();
rendered.insert(100u32, "</think>".to_string());
let tokenizer = BpeTokenizer::from_vocab_and_merges_with_config(
vocab,
Vec::new(),
HashMap::new(),
rendered,
DEFAULT_BPE_CACHE_CAPACITY,
DEFAULT_BPE_MAX_SEQ_LEN,
)
.expect("construct colliding-spelling tokenizer");
assert_eq!(tokenizer.token_ids_for_content("</think>"), vec![7, 100]);
assert_eq!(tokenizer.token_ids_for_content("</think>")[0], 7);
assert_eq!(tokenizer.token_ids_for_content("a"), vec![0]);
assert!(tokenizer.token_ids_for_content("<think>").is_empty());
}
#[test]
fn test_from_tokenizer_json_renders_real_qwen_think_tags() {
let json = r#"{
"model": {
"type": "BPE",
"vocab": {"a": 0, "b": 1},
"merges": []
},
"added_tokens": [
{"id": 248045, "content": "<|im_start|>", "special": true},
{"id": 248046, "content": "<|im_end|>", "special": true},
{"id": 248058, "content": "<tool_call>", "special": false},
{"id": 248068, "content": "<think>", "special": false},
{"id": 248069, "content": "</think>", "special": false}
]
}"#;
let tokenizer =
BpeTokenizer::from_tokenizer_json_str(json).expect("load tokenizer.json fixture");
assert_eq!(tokenizer.decode(&[248069]), Some("</think>".to_string()));
assert_eq!(tokenizer.decode(&[248068]), Some("<think>".to_string()));
assert_eq!(tokenizer.decode(&[248058]), Some("<tool_call>".to_string()));
assert_eq!(tokenizer.decode(&[248046]), Some(String::new()));
assert_eq!(tokenizer.decode(&[248045]), Some(String::new()));
assert_eq!(
tokenizer.decode(&[0, 248069, 1]),
Some("a</think>b".to_string())
);
}
#[test]
fn test_bpe_duplicate_merge_keeps_first_rank() {
let mut vocab = HashMap::new();
for (s, i) in [("a", 0u32), ("b", 1), ("c", 2), ("ab", 3), ("bc", 4)] {
vocab.insert(s.to_string(), i);
}
let merges = vec![
("a".to_string(), "b".to_string()),
("b".to_string(), "c".to_string()),
("a".to_string(), "b".to_string()),
];
let tokenizer = BpeTokenizer::from_vocab_and_merges(vocab, merges).unwrap();
assert_eq!(tokenizer.tokenize_to_ids("abc"), vec![3, 2]);
}
#[test]
fn test_bpe_partial_byte_fallback_rolls_back() {
let mut vocab = HashMap::new();
vocab.insert("a".to_string(), 0u32);
vocab.insert("<unk>".to_string(), 1u32);
let merges = vec![("a".to_string(), "b".to_string())];
let tokenizer = BpeTokenizer::from_vocab_and_merges(vocab, merges).unwrap();
assert_eq!(tokenizer.tokenize_to_ids("ab"), vec![1]);
}
}