use rustc_hash::FxHashMap;
use serde_json::Value;
use super::super::byte_level::byte_level_decode;
use super::super::normalizer::Normalizer;
use super::super::sentencepiece::SentencePieceTokenizer;
use super::super::tokenizer::Tokenizer;
use super::super::wordpiece::WordPieceTokenizer;
use super::super::any_tokenizer::{AnyTokenizer, Backend};
use super::super::policy;
use super::components::{
find_added_token, parse_bert_norm, parse_norm_ops, parse_pre_tokenizer,
parse_special_decode_ids, parse_special_tokens, parse_unk_id,
};
use super::HfJsonError;
pub fn from_json_path<P: AsRef<std::path::Path>>(path: P) -> Result<AnyTokenizer, HfJsonError> {
let bytes = std::fs::read(path)?;
from_json_bytes(&bytes)
}
pub fn from_json_bytes(data: &[u8]) -> Result<AnyTokenizer, HfJsonError> {
let root: Value = serde_json::from_slice(data)?;
let model = root
.get("model")
.ok_or(HfJsonError::MissingField("model"))?;
let backend = match model_family(model)? {
"BPE" => build_bpe(&root, model)?,
"Unigram" => build_unigram(&root, model)?,
"WordPiece" => build_wordpiece(&root, model)?,
other => return Err(HfJsonError::UnsupportedModelType(other.to_string())),
};
let policy = policy::parse(&root)?;
let decoder = super::super::decoder::parse(root.get("decoder"));
let special_decode = parse_special_decode_ids(&root);
Ok(AnyTokenizer {
backend,
policy,
decoder,
special_decode,
})
}
fn model_family(model: &Value) -> Result<&'static str, HfJsonError> {
if let Some(t) = model.get("type").and_then(Value::as_str) {
return match t {
"BPE" => Ok("BPE"),
"Unigram" => Ok("Unigram"),
"WordPiece" => Ok("WordPiece"),
other => Err(HfJsonError::UnsupportedModelType(other.to_string())),
};
}
let nonempty_prefix = model
.get("continuing_subword_prefix")
.and_then(Value::as_str)
.is_some_and(|s| !s.is_empty());
if model.get("vocab").map(Value::is_array).unwrap_or(false) {
Ok("Unigram")
} else if model.get("merges").is_some() {
Ok("BPE")
} else if model.get("max_input_chars_per_word").is_some() || nonempty_prefix {
Ok("WordPiece")
} else {
Ok("BPE")
}
}
fn build_bpe(root: &Value, model: &Value) -> Result<Backend, HfJsonError> {
let pre = parse_pre_tokenizer(root.get("pre_tokenizer"));
let specials = parse_special_tokens(root);
let vocab = model
.get("vocab")
.and_then(Value::as_object)
.ok_or(HfJsonError::MissingField("model.vocab"))?;
let mut encoder: FxHashMap<Vec<u8>, u32> = FxHashMap::default();
encoder.reserve(vocab.len());
for (token, id) in vocab {
let id = id
.as_u64()
.ok_or(HfJsonError::MissingField("model.vocab[*] = u32"))? as u32;
match specials.get(token) {
Some(added) if added.id != id => {
return Err(HfJsonError::AddedTokenIdConflict {
content: token.clone(),
vocab_id: id,
added_id: added.id,
});
}
Some(_) => {}
None => {
if pre.byte_level && byte_level_decode(token).is_none() {
return Err(HfJsonError::InvalidByteLevel(token.clone()));
}
}
}
encoder.insert(token.as_bytes().to_vec(), id);
}
let merge_ranks = parse_merge_ranks(model, vocab);
let engine = super::super::pretokenizer::parse(root.get("pre_tokenizer"))?;
if engine.is_none()
&& !pre.anchored
&& root.get("pre_tokenizer").is_some_and(|v| !v.is_null())
&& !pre.unknown.is_empty()
{
return Err(HfJsonError::UnsupportedPreTokenizer(pre.unknown.join(", ")));
}
let is_byte_level = engine.as_ref().map_or(pre.byte_level, |pt| pt.byte_level());
let declares_byte_fallback = model
.get("byte_fallback")
.and_then(Value::as_bool)
.unwrap_or(false);
let unk_id = parse_unk_id(model, vocab, None);
let fuse_unk = model
.get("fuse_unk")
.and_then(Value::as_bool)
.unwrap_or(false);
let byte_fallback = (!is_byte_level)
.then(|| Tokenizer::byte_fallback_from_encoder(&encoder, unk_id, declares_byte_fallback))
.flatten()
.map(|bf| bf.with_fuse_unk(fuse_unk));
let tok = match engine {
Some(pt) => {
let t = if pt.byte_level() {
Tokenizer::new_byte_level(encoder, specials, super::super::tokenizer::GPT2_PATTERN)?
} else {
Tokenizer::new(encoder, specials, super::super::tokenizer::GPT2_PATTERN)?
};
let t = match merge_ranks {
Some(ranks) => t.with_merge_ranks(ranks),
None => t,
};
t.with_pre_tokenizer(pt)
}
None => {
let t = if pre.byte_level {
Tokenizer::new_byte_level(encoder, specials, &pre.pattern)?
} else if pre.metaspace {
Tokenizer::new_with_metaspace_decoder(encoder, specials, &pre.pattern)?
} else {
Tokenizer::new(encoder, specials, &pre.pattern)?
};
let t = match merge_ranks {
Some(ranks) => t.with_merge_ranks(ranks),
None => t,
};
t.with_prefix_space(pre.add_prefix_space)
}
};
let tok = tok
.with_added_token_matching(true)
.with_special_decode_ids(parse_special_decode_ids(root))
.with_normalizer(Normalizer::new(parse_norm_ops(root.get("normalizer"))?))
.with_byte_fallback(byte_fallback);
Ok(Backend::Bpe(tok))
}
fn parse_merge_ranks(
model: &Value,
vocab: &serde_json::Map<String, Value>,
) -> Option<FxHashMap<Vec<u8>, u32>> {
let merges = model.get("merges").and_then(Value::as_array)?;
let mut merged: Vec<String> = Vec::with_capacity(merges.len());
for m in merges {
match m {
Value::Array(p) if p.len() == 2 => {
if let (Some(a), Some(b)) = (p[0].as_str(), p[1].as_str()) {
merged.push(format!("{a}{b}"));
}
}
Value::String(s) => merged.push(s.replacen(' ', "", 1)),
_ => {}
}
}
if merged.is_empty() {
return None;
}
let mut base: Vec<(&String, u64)> = vocab
.iter()
.filter_map(|(k, v)| v.as_u64().map(|id| (k, id)))
.collect();
base.sort_by_key(|&(_, id)| id);
Some(super::super::bpe::merge_ranks(
&merged,
base.iter().map(|(k, _)| k.as_str()),
))
}
fn build_unigram(root: &Value, model: &Value) -> Result<Backend, HfJsonError> {
let vocab = model
.get("vocab")
.and_then(Value::as_array)
.ok_or(HfJsonError::MissingField("model.vocab"))?;
let mut tokens = Vec::with_capacity(vocab.len());
let mut scores = Vec::with_capacity(vocab.len());
for entry in vocab {
let pair = entry
.as_array()
.ok_or(HfJsonError::MissingField("model.vocab[*] = [token, score]"))?;
let token = pair
.first()
.and_then(Value::as_str)
.ok_or(HfJsonError::MissingField("model.vocab[*][0] = token"))?;
let score = pair.get(1).and_then(Value::as_f64).unwrap_or(0.0);
tokens.push(token.to_string());
scores.push(score);
}
let find = |cands: &[&str]| -> Option<u32> {
find_added_token(root, cands).or_else(|| {
cands
.iter()
.find_map(|c| tokens.iter().position(|t| t == c).map(|i| i as u32))
})
};
let eos = find(policy::EOS_CANDIDATES)
.or_else(|| {
model
.get("unk_id")
.and_then(Value::as_u64)
.map(|n| n as u32)
})
.unwrap_or(0);
let ops = parse_norm_ops(root.get("normalizer"))?;
let pre = parse_pre_tokenizer(root.get("pre_tokenizer"));
let tok = SentencePieceTokenizer::new(tokens, scores, None, eos)?
.with_normalizer(Normalizer::new(ops))
.with_prefix_space(pre.add_prefix_space)
.with_added_tokens(parse_special_tokens(root))?
.with_special_decode_ids(parse_special_decode_ids(root));
Ok(Backend::Unigram(tok))
}
fn build_wordpiece(root: &Value, model: &Value) -> Result<Backend, HfJsonError> {
let vocab = model
.get("vocab")
.and_then(Value::as_object)
.ok_or(HfJsonError::MissingField("model.vocab"))?;
let max_id = vocab
.values()
.filter_map(Value::as_u64)
.max()
.ok_or(HfJsonError::MissingField("model.vocab (empty)"))? as usize;
let mut id_to_token = vec![String::new(); max_id + 1];
for (token, id) in vocab {
let id = id
.as_u64()
.ok_or(HfJsonError::MissingField("model.vocab[*] = u32"))? as usize;
id_to_token[id] = token.clone();
}
let unk_id =
parse_unk_id(model, vocab, Some("[UNK]")).ok_or(HfJsonError::MissingSpecial("unk"))?;
let max_word_len = model
.get("max_input_chars_per_word")
.and_then(Value::as_u64)
.unwrap_or(100) as usize;
let norm = parse_bert_norm(root.get("normalizer"));
let prefix = model
.get("continuing_subword_prefix")
.and_then(Value::as_str)
.unwrap_or("##")
.to_string();
let tok = WordPieceTokenizer::with_options(
id_to_token,
unk_id,
max_word_len,
norm.lowercase,
norm.handle_chinese_chars,
norm.clean_text,
prefix,
)
.with_strip_accents(norm.strip_accents)
.with_added_tokens(parse_special_tokens(root))?
.with_special_decode_ids(parse_special_decode_ids(root));
Ok(Backend::WordPiece(tok))
}