use std::collections::BTreeMap;
use std::path::Path;
use crate::error::{FocrError, FocrResult};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Stream {
Rhythm,
Pitch,
Lift,
Note,
}
#[derive(Debug)]
struct WordLevelTable {
id_to_token: Vec<String>,
}
impl WordLevelTable {
fn from_json(text: &str, origin: &str) -> FocrResult<Self> {
let doc: serde_json::Value = serde_json::from_str(text).map_err(|e| {
FocrError::FormatMismatch(format!("{origin}: not valid tokenizer JSON: {e}"))
})?;
let model = &doc["model"];
if model["type"].as_str() != Some("WordLevel") {
return Err(FocrError::FormatMismatch(format!(
"{origin}: model.type {:?} is not WordLevel",
model["type"]
)));
}
let vocab = model["vocab"].as_object().ok_or_else(|| {
FocrError::FormatMismatch(format!("{origin}: model.vocab is not an object"))
})?;
let mut by_id: BTreeMap<u64, &str> = BTreeMap::new();
for (token, id) in vocab {
let id = id.as_u64().ok_or_else(|| {
FocrError::FormatMismatch(format!("{origin}: id for {token:?} is not a u64"))
})?;
if by_id.insert(id, token).is_some() {
return Err(FocrError::FormatMismatch(format!(
"{origin}: duplicate id {id}"
)));
}
}
let n = by_id.len() as u64;
if by_id.keys().next_back().map(|&k| k + 1) != Some(n) {
return Err(FocrError::FormatMismatch(format!(
"{origin}: id space is not dense 0..{n}"
)));
}
Ok(Self {
id_to_token: by_id.into_values().map(str::to_owned).collect(),
})
}
fn token(&self, id: u32) -> Option<&str> {
self.id_to_token.get(id as usize).map(String::as_str)
}
}
#[derive(Debug)]
pub struct MusicTokenizer {
rhythm: WordLevelTable,
pitch: WordLevelTable,
lift: WordLevelTable,
note: WordLevelTable,
}
pub const PAD_ID: u32 = 0;
pub const BOS_ID: u32 = 1;
pub const EOS_ID: u32 = 2;
impl MusicTokenizer {
pub fn from_dir(dir: &Path) -> FocrResult<Self> {
let read = |stem: &str| -> FocrResult<(String, String)> {
let path = dir.join(format!("tokenizer_{stem}.json"));
let text = std::fs::read_to_string(&path).map_err(|e| {
FocrError::ModelNotFound(format!(
"TrOMR tokenizer table missing: {}: {e}",
path.display()
))
})?;
Ok((text, path.display().to_string()))
};
let (rhythm, rhythm_src) = read("rhythm")?;
let (pitch, pitch_src) = read("pitch")?;
let (lift, lift_src) = read("lift")?;
let (note, note_src) = read("note")?;
Self::from_json_tables(
[&rhythm, &pitch, &lift, ¬e],
[&rhythm_src, &pitch_src, &lift_src, ¬e_src],
)
}
pub fn from_json_tables(tables: [&str; 4], sources: [&str; 4]) -> FocrResult<Self> {
let load = |idx: usize, want: usize| -> FocrResult<WordLevelTable> {
let table = WordLevelTable::from_json(tables[idx], sources[idx])?;
if table.id_to_token.len() != want {
return Err(FocrError::FormatMismatch(format!(
"{}: {} ids, census expects {want} (spec §9)",
sources[idx],
table.id_to_token.len()
)));
}
Ok(table)
};
Ok(Self {
rhythm: load(0, 260)?,
pitch: load(1, 71)?,
lift: load(2, 7)?,
note: load(3, 2)?,
})
}
fn table(&self, stream: Stream) -> &WordLevelTable {
match stream {
Stream::Rhythm => &self.rhythm,
Stream::Pitch => &self.pitch,
Stream::Lift => &self.lift,
Stream::Note => &self.note,
}
}
#[must_use]
pub fn token(&self, stream: Stream, id: u32) -> Option<&str> {
self.table(stream).token(id)
}
#[must_use]
pub fn detokenize(&self, stream: Stream, ids: &[u32]) -> Vec<String> {
ids.iter()
.filter_map(|&id| {
let tok = self.token(stream, id).unwrap_or("");
let cleaned = tok.replace('Ġ', " ");
let cleaned = cleaned.trim();
(!matches!(cleaned, "[BOS]" | "[EOS]" | "[PAD]")).then(|| cleaned.to_owned())
})
.collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
fn fixture_dir() -> std::path::PathBuf {
std::path::Path::new(env!("CARGO_MANIFEST_DIR")).join("tests/fixtures/tromr")
}
fn load() -> MusicTokenizer {
MusicTokenizer::from_dir(&fixture_dir()).expect("committed tables load")
}
#[test]
fn tables_match_the_census_layout() {
let tk = load();
for (id, want) in [
(0, "[PAD]"),
(1, "[BOS]"),
(2, "[EOS]"),
(3, "+"),
(4, "|"),
(5, "barline"),
] {
assert_eq!(tk.token(Stream::Rhythm, id), Some(want));
}
assert_eq!(tk.token(Stream::Pitch, 0), Some("nonote"));
assert_eq!(tk.token(Stream::Pitch, 1), Some("note-C0"));
assert_eq!(tk.token(Stream::Pitch, 70), Some("note-B9"));
assert_eq!(tk.token(Stream::Lift, 0), Some("nonote"));
assert_eq!(tk.token(Stream::Lift, 1), Some("lift_null"));
assert_eq!(tk.token(Stream::Lift, 6), Some("lift_N"));
assert_eq!(tk.token(Stream::Note, 0), Some("nonote"));
assert_eq!(tk.token(Stream::Note, 1), Some("note"));
assert_eq!(tk.token(Stream::Note, 2), None);
assert_eq!(tk.token(Stream::Rhythm, 260), None);
assert_eq!(tk.token(Stream::Rhythm, PAD_ID), Some("[PAD]"));
assert_eq!(tk.token(Stream::Rhythm, BOS_ID), Some("[BOS]"));
assert_eq!(tk.token(Stream::Rhythm, EOS_ID), Some("[EOS]"));
}
#[test]
fn detokenize_matches_the_oracle_goldens() {
let tk = load();
let text = std::fs::read_to_string(fixture_dir().join("detokenize_goldens.json"))
.expect("goldens committed");
let gold: serde_json::Value = serde_json::from_str(&text).expect("goldens parse");
for (name, stream) in [
("rhythm", Stream::Rhythm),
("pitch", Stream::Pitch),
("lift", Stream::Lift),
("note", Stream::Note),
] {
let entry = &gold[name];
let ids: Vec<u32> = entry["probe_ids"]
.as_array()
.unwrap()
.iter()
.map(|v| u32::try_from(v.as_u64().unwrap()).unwrap())
.collect();
let want: Vec<String> = entry["detokenized"]
.as_array()
.unwrap()
.iter()
.map(|v| v.as_str().unwrap().to_owned())
.collect();
assert_eq!(tk.detokenize(stream, &ids), want, "{name}");
assert!(
entry["oor_token_is_none"].as_bool().unwrap(),
"{name}: oracle confirms out-of-range => None"
);
}
}
#[test]
fn detokenize_keeps_oor_empty_and_drops_specials_by_token() {
let tk = load();
assert_eq!(
tk.detokenize(Stream::Rhythm, &[BOS_ID, 5, 9999, EOS_ID, PAD_ID]),
vec!["barline".to_owned(), String::new()]
);
assert_eq!(
tk.detokenize(Stream::Pitch, &[1, 0, 2]),
vec![
"note-C0".to_owned(),
"nonote".to_owned(),
"note-D0".to_owned()
]
);
}
#[test]
fn loader_error_paths_are_clean() {
let err = MusicTokenizer::from_dir(std::path::Path::new("/nonexistent/tromr"))
.expect_err("missing tables must error");
assert!(matches!(err, FocrError::ModelNotFound(_)), "{err:?}");
for (bad, why) in [
(r#"{"model":{"type":"BPE","vocab":{}}}"#, "not WordLevel"),
(
r#"{"model":{"type":"WordLevel","vocab":{"a":0,"b":2}}}"#,
"gap",
),
(
r#"{"model":{"type":"WordLevel","vocab":"x"}}"#,
"not an object",
),
("not json", "not JSON"),
] {
let err = WordLevelTable::from_json(bad, "synthetic").expect_err(why);
assert!(
matches!(err, FocrError::FormatMismatch(_)),
"{why}: {err:?}"
);
}
}
}