use std::path::{Path, PathBuf};
use tokenizers::Tokenizer;
use super::*;
use crate::audio::align::acoustic::{
AcousticGeometry, Granularity, LetterCase, OutputKind, Tokenization,
};
const DICT_ENTRIES: [(&str, u32); VOCAB_SIZE] = [
("-", 0),
("|", 1),
("E", 2),
("T", 3),
("A", 4),
("O", 5),
("N", 6),
("I", 7),
("H", 8),
("S", 9),
("R", 10),
("D", 11),
("L", 12),
("U", 13),
("M", 14),
("W", 15),
("C", 16),
("F", 17),
("G", 18),
("Y", 19),
("P", 20),
("B", 21),
("V", 22),
("K", 23),
("'", 24),
("X", 25),
("J", 26),
("Q", 27),
("Z", 28),
];
fn asset_path() -> PathBuf {
Path::new(env!("CARGO_MANIFEST_DIR"))
.join("src/audio/align/assets/chordai_base960h_tokenizer.json")
}
fn load_tokenizer() -> Tokenizer {
Tokenizer::from_bytes(tokenizer_json_bytes()).expect("embedded asset must parse")
}
#[test]
fn vocab_size_is_29() {
assert_eq!(VOCAB_SIZE, 29);
assert_eq!(DICT_ENTRIES.len(), 29);
}
#[test]
fn blank_id_is_zero() {
assert_eq!(BLANK_ID, 0);
}
#[test]
fn word_delimiter_is_pipe() {
assert_eq!(WORD_DELIMITER, "|");
}
#[test]
fn embedded_bytes_match_the_committed_file_on_disk() {
let disk_bytes = std::fs::read(asset_path()).expect("committed asset must be readable");
assert_eq!(
disk_bytes,
tokenizer_json_bytes(),
"include_bytes! must reflect the committed asset exactly"
);
}
#[test]
fn on_disk_asset_round_trips_through_asrys_loader_shape() {
let bytes = std::fs::read(asset_path()).expect("read tokenizer asset");
let tok = Tokenizer::from_bytes(&bytes).expect("Tokenizer::from_bytes parses the asset");
assert_eq!(tok.get_vocab_size(true), VOCAB_SIZE);
}
#[test]
fn tokenizer_vocab_size_matches_vocab_size_exactly() {
let tok = load_tokenizer();
assert_eq!(tok.get_vocab_size(true), VOCAB_SIZE);
assert_eq!(tok.get_vocab_size(false), VOCAB_SIZE);
}
#[test]
fn blank_token_resolves_to_blank_id() {
let tok = load_tokenizer();
assert_eq!(tok.token_to_id("-"), Some(BLANK_ID));
}
#[test]
fn word_delimiter_resolves_via_token_to_id() {
let tok = load_tokenizer();
assert_eq!(tok.token_to_id(WORD_DELIMITER), Some(1));
}
#[test]
fn every_dict_entry_round_trips_through_token_to_id() {
let tok = load_tokenizer();
for (token, expected_id) in DICT_ENTRIES {
assert_eq!(
tok.token_to_id(token),
Some(expected_id),
"token {token:?} must resolve to id {expected_id}"
);
}
assert_eq!(
tok.get_vocab(true).len(),
VOCAB_SIZE,
"vocab must contain exactly these 29 entries, no more"
);
}
#[test]
fn truncated_asset_is_rejected_by_tokenizer_from_bytes() {
let bytes = tokenizer_json_bytes();
let truncated = &bytes[..bytes.len() / 2];
assert!(
std::str::from_utf8(truncated).is_ok(),
"fixture assumption: the halfway point must still land on an ASCII byte"
);
assert!(
Tokenizer::from_bytes(truncated).is_err(),
"truncated JSON must not parse as a valid tokenizer"
);
}
#[test]
fn corrupted_delimiter_id_is_caught_by_the_round_trip_check() {
let text = std::str::from_utf8(tokenizer_json_bytes()).expect("asset is UTF-8");
let needle = "\"|\": 1,";
assert!(
text.contains(needle),
"fixture assumption: the delimiter's exact `{needle}` line must be present in the asset \
for this mutation to actually corrupt it"
);
let mutated = text.replacen(needle, "\"|\": 91,", 1);
let tok = Tokenizer::from_bytes(mutated.as_bytes()).expect("still structurally valid JSON");
assert_eq!(tok.get_vocab_size(true), VOCAB_SIZE);
assert_ne!(tok.token_to_id(WORD_DELIMITER), Some(1));
assert_eq!(tok.token_to_id(WORD_DELIMITER), Some(91));
}
#[test]
fn corrupted_vocab_entry_removal_is_caught_by_vocab_size_check() {
let text = std::str::from_utf8(tokenizer_json_bytes()).expect("asset is UTF-8");
let needle = "\"Q\": 27,\n";
assert!(
text.contains(needle),
"fixture assumption: the `{needle:?}` line must be present in the asset for this \
mutation to actually corrupt it"
);
let mutated = text.replacen(needle, "", 1);
let tok = Tokenizer::from_bytes(mutated.as_bytes()).expect("still structurally valid JSON");
let corrupted_size = tok.get_vocab_size(true);
assert_ne!(
corrupted_size, VOCAB_SIZE,
"removing an entry must change the observed vocab size"
);
assert_eq!(corrupted_size, VOCAB_SIZE - 1);
assert_eq!(tok.token_to_id("Q"), None);
}
const STAGED_DICT_SHA256: &str = "ef41495ab958d4416ad2f81ea51a77d4a3c79cace96e92e978c443c7bfbdd2e5";
fn staged_dict() -> Vec<u8> {
let entries: Vec<String> = DICT_ENTRIES
.iter()
.map(|(token, id)| format!("\"{token}\": {id}"))
.collect();
let dict = format!("{{{}}}", entries.join(", ")).into_bytes();
assert_eq!(
sha256_hex(&dict),
STAGED_DICT_SHA256,
"the rebuilt table must be the staged file, byte for byte"
);
dict
}
fn sha256_hex(bytes: &[u8]) -> String {
use sha2::{Digest, Sha256};
Sha256::digest(bytes)
.iter()
.map(|byte| format!("{byte:02x}"))
.collect()
}
fn json(bytes: &[u8]) -> serde_json::Value {
serde_json::from_slice(bytes).expect("JSON")
}
#[test]
fn the_staged_dict_reads_back_as_the_bundled_table() {
let vocabulary = Vocabulary::from_json(&staged_dict()).expect("the staged table reads");
assert_eq!(vocabulary.size().get(), VOCAB_SIZE);
let document = vocabulary.tokenizer_json(&AcousticContract::BASE960H);
assert_eq!(
json(&document),
json(tokenizer_json_bytes()),
"the written document must be the committed asset"
);
let tok = Tokenizer::from_bytes(&document).expect("the document parses");
for (token, id) in DICT_ENTRIES {
assert_eq!(tok.token_to_id(token), Some(id), "{token:?}");
}
}
fn declared_specials(tokenizer: &Tokenizer) -> Vec<(u32, String)> {
let mut specials: Vec<(u32, String)> = tokenizer
.get_added_vocabulary()
.get_added_tokens_decoder()
.iter()
.filter(|(_, token)| token.special)
.map(|(&id, token)| (id, token.content.clone()))
.collect();
specials.sort_unstable();
specials
}
#[test]
fn the_bundled_document_declares_exactly_the_staged_contracts_specials() {
let staged = [(0u32, "-".to_owned()), (1, "|".to_owned())];
let written = Vocabulary::bundled().tokenizer_json(&AcousticContract::BASE960H);
for document in [written.as_slice(), tokenizer_json_bytes()] {
let value = json(document);
let added = value["added_tokens"].as_array().expect("an array");
assert_eq!(added.len(), 2);
for (entry, (id, content)) in added.iter().zip(&staged) {
assert_eq!(entry["id"], *id);
assert_eq!(entry["content"], content.as_str());
assert_eq!(entry["special"], true);
}
let tok = Tokenizer::from_bytes(document).expect("the document parses");
assert_eq!(declared_specials(&tok), staged);
assert_eq!(tok.get_vocab_size(true), VOCAB_SIZE);
assert_eq!(tok.token_to_id("-"), Some(0));
assert_eq!(tok.token_to_id("|"), Some(1));
}
assert_eq!(json(&written), json(tokenizer_json_bytes()));
}
fn contract_with(
blank: u32,
delimiter: WordDelimiter,
specials: &'static [&'static str],
) -> AcousticContract {
AcousticContract::new(
blank,
AcousticGeometry::WAV2VEC2,
Tokenization::new(
delimiter,
LetterCase::Upper,
Granularity::Character,
specials,
),
OutputKind::Logits,
)
}
#[test]
fn a_contracts_document_declares_its_non_lexical_tokens_by_statement() {
let table = br#"{"<pad>": 0, "<s>": 1, "</s>": 2, "<unk>": 3, "|": 4, "E": 5, "T": 6,
"A": 7, "O": 8, "N": 9, "I": 10, "H": 11, "S": 12, "R": 13, "D": 14, "L": 15, "U": 16,
"M": 17, "W": 18, "C": 19, "F": 20, "G": 21, "Y": 22, "P": 23, "B": 24, "V": 25, "K": 26,
"'": 27, "X": 28, "J": 29, "Q": 30, "Z": 31}"#;
let vocabulary = Vocabulary::from_json(table).expect("the table reads");
let parse = |contract: &AcousticContract| {
Tokenizer::from_bytes(vocabulary.tokenizer_json(contract)).expect("the document parses")
};
let hf = parse(&contract_with(
0,
WordDelimiter::Pipe,
&["<s>", "</s>", "<unk>", "<mask>"],
));
assert_eq!(
declared_specials(&hf),
[
(0, "<pad>".to_owned()),
(1, "<s>".to_owned()),
(2, "</s>".to_owned()),
(3, "<unk>".to_owned()),
(4, "|".to_owned()),
]
);
assert_eq!(hf.get_vocab_size(true), 32);
let bare = parse(&contract_with(0, WordDelimiter::Absent, &[]));
assert_eq!(declared_specials(&bare), [(0, "<pad>".to_owned())]);
assert_eq!(bare.get_vocab_size(true), 32);
}
#[test]
fn the_parse_drops_a_special_spelled_as_the_empty_string() {
let vocabulary =
Vocabulary::from_json(br#"{"<pad>": 0, "|": 1, "A": 2, "": 3}"#).expect("the table reads");
let document = vocabulary.tokenizer_json(&contract_with(0, WordDelimiter::Pipe, &[""]));
let value = json(&document);
let declared: Vec<(u64, &str)> = value["added_tokens"]
.as_array()
.expect("an array")
.iter()
.map(|entry| {
(
entry["id"].as_u64().expect("an id"),
entry["content"].as_str().expect("a content"),
)
})
.collect();
assert_eq!(declared, [(0, "<pad>"), (1, "|"), (3, "")]);
let tok = Tokenizer::from_bytes(&document).expect("the document parses");
assert_eq!(
declared_specials(&tok),
[(0, "<pad>".to_owned()), (1, "|".to_owned())]
);
assert_eq!(tok.token_to_id(""), Some(3), "the table still spells it");
}
#[test]
fn bundled_tokens_are_the_committed_tables() {
let vocabulary = Vocabulary::bundled();
let bundled: Vec<&str> = vocabulary.tokens().collect();
let staged: Vec<&str> = DICT_ENTRIES.iter().map(|(token, _)| *token).collect();
assert_eq!(bundled, staged);
let read = Vocabulary::from_json(&staged_dict()).expect("the staged table reads");
assert_eq!(read.tokens().collect::<Vec<_>>(), staged);
assert!(read.contains("|") && read.contains("A") && !read.contains("a"));
}
#[test]
fn bundled_is_the_committed_table() {
let bundled = Vocabulary::bundled();
assert_eq!(bundled.size().get(), VOCAB_SIZE);
assert_eq!(
json(&bundled.tokenizer_json(&AcousticContract::BASE960H)),
json(tokenizer_json_bytes())
);
}
#[test]
fn a_table_carries_no_blank_whatever_its_names() {
for (table, size) in [
(
r#"{"<pad>": 0, "<s>": 1, "</s>": 2, "<unk>": 3, "|": 4, "A": 5}"#,
6,
),
(r#"{"a": 0, "b": 1, "|": 2, "[UNK]": 3, "[PAD]": 4}"#, 5),
(r#"{"<blank>": 0, "<pad>": 1, "a": 2}"#, 3),
(r#"{"-": 0, "<pad>": 1, "a": 2}"#, 3),
(r#"{"a": 0, "-": 1, "b": 2}"#, 3),
(r#"{"a": 0, "b": 1}"#, 2),
] {
let vocabulary = Vocabulary::from_json(table.as_bytes()).expect("the table reads");
assert_eq!(vocabulary.size().get(), size, "{table}");
let last = size - 1;
let contract = contract_with(
u32::try_from(last).expect("a small id"),
WordDelimiter::Absent,
&[],
);
let document = json(&vocabulary.tokenizer_json(&contract));
assert_eq!(
document["model"]["vocab"],
json(table.as_bytes()),
"{table}: every token at its own id"
);
let added = document["added_tokens"].as_array().expect("an array");
assert_eq!(added.len(), 1, "{table}: the stated blank alone");
assert_eq!(added[0]["id"], last, "{table}");
assert_eq!(
added[0]["content"],
vocabulary.tokens().nth(last).expect("the last entry"),
"{table}"
);
}
}
#[test]
fn a_table_that_skips_or_repeats_an_id_is_refused_by_name() {
for (table, id) in [
(r#"{"-": 0, "|": 2}"#, 1),
(r#"{"-": 0, "a": 0}"#, 1),
(r#"{"<pad>": 1, "a": 2}"#, 0),
] {
let Err(VocabularyError::MissingId(missing)) = Vocabulary::from_json(table.as_bytes()) else {
panic!("{table} must be refused as a missing id");
};
assert_eq!((missing.id(), missing.entries()), (id, 2), "{table}");
}
}
#[test]
fn a_token_named_twice_is_refused_by_name() {
for (table, token) in [
(r#"{"<pad>": 0, "A": 0, "A": 1}"#, "A"),
(r#"{"-": 0, "a": 1, "a": 1}"#, "a"),
(r#"{"<pad>": 0, "A": 1, "A": 2}"#, "A"),
] {
let Err(VocabularyError::DuplicateToken(repeated)) = Vocabulary::from_json(table.as_bytes())
else {
panic!("{table} must be refused as a repeated token");
};
assert_eq!(repeated, token, "{table}");
}
}
#[test]
fn an_empty_table_is_refused_by_name() {
assert!(matches!(
Vocabulary::from_json(b"{}"),
Err(VocabularyError::Empty)
));
}
#[test]
fn bytes_that_are_not_a_token_table_are_refused_by_name() {
let inputs: [&[u8]; 7] = [
b"[1, 2]",
br#"{"-": -1}"#,
br#"{"-": "0"}"#,
br#"{"-": 1.5}"#,
br#"{"-": 4294967296}"#,
b"not json",
tokenizer_json_bytes(),
];
for input in inputs {
assert!(
matches!(Vocabulary::from_json(input), Err(VocabularyError::Parse(_))),
"{}",
String::from_utf8_lossy(input)
);
}
}
#[test]
fn from_file_reads_the_table_beside_a_model() {
let dir = tempfile::tempdir().expect("a temporary directory");
let path = dir.path().join("base960h_dict.json");
std::fs::write(&path, staged_dict()).expect("write the table");
let vocabulary = Vocabulary::from_file(&path).expect("the table reads");
assert_eq!(vocabulary.size().get(), VOCAB_SIZE);
let missing = dir.path().join("absent_dict.json");
let Err(VocabularyError::Read(read)) = Vocabulary::from_file(&missing) else {
panic!("an absent file must be refused as a read failure");
};
assert_eq!(read.path(), missing);
let source = std::error::Error::source(&read)
.and_then(|source| source.downcast_ref::<std::io::Error>())
.expect("the I/O failure is the source");
assert_eq!(source.kind(), std::io::ErrorKind::NotFound);
}
#[test]
fn debug_names_the_size() {
assert_eq!(
format!("{:?}", Vocabulary::bundled()),
"Vocabulary { size: 29, .. }"
);
}