use rustc_hash::{FxHashMap, FxHashSet};
use serde::de::{MapAccess, SeqAccess, Visitor};
use serde::{Deserialize, Deserializer};
use serde_json::value::RawValue;
use serde_json::Value;
use std::borrow::Cow;
use std::fmt;
use std::marker::PhantomData;
use super::super::byte_level::byte_level_decode;
use super::super::decode_table::Decoder;
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::encoder::Encoder;
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;
struct CowStr<'a>(Cow<'a, str>);
impl<'de: 'a, 'a> Deserialize<'de> for CowStr<'a> {
fn deserialize<D: Deserializer<'de>>(de: D) -> Result<Self, D::Error> {
struct V<'a>(PhantomData<&'a ()>);
impl<'de: 'a, 'a> Visitor<'de> for V<'a> {
type Value = CowStr<'a>;
fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("a string")
}
fn visit_borrowed_str<E>(self, v: &'de str) -> Result<Self::Value, E> {
Ok(CowStr(Cow::Borrowed(v)))
}
fn visit_str<E>(self, v: &str) -> Result<Self::Value, E> {
Ok(CowStr(Cow::Owned(v.to_owned())))
}
fn visit_string<E>(self, v: String) -> Result<Self::Value, E> {
Ok(CowStr(Cow::Owned(v)))
}
}
de.deserialize_str(V(PhantomData))
}
}
struct MergeList(Vec<(String, usize)>);
impl<'de> Deserialize<'de> for MergeList {
fn deserialize<D: Deserializer<'de>>(de: D) -> Result<Self, D::Error> {
struct Merged(Option<(String, usize)>);
impl<'de> Deserialize<'de> for Merged {
fn deserialize<D: Deserializer<'de>>(de: D) -> Result<Self, D::Error> {
struct V;
impl<'de> Visitor<'de> for V {
type Value = Merged;
fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("a merge pair or its joined string")
}
fn visit_seq<S: SeqAccess<'de>>(self, mut seq: S) -> Result<Merged, S::Error> {
let (a, b) = (seq.next_element::<CowStr>()?, seq.next_element::<CowStr>()?);
let extra = seq.next_element::<serde::de::IgnoredAny>()?.is_some();
Ok(Merged(match (a, b, extra) {
(Some(a), Some(b), false) if !a.0.is_empty() && !b.0.is_empty() => {
let mut s = String::with_capacity(a.0.len() + b.0.len());
s.push_str(&a.0);
s.push_str(&b.0);
let split = a.0.len();
Some((s, split))
}
_ => None,
}))
}
fn visit_str<E>(self, v: &str) -> Result<Merged, E> {
let Some(split) = v.find(' ') else {
return Ok(Merged(None));
};
if split == 0 || split + 1 == v.len() {
return Ok(Merged(None));
}
Ok(Merged(Some((v.replacen(' ', "", 1), split))))
}
}
de.deserialize_any(V)
}
}
struct V;
impl<'de> Visitor<'de> for V {
type Value = MergeList;
fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("a merges list")
}
fn visit_seq<S: SeqAccess<'de>>(self, mut seq: S) -> Result<MergeList, S::Error> {
let mut out = Vec::with_capacity(seq.size_hint().unwrap_or(0));
while let Some(Merged(m)) = seq.next_element::<Merged>()? {
out.extend(m);
}
Ok(MergeList(out))
}
}
de.deserialize_seq(V)
}
}
struct VocabPairs<'a>(Vec<(Cow<'a, str>, u32)>);
impl<'de: 'a, 'a> Deserialize<'de> for VocabPairs<'a> {
fn deserialize<D: Deserializer<'de>>(de: D) -> Result<Self, D::Error> {
struct V<'a>(PhantomData<&'a ()>);
impl<'de: 'a, 'a> Visitor<'de> for V<'a> {
type Value = VocabPairs<'a>;
fn expecting(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("a vocabulary object")
}
fn visit_map<M: MapAccess<'de>>(self, mut map: M) -> Result<Self::Value, M::Error> {
let mut out = Vec::with_capacity(map.size_hint().unwrap_or(0));
while let Some((token, id)) = map.next_entry::<CowStr<'a>, u32>()? {
out.push((token.0, id));
}
Ok(VocabPairs(out))
}
}
de.deserialize_map(V(PhantomData))
}
}
fn expand(raw: &RawValue) -> Result<Value, HfJsonError> {
Ok(serde_json::from_str(raw.get())?)
}
fn object_from(spans: &FxHashMap<&str, &RawValue>) -> Result<Value, HfJsonError> {
let mut map = serde_json::Map::with_capacity(spans.len());
for (key, raw) in spans {
map.insert((*key).to_string(), expand(raw)?);
}
Ok(Value::Object(map))
}
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 top: FxHashMap<&str, &RawValue> = serde_json::from_slice(data)?;
let model_raw = *top.get("model").ok_or(HfJsonError::MissingField("model"))?;
let mut model_spans: FxHashMap<&str, &RawValue> = serde_json::from_str(model_raw.get())?;
let vocab_raw = model_spans.remove("vocab");
let merges_raw = model_spans.remove("merges");
let mut root_spans = top;
root_spans.remove("model");
let root = object_from(&root_spans)?;
let model = object_from(&model_spans)?;
let backend = match model_family(&model, vocab_raw, merges_raw.is_some())? {
"BPE" => {
let raw = vocab_raw.ok_or(HfJsonError::MissingField("model.vocab"))?;
let vocab: VocabPairs<'_> = serde_json::from_str(raw.get())?;
build_bpe(&root, &model, &vocab.0, merges_raw)?
}
family => {
let mut model = model;
if let Value::Object(map) = &mut model {
if let Some(raw) = vocab_raw {
map.insert("vocab".to_string(), expand(raw)?);
}
if let Some(raw) = merges_raw {
map.insert("merges".to_string(), expand(raw)?);
}
}
match family {
"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,
vocab: Option<&RawValue>,
has_merges: bool,
) -> 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());
let vocab_is_array = vocab.is_some_and(|v| v.get().trim_start().starts_with('['));
if vocab_is_array {
Ok("Unigram")
} else if has_merges {
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,
vocab: &[(Cow<'_, str>, u32)],
merges: Option<&RawValue>,
) -> Result<Backend, HfJsonError> {
let pre = parse_pre_tokenizer(root.get("pre_tokenizer"));
let specials = parse_special_tokens(root);
let merge_rules = parse_merge_rules(merges);
let ignore_merges = model
.get("ignore_merges")
.and_then(Value::as_bool)
.unwrap_or(false);
if let Some(prefix) = model
.get("continuing_subword_prefix")
.and_then(Value::as_str)
.filter(|p| !p.is_empty())
{
return Err(HfJsonError::UnsupportedModelField(format!(
"model.continuing_subword_prefix = {prefix:?}"
)));
}
let suffix = model
.get("end_of_word_suffix")
.and_then(Value::as_str)
.filter(|s| !s.is_empty());
let merge_ranks = match &merge_rules {
Some(rules) => parse_merge_ranks(rules, vocab),
None => None,
};
let unreachable = match &merge_rules {
Some(rules) if !ignore_merges => unreachable_tokens(
rules,
vocab,
suffix,
merge_ranks.as_ref(),
!pre.byte_level,
),
_ => FxHashSet::default(),
};
let mut encoder: Encoder = Encoder::default();
encoder.reserve(vocab.len() - unreachable.len().min(vocab.len()));
let mut raw_encoder: Encoder = Encoder::default();
if pre.byte_level {
raw_encoder.reserve(vocab.len());
}
let mut fallback_ids: [Option<u32>; 256] = [None; 256];
for (token, id) in vocab {
let (token, id) = (token.as_ref(), *id);
if let Some(byte) = byte_fallback_byte(token.as_bytes()) {
fallback_ids[byte as usize].get_or_insert(id);
}
match specials.get(token) {
Some(added) if added.id != id => {
return Err(HfJsonError::AddedTokenIdConflict {
content: token.to_string(),
vocab_id: id,
added_id: added.id,
});
}
Some(_) => {}
None => {
if pre.byte_level {
match byte_level_decode(token) {
None => return Err(HfJsonError::InvalidByteLevel(token.to_string())),
Some(raw) => {
if !unreachable.contains(token) {
raw_encoder.insert(&raw, id);
}
}
}
}
}
}
if unreachable.contains(token) {
continue;
}
encoder.insert(token.as_bytes(), id);
}
let full_decoder = (!unreachable.is_empty()).then(|| {
let mut decoder = Decoder::with_capacity(vocab.len());
for (token, id) in vocab {
decoder.insert(*id, token.as_bytes());
}
decoder
});
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,
|name| {
vocab
.iter()
.find(|(token, _)| token.as_ref() == name)
.map(|(_, id)| *id)
},
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(
|spelling| byte_fallback_byte(spelling).and_then(|b| fallback_ids[b as usize]),
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::from_encoder(
encoder,
specials,
super::super::tokenizer::GPT2_PATTERN,
true,
)?
} else {
Tokenizer::from_encoder(
encoder,
specials,
super::super::tokenizer::GPT2_PATTERN,
false,
)?
};
let t = match merge_ranks {
Some(ranks) => t.with_merge_ranks(ranks),
None => t,
};
let t = t.with_raw_encoder(raw_encoder);
t.with_pre_tokenizer(pt)
}
None => {
let t = if pre.byte_level {
Tokenizer::from_encoder(encoder, specials, &pre.pattern, true)?
} else if pre.metaspace {
Tokenizer::from_encoder_with_metaspace_decoder(encoder, specials, &pre.pattern)?
} else {
Tokenizer::from_encoder(encoder, specials, &pre.pattern, false)?
};
let t = match merge_ranks {
Some(ranks) => t.with_merge_ranks(ranks),
None => t,
};
t.with_prefix_space(pre.add_prefix_space)
.with_metaspace_split(pre.metaspace_split)
}
};
let tok = match full_decoder {
Some(decoder) => tok.with_decode_table(decoder),
None => tok,
};
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)
.with_end_of_word_suffix(suffix);
Ok(Backend::Bpe(tok))
}
fn parse_merge_ranks(rules: &[(String, usize)], vocab: &[(Cow<'_, str>, u32)]) -> Option<Encoder> {
if rules.is_empty() {
return None;
}
let mut base: Vec<(&str, u32)> = vocab.iter().map(|(k, id)| (k.as_ref(), *id)).collect();
base.sort_by_key(|&(_, id)| id);
Some(super::super::bpe::merge_ranks(
rules.iter().map(|(merged, _)| merged.clone()).collect(),
base.iter().map(|(k, _)| *k),
))
}
fn parse_merge_rules(merges: Option<&RawValue>) -> Option<Vec<(String, usize)>> {
let MergeList(rules) = serde_json::from_str(merges?.get()).ok()?;
(!rules.is_empty()).then_some(rules)
}
fn byte_fallback_byte(token: &[u8]) -> Option<u8> {
let [b'<', b'0', b'x', hi, lo, b'>'] = token else {
return None;
};
let digit = |b: &u8| {
(*b as char)
.to_digit(16)
.filter(|_| !b.is_ascii_lowercase())
};
Some((digit(hi)? * 16 + digit(lo)?) as u8)
}
#[cfg(feature = "rayon")]
const MIN_PARALLEL_VOCAB: usize = 16_384;
fn unreachable_tokens<'v>(
rules: &[(String, usize)],
vocab: &'v [(Cow<'v, str>, u32)],
suffix: Option<&str>,
ranks: Option<&Encoder>,
char_granular: bool,
) -> FxHashSet<&'v str> {
let ranks = suffix.is_none().then_some(ranks).flatten();
let mut results: FxHashMap<&str, usize> = FxHashMap::default();
let mut named: FxHashSet<&str> = FxHashSet::default();
match ranks.is_some() {
true => {
results.reserve(rules.len());
for (merged, split) in rules {
results.insert(merged.as_str(), *split);
}
}
false => {
named.reserve(rules.len() * 2);
for (merged, split) in rules {
named.insert(merged.as_str());
named.insert(&merged[..*split]);
named.insert(&merged[*split..]);
}
}
}
let is_unreachable = |token: &&'v str| -> bool {
if is_seed_spelling(token, suffix) {
return false;
}
if super::super::vocab::is_byte_fallback_piece(token.as_bytes()) {
return ranks.is_some();
}
let Some(ranks) = ranks else {
return !named.contains(*token);
};
let Some(&split) = results.get(*token) else {
return true;
};
let symbols = |s: &str| match char_granular {
true => s.chars().count(),
false => s.len(),
};
if symbols(&token[..split]) == 1 && symbols(&token[split..]) == 1 {
return false;
}
!super::super::bpe::merges_to_whole(
token.as_bytes(),
super::super::bpe::RankLookup::new(ranks),
char_granular,
)
};
#[cfg(feature = "rayon")]
if ranks.is_some() && vocab.len() >= MIN_PARALLEL_VOCAB {
use rayon::prelude::*;
return vocab
.par_iter()
.map(|(token, _)| token.as_ref())
.filter(is_unreachable)
.collect::<Vec<&'v str>>()
.into_iter()
.collect();
}
vocab
.iter()
.map(|(token, _)| token.as_ref())
.filter(is_unreachable)
.collect()
}
fn is_seed_spelling(token: &str, suffix: Option<&str>) -> bool {
let bare = match suffix {
Some(suffix) => token.strip_suffix(suffix).unwrap_or(token),
None => token,
};
let mut chars = bare.chars();
chars.next().is_some() && chars.next().is_none()
}
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));
let tok = match super::super::pretokenizer::parse(root.get("pre_tokenizer"))? {
Some(pt) if !pt.byte_level() => tok.with_word_split(pt),
_ => tok,
};
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,
|name| vocab.get(name).and_then(Value::as_u64).map(|id| id as u32),
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))
}