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 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);
let encoding = tokenizer
.encode(tmp_buf.as_str(), false)
.map_err(|e| {
EmissionsError::Tokenization(EmissionsFailure::new(format_smolstr!(
"encode({:?}) failed: {e:?}",
projected
)))
})?;
let ids = encoding.get_ids();
let is_unk_or_empty = ids.is_empty()
|| match unk_token_id {
Some(unk) => ids.contains(&unk),
None => false,
};
if is_unk_or_empty {
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 encoding = tokenizer
.encode(tmp_buf.as_str(), false)
.map_err(|e| {
EmissionsError::Tokenization(EmissionsFailure::new(format_smolstr!(
"encode({:?}) failed: {e:?}",
projected
)))
})?;
let ids = encoding.get_ids();
let is_unk_or_empty = ids.is_empty()
|| match unk_token_id {
Some(unk) => ids.contains(&unk),
None => false,
};
if is_unk_or_empty {
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."
),
)));
}
}
} else {
for &id in ids {
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::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(_)));
}
}