use smol_str::format_smolstr;
use tokenizers::Tokenizer;
use crate::{
runner::aligner::algorithm::{
errors::{EmissionsError, EmissionsFailure},
trellis_beam::WILDCARD_TOKEN_ID,
},
types::Lang,
};
#[derive(Debug)]
pub struct TokenizedText {
token_ids: Vec<i32>,
word_idx_per_token: Vec<Option<usize>>,
separator_token_id: Option<u32>,
}
impl TokenizedText {
#[must_use]
pub const fn new(
token_ids: Vec<i32>,
word_idx_per_token: Vec<Option<usize>>,
separator_token_id: Option<u32>,
) -> Self {
Self {
token_ids,
word_idx_per_token,
separator_token_id,
}
}
#[must_use]
pub fn token_ids(&self) -> &[i32] {
&self.token_ids
}
#[must_use]
pub fn word_idx_per_token(&self) -> &[Option<usize>] {
&self.word_idx_per_token
}
#[must_use]
pub const fn separator_token_id(&self) -> Option<u32> {
self.separator_token_id
}
}
fn is_skippable_internal_punct(c: char) -> bool {
c == '.'
}
fn vocab_id(tokenizer: &Tokenizer, token: &str, unk_token_id: Option<u32>) -> Option<u32> {
tokenizer
.token_to_id(token)
.filter(|&id| Some(id) != unk_token_id)
}
fn consume_oov_decision(
oov_decisions: &[crate::core::ResolvedOov],
oov_consumed: &mut usize,
site_label: &str,
) -> Result<crate::core::OovDecision, EmissionsError> {
let decision = oov_decisions
.get(*oov_consumed)
.map(|r| r.decision())
.ok_or_else(|| {
EmissionsError::Tokenization(EmissionsFailure::new(format_smolstr!(
"oov_decisions ran out at index {} ({site_label}); call detect_oov_events \
first to size the decisions vec correctly",
*oov_consumed,
)))
})?;
*oov_consumed += 1;
Ok(decision)
}
fn boundary_fail_closed(position: &str) -> EmissionsError {
EmissionsError::SemanticOutOfVocab(EmissionsFailure::new(format_smolstr!(
"BoundaryPunct ({position}) resolved as FailClosed by caller policy; \
no word alignment produced for the supplied tokens."
)))
}
#[allow(
clippy::too_many_arguments,
reason = "8 args mirror the wav2vec2 tokenisation contract \
(tokenizer, text, word_count, delimiter flag, casing \
flag, unk id, wildcard map, output buffer); each is a \
distinct semantic input from a different upstream pass"
)]
pub fn detect_oov_events(
tokenizer: &Tokenizer,
normalized: &str,
word_count: usize,
uppercase_input: bool,
unk_token_id: Option<u32>,
language: &Lang,
wildcard_boundary_per_word: &[crate::runner::aligner::normalizer::WildcardBoundary],
) -> Result<Vec<crate::core::OovEvent>, EmissionsError> {
use crate::core::OovEvent;
let mut events: Vec<OovEvent> = Vec::new();
let words: Vec<&str> = normalized.split_whitespace().collect();
if words.len() != word_count {
return Err(EmissionsError::Tokenization(EmissionsFailure::new(
format_smolstr!(
"word_count mismatch: caller={}, normalized has {}",
word_count,
words.len(),
),
)));
}
if !wildcard_boundary_per_word.is_empty() && wildcard_boundary_per_word.len() != word_count {
return Err(EmissionsError::Tokenization(EmissionsFailure::new(
format_smolstr!(
"wildcard_boundary_per_word.len() = {} != word_count = {}",
wildcard_boundary_per_word.len(),
word_count,
),
)));
}
let mut tmp_buf = String::with_capacity(8);
let mut char_index: usize = 0;
for (word_index, word) in words.iter().enumerate() {
let boundary = wildcard_boundary_per_word
.get(word_index)
.copied()
.unwrap_or(crate::runner::aligner::normalizer::WildcardBoundary::NONE);
let prefix_wildcards = boundary.prefix();
let suffix_wildcards = boundary.suffix();
for _ in 0..prefix_wildcards {
events.push(OovEvent::new(
crate::core::OovKind::BoundaryPunct,
char_index,
word_index,
language.clone(),
));
}
for ch in word.chars() {
if is_skippable_internal_punct(ch) {
events.push(OovEvent::new(
crate::core::OovKind::InternalPunct(ch),
char_index,
word_index,
language.clone(),
));
char_index += 1;
continue;
}
let projected = if uppercase_input {
ch.to_ascii_uppercase()
} else {
ch
};
tmp_buf.clear();
tmp_buf.push(projected);
if vocab_id(tokenizer, &tmp_buf, unk_token_id).is_none() {
events.push(OovEvent::new(
crate::core::OovKind::Symbol(ch),
char_index,
word_index,
language.clone(),
));
}
char_index += 1;
}
for _ in 0..suffix_wildcards {
events.push(OovEvent::new(
crate::core::OovKind::BoundaryPunct,
char_index,
word_index,
language.clone(),
));
}
if word_index + 1 < words.len() {
char_index += 1;
}
}
Ok(events)
}
pub fn tokenize_with_word_map(
tokenizer: &Tokenizer,
normalized: &str,
word_count: usize,
use_word_delimiter: bool,
uppercase_input: bool,
unk_token_id: Option<u32>,
wildcard_boundary_per_word: &[crate::runner::aligner::normalizer::WildcardBoundary],
language: &Lang,
oov_decisions: &[crate::core::ResolvedOov],
) -> Result<TokenizedText, EmissionsError> {
let pre_events = detect_oov_events(
tokenizer,
normalized,
word_count,
uppercase_input,
unk_token_id,
language,
wildcard_boundary_per_word,
)?;
if pre_events.len() != oov_decisions.len() {
return Err(EmissionsError::Tokenization(EmissionsFailure::new(
format_smolstr!(
"oov_decisions length {} does not match the {} OOV events detected for this \
text; this typically means the caller passed decisions from a different \
chunk's text. Re-run `detect_oov_events` for the chunk's normalised text \
and re-decide before calling `tokenize_with_word_map`.",
oov_decisions.len(),
pre_events.len(),
),
)));
}
for (i, (pre, resolved)) in pre_events.iter().zip(oov_decisions.iter()).enumerate() {
if !resolved.event().matches_position(pre) {
return Err(EmissionsError::Tokenization(EmissionsFailure::new(
format_smolstr!(
"oov_decisions[{i}] was produced for a different OOV event than the one \
this chunk's text actually has at position {i}: supplied={:?} but \
detected={:?}. This typically means the caller reused decisions from a \
previous chunk whose OOV count happened to match. Re-run \
`detect_oov_events` for THIS chunk's normalised text and re-decide.",
resolved.event(),
pre,
),
)));
}
}
let mut oov_consumed: usize = 0;
let mut token_ids: Vec<i32> = Vec::with_capacity(normalized.len() + word_count * 2);
let mut word_idx_per_token: Vec<Option<usize>> = Vec::with_capacity(token_ids.capacity());
let words: Vec<&str> = normalized.split_whitespace().collect();
if words.len() != word_count {
return Err(EmissionsError::Tokenization(EmissionsFailure::new(
format_smolstr!(
"word_count mismatch: caller={}, normalized has {}",
word_count,
words.len()
),
)));
}
if !wildcard_boundary_per_word.is_empty() && wildcard_boundary_per_word.len() != word_count {
return Err(EmissionsError::Tokenization(EmissionsFailure::new(
format_smolstr!(
"wildcard_boundary_per_word.len() = {} != word_count = {}",
wildcard_boundary_per_word.len(),
word_count
),
)));
}
let mut per_word_tokens: Vec<Vec<i32>> = Vec::with_capacity(words.len());
let mut tmp_buf = String::with_capacity(8);
for (wi, word) in words.iter().enumerate() {
let boundary = wildcard_boundary_per_word
.get(wi)
.copied()
.unwrap_or(crate::runner::aligner::normalizer::WildcardBoundary::NONE);
let prefix_wildcards = boundary.prefix();
let suffix_wildcards = boundary.suffix();
let mut word_tokens: Vec<i32> = Vec::with_capacity(word.len());
for _ in 0..prefix_wildcards {
let decision =
consume_oov_decision(oov_decisions, &mut oov_consumed, "BoundaryPunct (prefix)")?;
match decision {
crate::core::OovDecision::Wildcard => word_tokens.push(WILDCARD_TOKEN_ID),
crate::core::OovDecision::FailClosed => {
return Err(boundary_fail_closed("prefix"));
}
}
}
for ch in word.chars() {
if is_skippable_internal_punct(ch) {
let decision = consume_oov_decision(oov_decisions, &mut oov_consumed, "InternalPunct")?;
match decision {
crate::core::OovDecision::Wildcard => word_tokens.push(WILDCARD_TOKEN_ID),
crate::core::OovDecision::FailClosed => {
return Err(EmissionsError::SemanticOutOfVocab(EmissionsFailure::new(
format_smolstr!(
"InternalPunct {ch:?} resolved as FailClosed by caller policy; \
no word alignment produced for the supplied tokens."
),
)));
}
}
continue;
}
let projected = if uppercase_input {
ch.to_ascii_uppercase()
} else {
ch
};
tmp_buf.clear();
tmp_buf.push(projected);
let Some(id) = vocab_id(tokenizer, &tmp_buf, unk_token_id) else {
let decision = consume_oov_decision(oov_decisions, &mut oov_consumed, "Symbol")?;
match decision {
crate::core::OovDecision::Wildcard => {
word_tokens.push(WILDCARD_TOKEN_ID);
}
crate::core::OovDecision::FailClosed => {
return Err(EmissionsError::SemanticOutOfVocab(EmissionsFailure::new(
format_smolstr!(
"OOV {ch:?} resolved as FailClosed by caller policy; \
no word alignment produced for the supplied tokens."
),
)));
}
}
continue;
};
let signed_id = i32::try_from(id).map_err(|_| {
EmissionsError::Tokenization(EmissionsFailure::new(format_smolstr!(
"tokenizer returned id {} which exceeds i32::MAX or aliases the wildcard \
sentinel; tokenizer / model mismatch?",
id
)))
})?;
if signed_id < 0 {
return Err(EmissionsError::Tokenization(EmissionsFailure::new(
format_smolstr!(
"tokenizer returned negative-after-cast id {} (raw {}); refusing to alias \
wildcard sentinel",
signed_id,
id
),
)));
}
word_tokens.push(signed_id);
}
for _ in 0..suffix_wildcards {
let decision =
consume_oov_decision(oov_decisions, &mut oov_consumed, "BoundaryPunct (suffix)")?;
match decision {
crate::core::OovDecision::Wildcard => word_tokens.push(WILDCARD_TOKEN_ID),
crate::core::OovDecision::FailClosed => {
return Err(boundary_fail_closed("suffix"));
}
}
}
per_word_tokens.push(word_tokens);
}
let delim_id = if use_word_delimiter {
tokenizer.token_to_id("|")
} else {
None
};
let mut last_emitted_word: Option<usize> = None;
for (word_idx, group) in per_word_tokens.iter().enumerate() {
if group.is_empty() {
continue;
}
if last_emitted_word.is_some()
&& let Some(d) = delim_id
{
let signed_d = i32::try_from(d).map_err(|_| {
EmissionsError::Tokenization(EmissionsFailure::new(format_smolstr!(
"tokenizer returned `|` delimiter id {} which exceeds i32::MAX",
d
)))
})?;
if signed_d < 0 {
return Err(EmissionsError::Tokenization(EmissionsFailure::new(
format_smolstr!(
"tokenizer returned negative-after-cast `|` delimiter id {} (raw {})",
signed_d,
d
),
)));
}
token_ids.push(signed_d);
word_idx_per_token.push(None);
}
for &id in group {
token_ids.push(id);
word_idx_per_token.push(Some(word_idx));
}
last_emitted_word = Some(word_idx);
}
if oov_consumed != oov_decisions.len() {
return Err(EmissionsError::Tokenization(EmissionsFailure::new(
format_smolstr!(
"oov_decisions length {} does not match the {} OOV chars actually \
encountered; this typically means the caller passed decisions \
from a different chunk's text. Re-run `detect_oov_events` for \
the chunk's normalised text and re-decide.",
oov_decisions.len(),
oov_consumed,
),
)));
}
Ok(TokenizedText {
token_ids,
word_idx_per_token,
separator_token_id: delim_id,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{
core::{OovEvent, OovKind},
runner::aligner::{
core::{detect_unk_token_id, load_tokenizer_bytes_with_compat},
normalizer::WildcardBoundary,
},
types::Lang,
};
const UPPERCASE_TOKENIZER_JSON: &str = r#"{
"version": "1.0",
"truncation": null,
"padding": null,
"added_tokens": [],
"normalizer": null,
"pre_tokenizer": {
"type": "Split",
"pattern": {"Regex": ""},
"behavior": "Isolated",
"invert": false
},
"post_processor": null,
"decoder": null,
"model": {
"type": "WordLevel",
"vocab": {
"<unk>": 0,
"<pad>": 1,
"|": 2,
"A": 3, "B": 4, "C": 5, "D": 6, "E": 7, "F": 8, "G": 9,
"H": 10, "I": 11, "J": 12, "K": 13, "L": 14, "M": 15,
"N": 16, "O": 17, "P": 18, "Q": 19, "R": 20, "S": 21,
"T": 22, "U": 23, "V": 24, "W": 25, "X": 26, "Y": 27, "Z": 28
},
"unk_token": "<unk>"
}
}"#;
fn uppercase_tokenizer() -> Tokenizer {
Tokenizer::from_bytes(UPPERCASE_TOKENIZER_JSON.as_bytes())
.expect("inline WordLevel tokenizer must parse")
}
fn tokenize_with_default_oov(
tokenizer: &Tokenizer,
normalized: &str,
word_count: usize,
use_word_delimiter: bool,
uppercase_input: bool,
unk_token_id: Option<u32>,
wildcard_boundary_per_word: &[crate::runner::aligner::normalizer::WildcardBoundary],
language: &Lang,
) -> Result<TokenizedText, EmissionsError> {
let events = detect_oov_events(
tokenizer,
normalized,
word_count,
uppercase_input,
unk_token_id,
language,
wildcard_boundary_per_word,
)?;
let decisions = crate::core::default_oov_decisions(&events);
tokenize_with_word_map(
tokenizer,
normalized,
word_count,
use_word_delimiter,
uppercase_input,
unk_token_id,
wildcard_boundary_per_word,
language,
&decisions,
)
}
#[test]
fn detect_oov_events_empty_for_in_vocab_text() {
let tok = uppercase_tokenizer();
let unk = tok.token_to_id("<unk>");
let events = detect_oov_events(
&tok,
"hello",
1,
true,
unk,
&Lang::En,
&[],
)
.expect("ok");
assert!(
events.is_empty(),
"in-vocab text should produce 0 events; got {events:?}"
);
}
#[test]
fn detect_oov_events_collects_in_source_order() {
let tok = uppercase_tokenizer();
let unk = tok.token_to_id("<unk>");
let events = detect_oov_events(&tok, "AT&T", 1, true, unk, &Lang::En, &[]).expect("ok");
assert_eq!(events.len(), 1);
assert_eq!(events[0].kind(), &crate::core::OovKind::Symbol('&'));
assert_eq!(events[0].word_index(), 0);
assert_eq!(events[0].language(), &Lang::En);
}
#[test]
fn tokenize_with_word_map_rejects_too_long_oov_decisions() {
let tok = uppercase_tokenizer();
let unk = tok.token_to_id("<unk>");
let real_event = detect_oov_events(&tok, "AT&T", 1, true, unk, &Lang::En, &[])
.expect("ok")
.pop()
.expect("AT&T has one OOV");
let extra_event =
crate::core::OovEvent::new(crate::core::OovKind::Symbol('?'), 99, 99, Lang::En);
let too_long = vec![
crate::core::ResolvedOov::new(real_event, crate::core::OovDecision::Wildcard),
crate::core::ResolvedOov::new(extra_event, crate::core::OovDecision::Wildcard),
];
let result =
tokenize_with_word_map(&tok, "AT&T", 1, true, true, unk, &[], &Lang::En, &too_long);
match result {
Err(EmissionsError::Tokenization(payload)) => {
assert!(
payload.message().contains("oov_decisions length 2")
&& payload.message().contains("1 OOV events detected"),
"diagnostic should cite the length mismatch; got {message}",
message = payload.message(),
);
}
other => panic!("expected TokenizationFailed mismatch; got {other:?}"),
}
}
#[test]
fn tokenize_with_word_map_rejects_too_long_decisions_even_when_first_is_fail_closed() {
let tok = uppercase_tokenizer();
let unk = tok.token_to_id("<unk>");
let real_event = detect_oov_events(&tok, "AT&T", 1, true, unk, &Lang::En, &[])
.expect("ok")
.pop()
.expect("AT&T has one OOV");
let extra_event =
crate::core::OovEvent::new(crate::core::OovKind::Symbol('?'), 99, 99, Lang::En);
let too_long = vec![
crate::core::ResolvedOov::new(real_event, crate::core::OovDecision::FailClosed),
crate::core::ResolvedOov::new(extra_event, crate::core::OovDecision::Wildcard),
];
let result =
tokenize_with_word_map(&tok, "AT&T", 1, true, true, unk, &[], &Lang::En, &too_long);
match result {
Err(EmissionsError::Tokenization(_)) => {
}
Err(EmissionsError::SemanticOutOfVocab(_)) => panic!(
"stale too-long decisions starting with FailClosed must surface as \
TokenizationFailed (the loud diagnostic); SemanticOutOfVocab is the \
silent recoverable path that masks the bug"
),
other => panic!("expected TokenizationFailed mismatch; got {other:?}"),
}
}
#[test]
fn tokenize_with_word_map_rejects_stale_same_length_decisions() {
let tok = uppercase_tokenizer();
let unk = tok.token_to_id("<unk>");
let stale_for_digit = detect_oov_events(&tok, "4", 1, true, unk, &Lang::En, &[]).expect("ok");
assert_eq!(stale_for_digit.len(), 1);
let stale_resolved = vec![crate::core::ResolvedOov::new(
stale_for_digit[0].clone(),
crate::core::OovDecision::Wildcard,
)];
let result = tokenize_with_word_map(
&tok,
"AT&T",
1,
true,
true,
unk,
&[],
&Lang::En,
&stale_resolved,
);
match result {
Err(EmissionsError::Tokenization(payload)) => {
assert!(
payload.message().contains("different OOV event"),
"diagnostic should cite the per-position identity mismatch; got {message}",
message = payload.message(),
);
}
other => panic!("expected TokenizationFailed identity mismatch; got {other:?}"),
}
}
#[test]
fn tokenize_with_word_map_accepts_mismatched_language_under_any_fallback() {
let tok = uppercase_tokenizer();
let unk = tok.token_to_id("<unk>");
let pre = detect_oov_events(&tok, "AT&T", 1, true, unk, &Lang::En, &[])
.expect("ok")
.pop()
.expect("AT&T has one OOV");
assert_eq!(pre.language(), &Lang::En);
let resolved = vec![crate::core::ResolvedOov::new(
crate::core::OovEvent::new(
pre.kind().clone(),
pre.char_index(),
pre.word_index(),
Lang::Ko,
),
crate::core::OovDecision::Wildcard,
)];
let result =
tokenize_with_word_map(&tok, "AT&T", 1, true, true, unk, &[], &Lang::En, &resolved);
assert!(
result.is_ok(),
"Any-fallback identity check must compare positional fields \
only (kind/char_index/word_index), not language. Got: {result:?}",
);
}
#[test]
fn detect_oov_events_tracks_word_index() {
let tok = uppercase_tokenizer();
let unk = tok.token_to_id("<unk>");
let events = detect_oov_events(&tok, "AT&T cost 43", 3, true, unk, &Lang::En, &[]).expect("ok");
let chars: Vec<Option<char>> = events.iter().map(|e| e.char()).collect();
let words: Vec<usize> = events.iter().map(|e| e.word_index()).collect();
assert_eq!(chars, vec![Some('&'), Some('4'), Some('3')]);
assert_eq!(words, vec![0, 2, 2]);
}
#[test]
fn detect_oov_events_surfaces_internal_punct() {
let tok = uppercase_tokenizer();
let unk = tok.token_to_id("<unk>");
let events = detect_oov_events(&tok, "U.S.A", 1, true, unk, &Lang::En, &[]).expect("ok");
let kinds: Vec<crate::core::OovKind> = events.iter().map(|e| e.kind().clone()).collect();
assert_eq!(
kinds,
vec![
crate::core::OovKind::InternalPunct('.'),
crate::core::OovKind::InternalPunct('.'),
],
"U.S.A should surface 2 InternalPunct events; got {events:?}",
);
}
#[test]
fn detect_oov_events_word_count_mismatch_errors() {
let tok = uppercase_tokenizer();
let unk = tok.token_to_id("<unk>");
let result = detect_oov_events(&tok, "hello world", 1, true, unk, &Lang::En, &[]);
assert!(matches!(result, Err(EmissionsError::Tokenization(_))));
}
#[test]
fn english_lowercase_word_uppercases_for_uppercase_only_vocab() {
let tok = uppercase_tokenizer();
let unk = tok.token_to_id("<unk>");
let result = tokenize_with_word_map(
&tok,
"hello",
1,
true,
true,
unk,
&[],
&Lang::En,
&[],
)
.expect("tokenisation must succeed with uppercase projection");
assert_eq!(result.token_ids.len(), 5);
let unk_i32 = unk.unwrap() as i32;
assert!(
result.token_ids.iter().all(|&id| id != unk_i32),
"no <unk> ids; got {:?}",
result.token_ids
);
let expected = ['H', 'E', 'L', 'L', 'O'].map(|c| {
tok
.token_to_id(&c.to_string())
.expect("uppercase letter in vocab") as i32
});
assert_eq!(result.token_ids, expected.to_vec());
}
#[test]
fn skippable_punctuation_only_yields_one_wildcard_token() {
let tok = uppercase_tokenizer();
let unk = tok.token_to_id("<unk>");
let result =
tokenize_with_default_oov(&tok, ".", 1, true, true, unk, &[], &Lang::En).expect("ok");
assert_eq!(result.token_ids, vec![WILDCARD_TOKEN_ID]);
assert_eq!(result.word_idx_per_token, vec![Some(0)]);
}
#[test]
fn internal_periods_in_abbreviation_strip_to_letters() {
let tok = uppercase_tokenizer();
let unk = tok.token_to_id("<unk>");
let result = tokenize_with_default_oov(
&tok,
"U.S.A",
1,
true,
true,
unk,
&[],
&Lang::En,
)
.expect("U.S.A. must tokenise via per-char strip");
assert_eq!(result.token_ids.len(), 5);
let unk_i32 = unk.unwrap() as i32;
assert!(
result.token_ids.iter().all(|&id| id != unk_i32),
"no <unk> ids must reach the lattice; got {:?}",
result.token_ids
);
let id_of = |c: char| tok.token_to_id(&c.to_string()).unwrap() as i32;
assert_eq!(
result.token_ids,
vec![
id_of('U'),
WILDCARD_TOKEN_ID,
id_of('S'),
WILDCARD_TOKEN_ID,
id_of('A'),
],
"internal-punct wildcards must land in source order, not appended at end"
);
assert_eq!(result.word_idx_per_token, vec![Some(0); 5]);
}
#[test]
fn partial_oov_alphanumeric_word_uses_wildcard() {
let tok = uppercase_tokenizer();
let unk = tok.token_to_id("<unk>");
let result =
tokenize_with_default_oov(&tok, "B2B", 1, true, true, unk, &[], &Lang::En).expect("ok");
assert_eq!(result.token_ids.len(), 3);
let b_id = tok.token_to_id("B").unwrap() as i32;
assert_eq!(result.token_ids[0], b_id);
assert_eq!(result.token_ids[1], WILDCARD_TOKEN_ID);
assert_eq!(result.token_ids[2], b_id);
assert_eq!(result.word_idx_per_token, vec![Some(0); 3]);
}
#[test]
fn all_digit_word_against_uppercase_vocab_uses_wildcards() {
let tok = uppercase_tokenizer();
let unk = tok.token_to_id("<unk>");
let result =
tokenize_with_default_oov(&tok, "1000", 1, true, true, unk, &[], &Lang::En).expect("ok");
assert_eq!(result.token_ids.len(), 4);
assert!(
result.token_ids.iter().all(|&id| id == WILDCARD_TOKEN_ID),
"every digit must become a wildcard; got {:?}",
result.token_ids
);
}
#[test]
fn ampersand_oov_drops_chunk() {
let tok = uppercase_tokenizer();
let unk = tok.token_to_id("<unk>");
let outcome = tokenize_with_default_oov(&tok, "AT&T", 1, true, true, unk, &[], &Lang::En);
match outcome {
Err(EmissionsError::SemanticOutOfVocab(payload)) => {
let message = payload.message();
assert!(
message.contains("'&'") || message.contains("\"&\""),
"diagnostic should cite the offending char; got {message}",
);
for banned in [
"ORT",
"worker",
"pool",
"Event::Error",
"ASR text preserved",
] {
assert!(
!message.contains(banned),
"OOV Display leaked {banned:?}: {message}"
);
}
}
other => panic!("expected SemanticOutOfVocab; got {other:?}"),
}
}
#[test]
fn accented_letter_uses_wildcard() {
let tok = uppercase_tokenizer();
let unk = tok.token_to_id("<unk>");
let result =
tokenize_with_default_oov(&tok, "café", 1, true, true, unk, &[], &Lang::En).expect("ok");
assert_eq!(result.token_ids.len(), 4);
let expected_letters = ['C', 'A', 'F'];
for (i, c) in expected_letters.iter().enumerate() {
assert_eq!(
result.token_ids[i],
tok.token_to_id(&c.to_string()).unwrap() as i32
);
}
assert_eq!(result.token_ids[3], WILDCARD_TOKEN_ID);
}
#[test]
fn middle_digit_word_no_longer_drops_chunk() {
let tok = uppercase_tokenizer();
let unk = tok.token_to_id("<unk>");
let result =
tokenize_with_default_oov(&tok, "hi 1000 world", 3, true, true, unk, &[], &Lang::En)
.expect("ok");
assert_eq!(result.token_ids.len(), 13);
let word_indices: std::collections::BTreeSet<usize> = result
.word_idx_per_token
.iter()
.filter_map(|w| *w)
.collect();
assert_eq!(word_indices.len(), 3);
}
#[test]
fn separator_token_id_is_returned() {
let tok = uppercase_tokenizer();
let unk = tok.token_to_id("<unk>");
let pipe = tok.token_to_id("|").expect("|");
let result = tokenize_with_word_map(
&tok,
"hello world",
2,
true,
true,
unk,
&[],
&Lang::En,
&[],
)
.expect("ok");
assert_eq!(result.separator_token_id, Some(pipe));
}
#[test]
fn separator_token_id_none_when_normaliser_opts_out() {
let tok = uppercase_tokenizer();
let unk = tok.token_to_id("<unk>");
let result = tokenize_with_word_map(
&tok,
"hello world",
2,
false,
true,
unk,
&[],
&Lang::En,
&[],
)
.expect("ok");
assert_eq!(result.separator_token_id, None);
}
#[test]
fn trailing_wildcards_land_after_encoded_chars() {
let tok = uppercase_tokenizer();
let unk = tok.token_to_id("<unk>");
let result = tokenize_with_default_oov(
&tok,
"hello",
1,
true,
true,
unk,
&[crate::runner::aligner::normalizer::WildcardBoundary::new(
0, 1,
)],
&Lang::En,
)
.expect("ok");
assert_eq!(result.token_ids.len(), 6);
assert_eq!(
result.token_ids[5], WILDCARD_TOKEN_ID,
"suffix wildcard must land at the END"
);
assert!(
result.token_ids[..5]
.iter()
.all(|&id| id != WILDCARD_TOKEN_ID),
"no leading wildcards expected when prefix=0; got tokens {:?}",
result.token_ids
);
assert_eq!(result.word_idx_per_token, vec![Some(0); 6]);
}
#[test]
fn leading_wildcards_land_before_encoded_chars() {
let tok = uppercase_tokenizer();
let unk = tok.token_to_id("<unk>");
let result = tokenize_with_default_oov(
&tok,
"hello",
1,
true,
true,
unk,
&[crate::runner::aligner::normalizer::WildcardBoundary::new(
1, 0,
)],
&Lang::En,
)
.expect("ok");
assert_eq!(result.token_ids.len(), 6);
assert_eq!(
result.token_ids[0], WILDCARD_TOKEN_ID,
"prefix wildcard must land at the START; got tokens {:?}",
result.token_ids
);
assert!(
result.token_ids[1..]
.iter()
.all(|&id| id != WILDCARD_TOKEN_ID),
"no trailing wildcards expected when suffix=0; got tokens {:?}",
result.token_ids
);
assert_eq!(result.word_idx_per_token, vec![Some(0); 6]);
}
#[test]
fn paired_wildcards_bracket_encoded_chars() {
let tok = uppercase_tokenizer();
let unk = tok.token_to_id("<unk>");
let result = tokenize_with_default_oov(
&tok,
"hello",
1,
true,
true,
unk,
&[crate::runner::aligner::normalizer::WildcardBoundary::new(
1, 1,
)],
&Lang::En,
)
.expect("ok");
assert_eq!(result.token_ids.len(), 7);
assert_eq!(result.token_ids[0], WILDCARD_TOKEN_ID, "prefix at start");
assert_eq!(result.token_ids[6], WILDCARD_TOKEN_ID, "suffix at end");
assert!(
result.token_ids[1..6]
.iter()
.all(|&id| id != WILDCARD_TOKEN_ID),
"interior must be encoded chars only; got {:?}",
result.token_ids
);
}
#[test]
fn wildcard_boundary_per_word_length_mismatch_errors() {
let tok = uppercase_tokenizer();
let unk = tok.token_to_id("<unk>");
let err = tokenize_with_word_map(
&tok,
"hello world",
2,
true,
true,
unk,
&[
crate::runner::aligner::normalizer::WildcardBoundary::new(1, 0),
crate::runner::aligner::normalizer::WildcardBoundary::new(2, 1),
crate::runner::aligner::normalizer::WildcardBoundary::new(3, 0),
], &Lang::En,
&[],
)
.expect_err("length mismatch must surface TokenizationFailed");
assert!(matches!(err, EmissionsError::Tokenization(_)));
}
const LETTERS: [&str; 26] = [
"A", "B", "C", "D", "E", "F", "G", "H", "I", "J", "K", "L", "M", "N", "O", "P", "Q", "R", "S",
"T", "U", "V", "W", "X", "Y", "Z",
];
fn word_level_tokenizer(alphabet: &[&str], unk_entry: bool) -> Tokenizer {
let mut vocab: Vec<String> = alphabet
.iter()
.enumerate()
.map(|(id, token)| format!("{token:?}: {id}"))
.collect();
if unk_entry {
vocab.push(format!("\"<unk>\": {}", alphabet.len()));
}
let json = format!(
r#"{{"version": "1.0", "truncation": null, "padding": null, "added_tokens": [],
"normalizer": null, "pre_tokenizer": null, "post_processor": null, "decoder": null,
"model": {{"type": "WordLevel", "vocab": {{{}}}, "unk_token": "<unk>"}}}}"#,
vocab.join(", ")
);
Tokenizer::from_bytes(json.as_bytes()).expect("a WordLevel tokenizer must parse")
}
#[test]
fn detect_oov_events_reports_what_a_vocabulary_without_unk_cannot_spell() {
let tok = word_level_tokenizer(&LETTERS, false);
assert!(
tok.encode("4", false).is_err(),
"precondition: this vocabulary cannot encode a character outside it"
);
let unk = detect_unk_token_id(&tok);
assert_eq!(unk, None);
let events = detect_oov_events(&tok, "b4d", 1, true, unk, &Lang::En, &[])
.expect("an unspellable character is an event, never an error");
assert_eq!(
events,
vec![OovEvent::new(OovKind::Symbol('4'), 1, 0, Lang::En)]
);
}
#[test]
fn tokenize_with_word_map_applies_the_decision_on_a_vocabulary_without_unk() {
let tok = word_level_tokenizer(&LETTERS, false);
let events = detect_oov_events(&tok, "b4d", 1, true, None, &Lang::En, &[]).expect("detect");
let wildcard = crate::core::wildcard_all_decisions(&events);
let tokenized =
tokenize_with_word_map(&tok, "b4d", 1, false, true, None, &[], &Lang::En, &wildcard)
.expect("a character the caller decided tokenizes");
let id_of = |token: &str| tok.token_to_id(token).expect("in the alphabet") as i32;
assert_eq!(
tokenized.token_ids(),
[id_of("B"), WILDCARD_TOKEN_ID, id_of("D")]
);
let refused = crate::core::fail_closed_all_decisions(&events);
let err = tokenize_with_word_map(&tok, "b4d", 1, false, true, None, &[], &Lang::En, &refused)
.expect_err("a refused character refuses the chunk");
assert!(
matches!(err, EmissionsError::SemanticOutOfVocab(_)),
"the policy's refusal, not a tokenization failure; got {err:?}"
);
}
#[test]
fn detect_oov_events_reports_only_what_a_letters_only_alphabet_cannot_spell() {
for unk_entry in [false, true] {
let tok = word_level_tokenizer(&LETTERS, unk_entry);
let unk = detect_unk_token_id(&tok);
let events = detect_oov_events(&tok, "Good <morning>", 2, true, unk, &Lang::En, &[])
.expect("an unspellable character is an event, never an error");
assert_eq!(
events,
vec![
OovEvent::new(OovKind::Symbol('<'), 5, 1, Lang::En),
OovEvent::new(OovKind::Symbol('>'), 13, 1, Lang::En),
],
"unk_entry = {unk_entry}"
);
}
}
fn encode_probe(tok: &Tokenizer, projected: char, unk: Option<u32>) -> Option<Vec<u32>> {
let encoding = tok
.encode(projected.to_string().as_str(), false)
.expect("a vocabulary that holds its unknown token encodes every character");
let ids = encoding.get_ids();
let unspellable = ids.is_empty() || unk.is_some_and(|unk| ids.contains(&unk));
(!unspellable).then(|| ids.to_vec())
}
fn events_before(
tok: &Tokenizer,
normalized: &str,
uppercase_input: bool,
unk: Option<u32>,
boundaries: &[WildcardBoundary],
) -> Vec<OovEvent> {
let words: Vec<&str> = normalized.split_whitespace().collect();
let mut events = Vec::new();
let mut char_index = 0;
for (word_index, word) in words.iter().enumerate() {
let boundary = boundaries
.get(word_index)
.copied()
.unwrap_or(WildcardBoundary::NONE);
for _ in 0..boundary.prefix() {
events.push(OovEvent::new(
OovKind::BoundaryPunct,
char_index,
word_index,
Lang::En,
));
}
for ch in word.chars() {
let projected = if uppercase_input {
ch.to_ascii_uppercase()
} else {
ch
};
let kind = if is_skippable_internal_punct(ch) {
Some(OovKind::InternalPunct(ch))
} else if encode_probe(tok, projected, unk).is_none() {
Some(OovKind::Symbol(ch))
} else {
None
};
if let Some(kind) = kind {
events.push(OovEvent::new(kind, char_index, word_index, Lang::En));
}
char_index += 1;
}
for _ in 0..boundary.suffix() {
events.push(OovEvent::new(
OovKind::BoundaryPunct,
char_index,
word_index,
Lang::En,
));
}
if word_index + 1 < words.len() {
char_index += 1;
}
}
events
}
#[test]
fn events_and_tokens_are_unchanged_when_the_vocabulary_holds_its_unk_token() {
let bundled = load_tokenizer_bytes_with_compat(
include_bytes!(concat!(
env!("CARGO_MANIFEST_DIR"),
"/assets/wav2vec2_base_960h_tokenizer.json"
)),
"bundled wav2vec2-base-960h",
)
.expect("the bundled tokenizer loads");
let vocabularies = [
("bundled wav2vec2-base-960h", bundled),
("uppercase", uppercase_tokenizer()),
("letters", word_level_tokenizer(&LETTERS, true)),
];
let chars: Vec<char> = ('\u{21}'..='\u{17f}')
.chain("|'.-ßẞΣςıİfi東京서울ひらがなアイ\u{301}\u{200b}\u{feff}\u{0}🎉".chars())
.filter(|c| !c.is_whitespace())
.collect();
let texts = [
"hello world",
"AT&T cost 43",
"U.S.A",
"café naïve",
"don't stop",
"Good <morning>",
"B2B 1000 ok",
"東京 서울 ひらがな",
"emoji 🎉 time",
"tab|pipe e\u{301}clair",
];
for (name, tok) in &vocabularies {
let unk = detect_unk_token_id(tok);
assert!(unk.is_some(), "{name}: holds its unknown token");
for uppercase_input in [false, true] {
for &ch in &chars {
let text = ch.to_string();
let events =
detect_oov_events(tok, &text, 1, uppercase_input, unk, &Lang::En, &[]).expect("detect");
assert_eq!(
events,
events_before(tok, &text, uppercase_input, unk, &[]),
"{name}: {ch:?}, uppercase_input = {uppercase_input}"
);
let decisions = crate::core::wildcard_all_decisions(&events);
let tokenized = tokenize_with_word_map(
tok,
&text,
1,
false,
uppercase_input,
unk,
&[],
&Lang::En,
&decisions,
)
.expect("tokenize");
let projected = if uppercase_input {
ch.to_ascii_uppercase()
} else {
ch
};
let before: Vec<i32> = match encode_probe(tok, projected, unk) {
Some(ids) if !is_skippable_internal_punct(ch) => {
ids.iter().map(|&id| id as i32).collect()
}
_ => vec![WILDCARD_TOKEN_ID],
};
assert_eq!(
tokenized.token_ids(),
before,
"{name}: {ch:?}, uppercase_input = {uppercase_input}"
);
}
for text in texts {
let word_count = text.split_whitespace().count();
for boundaries in [Vec::new(), vec![WildcardBoundary::new(1, 2); word_count]] {
let events = detect_oov_events(
tok,
text,
word_count,
uppercase_input,
unk,
&Lang::En,
&boundaries,
)
.expect("detect");
assert_eq!(
events,
events_before(tok, text, uppercase_input, unk, &boundaries),
"{name}: {text:?}, uppercase_input = {uppercase_input}"
);
}
}
}
}
}
}