use rustc_hash::FxHashMap;
use serde::de::{MapAccess, 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::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;
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 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 mut encoder: FxHashMap<Vec<u8>, u32> = FxHashMap::default();
encoder.reserve(vocab.len());
let mut raw_encoder: FxHashMap<Vec<u8>, u32> = FxHashMap::default();
if pre.byte_level {
raw_encoder.reserve(vocab.len());
}
for (token, id) in vocab {
let (token, id) = (token.as_ref(), *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) => {
raw_encoder.insert(raw, id);
}
}
}
}
}
encoder.insert(token.as_bytes().to_vec(), id);
}
let merge_ranks = parse_merge_ranks(merges, 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,
|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_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,
};
let t = t.with_raw_encoder(raw_encoder);
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(
merges: Option<&RawValue>,
vocab: &[(Cow<'_, str>, u32)],
) -> Option<FxHashMap<Vec<u8>, u32>> {
let merges: Vec<Value> = serde_json::from_str(merges?.get()).ok()?;
let merges = &merges;
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<(&str, u32)> = vocab.iter().map(|(k, id)| (k.as_ref(), *id)).collect();
base.sort_by_key(|&(_, id)| id);
Some(super::super::bpe::merge_ranks(
&merged,
base.iter().map(|(k, _)| *k),
))
}
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,
|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))
}