#[derive(Debug, Clone)]
pub struct SpecialTokens {
pub eos: u32,
pub eos_token_ids: Vec<u32>,
pub image_start: u32,
pub image_end: u32,
pub image_token: u32,
}
#[cfg(not(target_arch = "wasm32"))]
mod imp {
use tokenizers::Tokenizer;
use super::SpecialTokens;
use crate::CandleOcrError;
use crate::error::Result;
pub fn resolve_special_tokens(tokenizer: &Tokenizer) -> Result<SpecialTokens> {
let eos = tokenizer
.token_to_id("<|endoftext|>")
.ok_or_else(|| CandleOcrError::Tokenizer("<|endoftext|> not in vocab".to_string()))?;
let image_start = tokenizer
.token_to_id("<|begin_of_image|>")
.ok_or_else(|| CandleOcrError::Tokenizer("<|begin_of_image|> not in vocab".to_string()))?;
let image_end = tokenizer
.token_to_id("<|end_of_image|>")
.ok_or_else(|| CandleOcrError::Tokenizer("<|end_of_image|> not in vocab".to_string()))?;
let image_token = tokenizer
.token_to_id("<|image|>")
.ok_or_else(|| CandleOcrError::Tokenizer("<|image|> not in vocab".to_string()))?;
let mut eos_token_ids = vec![eos];
if let Some(user_token) = tokenizer.token_to_id("<|user|>") {
if !eos_token_ids.contains(&user_token) {
eos_token_ids.push(user_token);
}
} else {
let fallback_user_token: u32 = 59253;
if !eos_token_ids.contains(&fallback_user_token) {
eos_token_ids.push(fallback_user_token);
}
}
Ok(SpecialTokens {
eos,
eos_token_ids,
image_start,
image_end,
image_token,
})
}
pub fn build_input_ids(
special: &SpecialTokens,
tokenizer: &Tokenizer,
task_prompt: &str,
num_image_tokens: usize,
) -> Result<(Vec<u32>, usize)> {
let prompt_string = format!(
"[gMASK]<sop><|user|>\n<|begin_of_image|>{placeholders}<|end_of_image|>{task}<|assistant|>\n",
placeholders = "<|image|>".repeat(num_image_tokens),
task = task_prompt,
);
let encoding = tokenizer
.encode(prompt_string, false)
.map_err(|e| CandleOcrError::Tokenizer(format!("Encode prompt: {}", e)))?;
let ids: Vec<u32> = encoding.get_ids().to_vec();
let image_tokens_start = ids
.iter()
.position(|&id| id == special.image_token)
.ok_or_else(|| CandleOcrError::Tokenizer("No <|image|> placeholder in encoded prompt".to_string()))?;
let observed = ids[image_tokens_start..]
.iter()
.take_while(|&&id| id == special.image_token)
.count();
if observed != num_image_tokens {
return Err(CandleOcrError::Tokenizer(format!(
"Expected {} <|image|> placeholders, encoded prompt has {}",
num_image_tokens, observed
)));
}
Ok((ids, image_tokens_start))
}
pub fn decode_output(tokenizer: &Tokenizer, ids: &[u32]) -> Result<String> {
let text = tokenizer
.decode(ids, false)
.map_err(|e| CandleOcrError::Tokenizer(format!("Decode error: {}", e)))?;
let cleaned = text
.replace("<|endoftext|>", "")
.replace("<|user|>", "")
.replace("<|assistant|>", "")
.replace("<|begin_of_image|>", "")
.replace("<|end_of_image|>", "")
.replace("<|image|>", "")
.replace("[gMASK]", "")
.replace("<sop>", "")
.trim()
.to_string();
Ok(cleaned)
}
}
#[cfg(not(target_arch = "wasm32"))]
pub use imp::{build_input_ids, decode_output, resolve_special_tokens};