use anyhow::{Result, anyhow};
use std::path::Path;
use tokenizers::{AddedToken, Tokenizer};
pub const NUM_LOC_BINS: usize = 1000;
pub struct Florence2Tokenizer {
tk: Tokenizer,
loc_base: u32,
}
impl Florence2Tokenizer {
pub fn from_file(path: &Path) -> Result<Self> {
let mut tk = Tokenizer::from_file(path).map_err(|e| anyhow!("load tokenizer: {e}"))?;
let base = tk.get_vocab_size(true) as u32;
let mut added: Vec<AddedToken> = Vec::with_capacity(1024);
for s in ["<od>", "</od>", "<ocr>", "</ocr>"] {
added.push(AddedToken::from(s, true));
}
for x in 0..NUM_LOC_BINS {
added.push(AddedToken::from(format!("<loc_{x}>"), true));
}
for s in EXTRA_TOKENS {
added.push(AddedToken::from(*s, true));
}
tk.add_special_tokens(&added);
let loc_base = base + 4;
Ok(Self { tk, loc_base })
}
pub fn encode_prompt(&self, text: &str) -> Result<Vec<u32>> {
let prompt = crate::config::construct_prompt(text);
let enc = self
.tk
.encode(prompt, true)
.map_err(|e| anyhow!("encode: {e}"))?;
Ok(enc.get_ids().to_vec())
}
pub fn decode_keep_special(&self, ids: &[u32]) -> Result<String> {
self.tk
.decode(ids, false)
.map_err(|e| anyhow!("decode: {e}"))
}
pub fn decode(&self, ids: &[u32]) -> Result<String> {
self.tk
.decode(ids, true)
.map_err(|e| anyhow!("decode: {e}"))
}
pub fn loc_index(&self, id: u32) -> Option<usize> {
if id >= self.loc_base && (id as usize) < self.loc_base as usize + NUM_LOC_BINS {
Some((id - self.loc_base) as usize)
} else {
None
}
}
}
const EXTRA_TOKENS: &[&str] = &[
"<cap>",
"</cap>",
"<ncap>",
"</ncap>",
"<dcap>",
"</dcap>",
"<grounding>",
"</grounding>",
"<seg>",
"</seg>",
"<sep>",
"<region_cap>",
"</region_cap>",
"<region_to_desciption>",
"</region_to_desciption>",
"<proposal>",
"</proposal>",
"<poly>",
"</poly>",
"<and>",
];