use std::collections::{BTreeSet, HashMap};
use std::path::Path;
use crate::error::{ForgeError, Result};
use crate::tokenizer::Tokenizer;
pub struct CharTokenizer {
itos: Vec<char>,
stoi: HashMap<char, u32>,
}
impl CharTokenizer {
pub fn from_corpus(text: &str) -> Self {
Self::from_chars(text.chars().collect::<BTreeSet<char>>())
}
fn from_chars(sorted: BTreeSet<char>) -> Self {
let itos: Vec<char> = sorted.into_iter().collect();
let stoi = itos
.iter()
.enumerate()
.map(|(i, &c)| (c, i as u32))
.collect();
CharTokenizer { itos, stoi }
}
pub fn chars(&self) -> &[char] {
&self.itos
}
pub fn unknown_chars(&self, text: &str) -> Vec<char> {
let mut seen = Vec::new();
for c in text.chars() {
if !self.stoi.contains_key(&c) && !seen.contains(&c) {
seen.push(c);
}
}
seen
}
pub fn encode_lossy(&self, text: &str) -> Vec<u32> {
text.chars()
.filter_map(|c| self.stoi.get(&c).copied())
.collect()
}
pub fn to_json(&self) -> String {
let map: HashMap<String, u32> = self
.stoi
.iter()
.map(|(&c, &i)| (c.to_string(), i))
.collect();
serde_json::to_string(&map).expect("char vocab serializes")
}
pub fn from_json(s: &str) -> Result<Self> {
let map: HashMap<String, u32> = serde_json::from_str(s)?;
let mut itos = vec!['\0'; map.len()];
let mut filled = vec![false; map.len()];
for (k, id) in &map {
let mut cs = k.chars();
let (Some(c), None) = (cs.next(), cs.next()) else {
return Err(ForgeError::Tokenizer(format!(
"char vocab entry {k:?} is not a single character"
)));
};
let idx = *id as usize;
if idx >= itos.len() {
return Err(ForgeError::Tokenizer(format!(
"char vocab id {id} out of range for a {}-entry vocab",
itos.len()
)));
}
if filled[idx] {
return Err(ForgeError::Tokenizer(format!(
"char vocab id {id} assigned twice"
)));
}
itos[idx] = c;
filled[idx] = true;
}
let stoi = itos
.iter()
.enumerate()
.map(|(i, &c)| (c, i as u32))
.collect();
Ok(CharTokenizer { itos, stoi })
}
pub fn save_json(&self, path: impl AsRef<Path>) -> Result<()> {
std::fs::write(path, self.to_json())?;
Ok(())
}
pub fn from_json_file(path: impl AsRef<Path>) -> Result<Self> {
Self::from_json(&std::fs::read_to_string(path)?)
}
}
impl Tokenizer for CharTokenizer {
fn encode(&self, text: &str) -> Result<Vec<u32>> {
text.chars()
.map(|c| {
self.stoi.get(&c).copied().ok_or_else(|| {
ForgeError::Tokenizer(format!(
"character {c:?} is outside the {}-token vocabulary",
self.itos.len()
))
})
})
.collect()
}
fn decode(&self, ids: &[u32]) -> String {
ids.iter()
.filter_map(|&i| self.itos.get(i as usize))
.collect()
}
fn decode_bytes(&self, ids: &[u32]) -> Vec<u8> {
self.decode(ids).into_bytes()
}
fn vocab_size(&self) -> usize {
self.itos.len()
}
}