use crate::error::InferenceError;
use crate::tokenizer::bpe::parse_merges_json;
use crate::tokenizer::common::{
JsonValue, TokenizedInput, Tokenizer, invert_vocab, json_object_to_vocab, json_path,
known_special_id, pad_ids, parse_added_tokens, parse_json, parse_rendered_added_tokens,
};
use std::cmp::Ordering;
use std::collections::{BinaryHeap, HashMap, HashSet};
use std::fs;
use std::path::Path;
use std::sync::Arc;
const METASPACE: char = '\u{2581}';
const DEFAULT_MAX_SEQ_LEN: usize = 4_096;
#[derive(Debug, Clone)]
pub struct GemmaBpeTokenizer {
inner: Arc<GemmaBpeInner>,
}
#[derive(Debug, Clone)]
struct GemmaBpeInner {
vocab: HashMap<String, u32>,
id_to_token: Vec<String>,
merges: HashMap<String, HashMap<String, usize>>,
special_tokens: HashMap<String, u32>,
special_tokens_sorted: Vec<String>,
added_render: HashMap<u32, String>,
special_skip_ids: HashSet<u32>,
pad_id: u32,
unk_id: Option<u32>,
max_seq_len: usize,
}
#[derive(Debug, Clone)]
struct MergeNode {
token: String,
prev: Option<usize>,
next: Option<usize>,
alive: bool,
version: u64,
}
#[derive(Debug, Clone, Eq, PartialEq)]
struct MergeCandidate {
rank: usize,
left: usize,
right: usize,
left_version: u64,
right_version: u64,
}
impl Ord for MergeCandidate {
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 MergeCandidate {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
impl GemmaBpeTokenizer {
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)?;
validate_gemma_bpe_shape(&root)?;
let vocab =
json_object_to_vocab(json_path(&root, &["model", "vocab"]).ok_or_else(|| {
InferenceError::Tokenizer("tokenizer.json missing model.vocab".into())
})?)?;
let id_to_token = invert_vocab(&vocab)?;
let merges_value = json_path(&root, &["model", "merges"]).ok_or_else(|| {
InferenceError::Tokenizer("tokenizer.json missing model.merges".into())
})?;
let merges_list = parse_merges_json(merges_value)?;
let mut merge_ranks: HashMap<String, HashMap<String, usize>> = HashMap::new();
for (rank, (left, right)) in merges_list.into_iter().enumerate() {
merge_ranks
.entry(left)
.or_default()
.entry(right)
.or_insert(rank);
}
let mut added = parse_added_tokens(&root);
added.retain(|name, _| !name.is_empty());
let added_render: HashMap<u32, String> = parse_rendered_added_tokens(&root);
let rendered_ids: HashSet<u32> = added_render.keys().copied().collect();
let special_skip_ids: HashSet<u32> = added
.values()
.copied()
.filter(|id| !rendered_ids.contains(id))
.collect();
let mut special_tokens = added;
for name in ["<pad>", "<eos>", "<bos>", "<unk>", "<mask>"] {
if let Some(&id) = vocab.get(name) {
special_tokens.entry(name.to_string()).or_insert(id);
}
}
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 pad_id = known_special_id(&vocab, &["<pad>"]).unwrap_or(0);
let unk_id = known_special_id(&vocab, &["<unk>"]);
Ok(Self {
inner: Arc::new(GemmaBpeInner {
vocab,
id_to_token,
merges: merge_ranks,
special_tokens,
special_tokens_sorted,
added_render,
special_skip_ids,
pad_id,
unk_id,
max_seq_len: DEFAULT_MAX_SEQ_LEN,
}),
})
}
pub fn with_max_seq_len(self, max_seq_len: usize) -> Self {
let mut inner = (*self.inner).clone();
inner.max_seq_len = max_seq_len;
Self {
inner: Arc::new(inner),
}
}
fn tokenize_to_ids(&self, text: &str) -> Vec<u32> {
let mut ids = Vec::new();
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.encode_segment(&text[segment_start..pos], &mut ids);
}
ids.push(special_id);
pos = special_end;
segment_start = pos;
continue;
}
let ch = text[pos..]
.chars()
.next()
.expect("invariant: pos is inside non-empty UTF-8 text");
pos += ch.len_utf8();
}
if segment_start < text.len() {
self.encode_segment(&text[segment_start..], &mut ids);
}
if ids.len() > self.inner.max_seq_len {
ids.truncate(self.inner.max_seq_len);
}
ids
}
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.as_str())
&& let Some(&id) = self.inner.special_tokens.get(token)
{
return Some((pos + token.len(), id));
}
}
None
}
fn encode_segment(&self, text: &str, out: &mut Vec<u32>) {
if text.is_empty() {
return;
}
let mut symbols: Vec<String> = Vec::with_capacity(text.len());
for ch in text.chars() {
let mapped = if ch == ' ' { METASPACE } else { ch };
let single = mapped.to_string();
if self.inner.vocab.contains_key(single.as_str()) {
symbols.push(single);
continue;
}
let mut buf = [0u8; 4];
for &byte in mapped.encode_utf8(&mut buf).as_bytes() {
symbols.push(byte_fallback_token(byte));
}
}
for token in self.bpe_merge(symbols) {
if let Some(&id) = self.inner.vocab.get(token.as_str()) {
out.push(id);
} else if let Some(unk_id) = self.inner.unk_id {
out.push(unk_id);
}
}
}
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: &[MergeNode],
heap: &mut BinaryHeap<MergeCandidate>,
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(MergeCandidate {
rank,
left,
right,
left_version: nodes[left].version,
right_version: nodes[right].version,
});
}
fn bpe_merge(&self, symbols: Vec<String>) -> Vec<String> {
let mut nodes: Vec<MergeNode> = symbols
.into_iter()
.enumerate()
.map(|(idx, token)| MergeNode {
token,
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
}
fn append_token(&self, id: u32, pending_fallback: &mut Vec<u8>, out: &mut String) {
if self.inner.special_skip_ids.contains(&id) {
return;
}
if let Some(content) = self.inner.added_render.get(&id) {
flush_byte_fallback_run(pending_fallback, out);
out.push_str(content);
return;
}
let Some(tok) = self.inner.id_to_token.get(id as usize) else {
return;
};
if tok.is_empty() {
return;
}
if let Some(byte) = parse_byte_fallback_token(tok) {
pending_fallback.push(byte);
return;
}
flush_byte_fallback_run(pending_fallback, out);
for ch in tok.chars() {
if ch == METASPACE {
out.push(' ');
} else {
out.push(ch);
}
}
}
}
fn flush_byte_fallback_run(pending: &mut Vec<u8>, out: &mut String) {
if pending.is_empty() {
return;
}
match String::from_utf8(std::mem::take(pending)) {
Ok(valid) => out.push_str(&valid),
Err(err) => {
for _ in 0..err.into_bytes().len() {
out.push('\u{FFFD}');
}
}
}
}
impl Tokenizer for GemmaBpeTokenizer {
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 max_len = 0usize;
let mut all = Vec::with_capacity(texts.len());
for text in texts {
let ids = self.tokenize_to_ids(text);
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 mut out = String::new();
let mut pending_fallback: Vec<u8> = Vec::new();
for &id in ids {
self.append_token(id, &mut pending_fallback, &mut out);
}
flush_byte_fallback_run(&mut pending_fallback, &mut out);
Some(out)
}
fn vocab_size(&self) -> usize {
self.inner.id_to_token.len()
}
fn max_seq_len(&self) -> usize {
self.inner.max_seq_len
}
}
fn byte_fallback_token(byte: u8) -> String {
format!("<0x{byte:02X}>")
}
fn parse_byte_fallback_token(tok: &str) -> Option<u8> {
let hex = tok.strip_prefix("<0x")?.strip_suffix('>')?;
if hex.len() != 2 {
return None;
}
u8::from_str_radix(hex, 16).ok()
}
fn validate_gemma_bpe_shape(root: &JsonValue) -> Result<(), InferenceError> {
let model_type = json_path(root, &["model", "type"])
.and_then(JsonValue::as_str)
.unwrap_or("");
if model_type != "BPE" {
return Err(InferenceError::Tokenizer(format!(
"GemmaBpeTokenizer expects tokenizer.json model.type == \"BPE\", found {model_type:?}"
)));
}
let byte_fallback = json_path(root, &["model", "byte_fallback"])
.and_then(JsonValue::as_bool)
.unwrap_or(false);
if !byte_fallback {
return Err(InferenceError::Tokenizer(
"GemmaBpeTokenizer expects tokenizer.json model.byte_fallback == true".into(),
));
}
let normalizer = root.get("normalizer");
let is_gemma_normalizer = normalizer.is_some_and(|n| {
n.get("type").and_then(JsonValue::as_str) == Some("Replace")
&& json_path(n, &["pattern", "String"]).and_then(JsonValue::as_str) == Some(" ")
&& n.get("content").and_then(JsonValue::as_str) == Some("\u{2581}")
});
if !is_gemma_normalizer {
return Err(InferenceError::Tokenizer(
"GemmaBpeTokenizer expects a Replace(\" \" -> \"\u{2581}\") normalizer; \
tokenizer.json declares a different shape, refusing to guess"
.into(),
));
}
let pre_tokenizer = root.get("pre_tokenizer");
let is_gemma_pretokenizer = pre_tokenizer.is_some_and(|pt| {
pt.get("type").and_then(JsonValue::as_str) == Some("Split")
&& json_path(pt, &["pattern", "String"]).and_then(JsonValue::as_str) == Some(" ")
&& pt.get("behavior").and_then(JsonValue::as_str) == Some("MergedWithPrevious")
&& pt.get("invert").and_then(JsonValue::as_bool) == Some(false)
});
if !is_gemma_pretokenizer {
return Err(InferenceError::Tokenizer(
"GemmaBpeTokenizer expects a literal-string Split pre_tokenizer \
(pattern.String == \" \", behavior == \"MergedWithPrevious\", \
invert == false); tokenizer.json declares a different shape, \
refusing to guess"
.into(),
));
}
let decoders = root
.get("decoder")
.filter(|d| d.get("type").and_then(JsonValue::as_str) == Some("Sequence"))
.and_then(|d| d.get("decoders"))
.and_then(JsonValue::as_array);
let Some(decoders) = decoders else {
return Err(InferenceError::Tokenizer(
"GemmaBpeTokenizer expects decoder == Sequence[Replace(\"\u{2581}\" -> \
\" \"), ByteFallback, Fuse]; tokenizer.json declares a different \
shape, refusing to guess"
.into(),
));
};
validate_gemma_decoder_sequence(decoders)?;
let model = root.get("model").unwrap_or(&JsonValue::Null);
if !is_null_or_absent(model.get("dropout")) {
return Err(InferenceError::Tokenizer(
"GemmaBpeTokenizer expects model.dropout == null (this loader \
always applies BPE merges); tokenizer.json declares a different \
shape, refusing to guess"
.into(),
));
}
if !is_null_or_absent(model.get("continuing_subword_prefix")) {
return Err(InferenceError::Tokenizer(
"GemmaBpeTokenizer expects model.continuing_subword_prefix == null \
(this loader never affix-wraps a subword); tokenizer.json declares \
a different shape, refusing to guess"
.into(),
));
}
if !is_null_or_absent(model.get("end_of_word_suffix")) {
return Err(InferenceError::Tokenizer(
"GemmaBpeTokenizer expects model.end_of_word_suffix == null (this \
loader never affix-wraps a subword); tokenizer.json declares a \
different shape, refusing to guess"
.into(),
));
}
if model.get("fuse_unk").and_then(JsonValue::as_bool) != Some(true) {
return Err(InferenceError::Tokenizer(
"GemmaBpeTokenizer expects model.fuse_unk == true; tokenizer.json \
declares a different shape, refusing to guess"
.into(),
));
}
if model.get("ignore_merges").and_then(JsonValue::as_bool) != Some(false) {
return Err(InferenceError::Tokenizer(
"GemmaBpeTokenizer expects model.ignore_merges == false (this \
loader always applies BPE merges); tokenizer.json declares a \
different shape, refusing to guess"
.into(),
));
}
Ok(())
}
fn is_null_or_absent(value: Option<&JsonValue>) -> bool {
matches!(value, None | Some(JsonValue::Null))
}
fn validate_gemma_decoder_sequence(decoders: &[JsonValue]) -> Result<(), InferenceError> {
let [replace, byte_fallback, fuse] = decoders else {
return Err(InferenceError::Tokenizer(
"GemmaBpeTokenizer expects exactly three decoder stages \
[Replace, ByteFallback, Fuse]; tokenizer.json declares a \
different shape, refusing to guess"
.into(),
));
};
if replace.get("type").and_then(JsonValue::as_str) != Some("Replace") {
return Err(InferenceError::Tokenizer(
"GemmaBpeTokenizer expects decoder[0].type == \"Replace\"; \
tokenizer.json declares a different shape, refusing to guess"
.into(),
));
}
if json_path(replace, &["pattern", "String"]).and_then(JsonValue::as_str) != Some("\u{2581}") {
return Err(InferenceError::Tokenizer(
"GemmaBpeTokenizer expects decoder[0].pattern.String == \
\"\u{2581}\"; tokenizer.json declares a different shape, \
refusing to guess"
.into(),
));
}
if replace.get("content").and_then(JsonValue::as_str) != Some(" ") {
return Err(InferenceError::Tokenizer(
"GemmaBpeTokenizer expects decoder[0].content == \" \"; \
tokenizer.json declares a different shape, refusing to guess"
.into(),
));
}
if byte_fallback.get("type").and_then(JsonValue::as_str) != Some("ByteFallback") {
return Err(InferenceError::Tokenizer(
"GemmaBpeTokenizer expects decoder[1].type == \"ByteFallback\"; \
tokenizer.json declares a different shape, refusing to guess"
.into(),
));
}
if fuse.get("type").and_then(JsonValue::as_str) != Some("Fuse") {
return Err(InferenceError::Tokenizer(
"GemmaBpeTokenizer expects decoder[2].type == \"Fuse\"; \
tokenizer.json declares a different shape, refusing to guess"
.into(),
));
}
Ok(())
}
pub const GEMMA4_IMAGE_SOFT_TOKENS_PER_IMAGE: u32 = 280;
pub const GEMMA4_AUDIO_MS_PER_SOFT_TOKEN: u32 = 40;
pub const GEMMA4_AUDIO_MAX_SOFT_TOKENS: u32 = 750;
pub const GEMMA4_AUDIO_SAMPLING_RATE_HZ: u32 = 16_000;
pub const GEMMA4_AUDIO_FRAME_LENGTH_SAMPLES: u32 = 320;
pub const GEMMA4_AUDIO_HOP_LENGTH_SAMPLES: u32 = 160;
pub fn image_marker_expansion_tokens(image_marker_count: usize) -> usize {
image_marker_count * GEMMA4_IMAGE_SOFT_TOKENS_PER_IMAGE as usize
}
pub fn audio_marker_expansion_tokens_from_samples(num_samples: u32) -> u32 {
let frame_size_for_unfold = i64::from(GEMMA4_AUDIO_FRAME_LENGTH_SAMPLES) + 1;
let pad_left = i64::from(GEMMA4_AUDIO_FRAME_LENGTH_SAMPLES / 2);
let hop = i64::from(GEMMA4_AUDIO_HOP_LENGTH_SAMPLES);
let mel_frame_numerator = i64::from(num_samples) + pad_left - frame_size_for_unfold;
let num_mel_frames = mel_frame_numerator.div_euclid(hop) + 1;
if num_mel_frames <= 0 {
return 0;
}
const SSCP_KERNEL: i64 = 3;
const SSCP_STRIDE: i64 = 2;
const SSCP_PADDING: i64 = 1;
const SSCP_LAYERS: u32 = 2;
let mut t = num_mel_frames;
for _ in 0..SSCP_LAYERS {
t = (t + 2 * SSCP_PADDING - SSCP_KERNEL).div_euclid(SSCP_STRIDE) + 1;
}
u32::try_from(t)
.unwrap_or(0)
.min(GEMMA4_AUDIO_MAX_SOFT_TOKENS)
}
pub fn audio_marker_expansion_tokens(duration_ms: u32) -> u32 {
let num_samples = u64::from(duration_ms) * u64::from(GEMMA4_AUDIO_SAMPLING_RATE_HZ) / 1000;
audio_marker_expansion_tokens_from_samples(u32::try_from(num_samples).unwrap_or(u32::MAX))
}
pub fn total_audio_marker_expansion_tokens(durations_ms: &[u32]) -> usize {
durations_ms
.iter()
.copied()
.map(|ms| audio_marker_expansion_tokens(ms) as usize)
.sum()
}