use core::{
num::{NonZeroU32, NonZeroU64, NonZeroUsize},
sync::atomic::{AtomicBool, AtomicU64, Ordering},
time::Duration,
};
use std::path::Path;
use mediatime::TimeRange;
use smol_str::{SmolStr, format_smolstr};
use tokenizers::Tokenizer;
use crate::{
core::AlignmentResult,
runner::aligner::{
algorithm::{
compose::{build_speech_frames, compose_words, effective_samples_per_frame},
encode::{LogProbsTV, validate_stride_extent, validate_vocab_dim},
tokenize::{TokenizedText, detect_oov_events, tokenize_with_word_map},
trellis_beam::align_to_word_segments,
},
emissions_api::{SpeechCoverage, SpeechSpans},
normalizer::{DynTextNormalizer, NormalizationError, NormalizedText},
},
types::{AlignmentError, AlignmentFailure, Lang, WorkFailure, WorkerHangTimeout, WorkerKind},
};
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
#[error("{message}")]
pub(crate) struct AlignerCoreLoadError {
message: SmolStr,
}
impl AlignerCoreLoadError {
pub(crate) const fn new(message: SmolStr) -> Self {
Self { message }
}
pub(crate) fn message(&self) -> &SmolStr {
&self.message
}
}
pub(crate) fn detect_blank_token_id(tok: &Tokenizer) -> Option<u32> {
if let Some(id) = tok.token_to_id("<pad>") {
return Some(id);
}
if let Some(id) = tok.token_to_id("[PAD]") {
return Some(id);
}
if let Some(id) = tok.token_to_id("<blank>") {
return Some(id);
}
None
}
pub(crate) fn detect_unk_token_id(tok: &Tokenizer) -> Option<u32> {
tok
.token_to_id("<unk>")
.or_else(|| tok.token_to_id("[UNK]"))
}
pub(crate) fn detect_vocab_uppercase_only(tok: &Tokenizer) -> bool {
tok.token_to_id("A").is_some() && tok.token_to_id("a").is_none()
}
pub(crate) fn capture_vocab_size(tok: &Tokenizer) -> Option<NonZeroUsize> {
NonZeroUsize::new(tok.get_vocab_size(true))
}
pub(crate) fn validate_word_delimiter_present(
tokenizer: &Tokenizer,
use_word_delimiter: bool,
) -> Result<(), AlignerCoreLoadError> {
if !use_word_delimiter {
return Ok(());
}
if tokenizer.token_to_id("|").is_some() {
return Ok(());
}
Err(AlignerCoreLoadError::new(SmolStr::from(
"tokenizer is missing the `|` word-delimiter token, but the language's normaliser \
declared `use_word_delimiter = true`. wav2vec2 word-segmented vocabularies require \
a `|` token between spoken words. Either swap to a tokenizer that exposes `|`, or \
supply a normaliser whose `use_word_delimiter` returns false (char-level segmentation).",
)))
}
pub(crate) fn validate_decision_languages(
oov_decisions: &[crate::core::ResolvedOov],
expected_lang: &Lang,
) -> Result<(), WorkFailure> {
for (i, resolved) in oov_decisions.iter().enumerate() {
if resolved.event().language() != expected_lang {
return Err(WorkFailure::Alignment(AlignmentError::Tokenization(
AlignmentFailure::new(
format_smolstr!(
"oov_decisions[{i}].event.language = {:?} but the decisions for this chunk \
must carry {:?}. A ResolvedOov's positional identity (kind, char_index, word_index) \
deliberately ignores language, so a foreign-language decision landing at a matching \
position would silently apply ANOTHER language's wildcard / fail-closed policy. \
Recompute via `detect_oov(text)` + a policy helper from `crate::core::oov`.",
resolved.event().language(),
expected_lang,
),
expected_lang.clone(),
),
)));
}
}
Ok(())
}
pub(crate) const fn coerce_speech_coverage(value: f32) -> f32 {
if value.is_nan() {
crate::runner::aligner::algorithm::compose::DEFAULT_MIN_SPEECH_COVERAGE
} else if value < 0.0 {
0.0
} else if value > 1.0 {
1.0
} else {
value
}
}
pub(crate) fn load_tokenizer_with_compat(path: &Path) -> Result<Tokenizer, AlignerCoreLoadError> {
let bytes = std::fs::read(path).map_err(|e| {
AlignerCoreLoadError::new(format_smolstr!("read tokenizer {}: {e}", path.display()))
})?;
load_tokenizer_bytes_with_compat(&bytes, &path.display().to_string())
}
pub(crate) fn load_tokenizer_bytes_with_compat(
bytes: &[u8],
origin: &str,
) -> Result<Tokenizer, AlignerCoreLoadError> {
let original_err = match Tokenizer::from_bytes(bytes) {
Ok(tok) => return Ok(tok),
Err(e) => format_smolstr!("{e:?}"),
};
if let Some(patched) = inject_wordlevel_model_type(bytes)
&& let Ok(tok) = Tokenizer::from_bytes(&patched)
{
return Ok(tok);
}
Err(AlignerCoreLoadError::new(format_smolstr!(
"Tokenizer::from_file({origin}) failed: {original_err}"
)))
}
fn inject_wordlevel_model_type(bytes: &[u8]) -> Option<Vec<u8>> {
let _ = core::str::from_utf8(bytes).ok()?;
let model_open = find_top_level_object_value_open(bytes, b"model")?;
let model_close = find_matching_close_brace(bytes, model_open)?;
if has_top_level_key(bytes, model_open + 1, model_close, b"type") {
return None;
}
let injection = b"\n \"type\": \"WordLevel\",\n \"unk_token\": \"<unk>\",";
let mut out: Vec<u8> = Vec::with_capacity(bytes.len() + injection.len());
out.extend_from_slice(&bytes[..=model_open]);
out.extend_from_slice(injection);
out.extend_from_slice(&bytes[model_open + 1..]);
Some(out)
}
fn find_top_level_object_value_open(bytes: &[u8], key: &[u8]) -> Option<usize> {
let mut in_string = false;
let mut escape = false;
let mut depth = 0_i32;
let mut i = 0;
while i < bytes.len() {
let c = bytes[i];
if escape {
escape = false;
i += 1;
continue;
}
if in_string {
match c {
b'\\' => escape = true,
b'"' => in_string = false,
_ => {}
}
i += 1;
continue;
}
match c {
b'"' => {
let key_end = i + 1 + key.len();
if depth == 1
&& key_end < bytes.len()
&& &bytes[i + 1..key_end] == key
&& bytes[key_end] == b'"'
{
let mut j = key_end + 1;
while j < bytes.len() && (bytes[j] as char).is_ascii_whitespace() {
j += 1;
}
if j >= bytes.len() || bytes[j] != b':' {
return None;
}
j += 1;
while j < bytes.len() && (bytes[j] as char).is_ascii_whitespace() {
j += 1;
}
if j < bytes.len() && bytes[j] == b'{' {
return Some(j);
}
return None;
}
in_string = true;
}
b'{' | b'[' => depth += 1,
b'}' | b']' => depth -= 1,
_ => {}
}
i += 1;
}
None
}
fn find_matching_close_brace(bytes: &[u8], open: usize) -> Option<usize> {
if bytes.get(open) != Some(&b'{') {
return None;
}
let mut in_string = false;
let mut escape = false;
let mut depth = 1_i32;
let mut i = open + 1;
while i < bytes.len() {
let c = bytes[i];
if escape {
escape = false;
i += 1;
continue;
}
if in_string {
match c {
b'\\' => escape = true,
b'"' => in_string = false,
_ => {}
}
i += 1;
continue;
}
match c {
b'"' => in_string = true,
b'{' => depth += 1,
b'}' => {
depth -= 1;
if depth == 0 {
return Some(i);
}
}
_ => {}
}
i += 1;
}
None
}
fn has_top_level_key(bytes: &[u8], start: usize, end: usize, key: &[u8]) -> bool {
let mut in_string = false;
let mut escape = false;
let mut depth = 0_i32;
let mut i = start;
while i < end {
let c = bytes[i];
if escape {
escape = false;
i += 1;
continue;
}
if in_string {
match c {
b'\\' => escape = true,
b'"' => in_string = false,
_ => {}
}
i += 1;
continue;
}
match c {
b'"' => {
let key_end = i + 1 + key.len();
if depth == 0 && key_end < end && &bytes[i + 1..key_end] == key && bytes[key_end] == b'"' {
let mut j = key_end + 1;
while j < end && (bytes[j] as char).is_ascii_whitespace() {
j += 1;
}
if j < end && bytes[j] == b':' {
return true;
}
}
in_string = true;
}
b'{' | b'[' => depth += 1,
b'}' | b']' => depth -= 1,
_ => {}
}
i += 1;
}
false
}
#[derive(Clone, Copy, PartialEq, Eq, Debug)]
pub(crate) struct AlignerId(NonZeroU64);
impl AlignerId {
fn next() -> Self {
static COUNTER: AtomicU64 = AtomicU64::new(1);
let raw = COUNTER.fetch_add(1, Ordering::Relaxed);
Self(NonZeroU64::new(raw).expect("AlignerId counter overflowed u64"))
}
}
pub(crate) struct AlignerCore {
id: AlignerId,
tokenizer: Tokenizer,
language: Lang,
normalizer: DynTextNormalizer,
hop_samples: NonZeroU32,
blank_token_id: u32,
unk_token_id: Option<u32>,
vocab_uppercase_only: bool,
tokenizer_vocab_size: NonZeroUsize,
min_speech_coverage: SpeechCoverage,
max_intra_silent_run: Duration,
}
pub struct PreparedChunk<'a> {
owner: AlignerId,
inner: Option<PreparedInner<'a>>,
}
struct PreparedInner<'a> {
encoder_input: Vec<f32>,
real_samples: usize,
speech: SpeechSpans,
normalized: NormalizedText<'a>,
tokenized: TokenizedText,
}
impl PreparedChunk<'_> {
#[must_use]
pub fn encoder_input(&self) -> &[f32] {
self.inner.as_ref().map_or(&[], |i| &i.encoder_input)
}
#[must_use]
pub const fn is_trivial(&self) -> bool {
self.inner.is_none()
}
#[must_use]
pub fn real_samples(&self) -> usize {
self.inner.as_ref().map_or(0, |i| i.real_samples)
}
}
impl AlignerCore {
#[allow(
clippy::too_many_arguments,
reason = "one field per argument; the guards that produce them run \
in the front ends' constructors, and bundling them into a struct \
would just move the same list one level out"
)]
pub(crate) fn from_parts(
tokenizer: Tokenizer,
language: Lang,
normalizer: DynTextNormalizer,
hop_samples: NonZeroU32,
blank_token_id: u32,
unk_token_id: Option<u32>,
vocab_uppercase_only: bool,
tokenizer_vocab_size: NonZeroUsize,
min_speech_coverage: SpeechCoverage,
max_intra_silent_run: Duration,
) -> Self {
Self {
id: AlignerId::next(),
tokenizer,
language,
normalizer,
hop_samples,
blank_token_id,
unk_token_id,
vocab_uppercase_only,
tokenizer_vocab_size,
min_speech_coverage,
max_intra_silent_run,
}
}
pub(crate) fn owns(&self, prepared: &PreparedChunk<'_>) -> bool {
prepared.owner == self.id
}
pub(crate) const fn language(&self) -> &Lang {
&self.language
}
pub(crate) const fn hop_samples(&self) -> NonZeroU32 {
self.hop_samples
}
pub(crate) const fn set_hop_samples(&mut self, value: NonZeroU32) {
self.hop_samples = value;
}
pub(crate) const fn blank_token_id(&self) -> u32 {
self.blank_token_id
}
pub(crate) const fn vocab_size(&self) -> NonZeroUsize {
self.tokenizer_vocab_size
}
pub(crate) const fn min_speech_coverage(&self) -> SpeechCoverage {
self.min_speech_coverage
}
pub(crate) const fn set_min_speech_coverage(&mut self, value: SpeechCoverage) {
self.min_speech_coverage = value;
}
pub(crate) const fn max_intra_silent_run(&self) -> Duration {
self.max_intra_silent_run
}
pub(crate) const fn set_max_intra_silent_run(&mut self, value: Duration) {
self.max_intra_silent_run = value;
}
pub(crate) fn detect_oov(&self, text: &str) -> Result<Vec<crate::core::OovEvent>, WorkFailure> {
let normalized = match self.normalizer.normalize(text) {
Ok(n) => n,
Err(NormalizationError::EmptyText) => {
return Ok(Vec::new());
}
Err(e) => {
return Err(WorkFailure::Alignment(AlignmentError::Normalization(
AlignmentFailure::new(
format_smolstr!("normalize failed: {e}"),
self.language.clone(),
),
)));
}
};
let n_words = normalized.normalized().split_whitespace().count();
detect_oov_events(
&self.tokenizer,
normalized.normalized(),
n_words,
self.vocab_uppercase_only,
self.unk_token_id,
&self.language,
normalized.wildcard_boundary_per_word(),
)
.map_err(|e| e.into_work_failure(&self.language))
}
pub(crate) fn prepare<'a>(
&self,
samples: &[f32],
speech: &SpeechSpans,
text: &'a str,
oov_decisions: &[crate::core::ResolvedOov],
expected_decision_language: &Lang,
abort_flag: &AtomicBool,
) -> Result<PreparedChunk<'a>, WorkFailure> {
validate_decision_languages(oov_decisions, expected_decision_language)?;
if abort_flag.load(Ordering::Relaxed) {
return Err(timed_out());
}
if let Some((idx, val)) = samples
.iter()
.copied()
.enumerate()
.find(|(_, s)| !s.is_finite())
{
return Err(WorkFailure::Alignment(AlignmentError::ModelInference(
AlignmentFailure::new(
format_smolstr!(
"non-finite sample at index {idx} (value {val:?}); upstream audio corruption — \
refuse to encode, masking-as-silence would only hide the bug"
),
self.language.clone(),
),
)));
}
let speech_mask = build_speech_mask(samples.len(), speech);
if abort_flag.load(Ordering::Relaxed) {
return Err(timed_out());
}
let normalized = match self.normalizer.normalize(text) {
Ok(nt) => nt,
Err(NormalizationError::EmptyText) => {
return Ok(PreparedChunk {
owner: self.id,
inner: None,
});
}
Err(NormalizationError::RuleFailed { detail }) => {
return Err(WorkFailure::Alignment(AlignmentError::Normalization(
AlignmentFailure::new(detail, self.language.clone()),
)));
}
};
let n_words = normalized.original_words().len();
if abort_flag.load(Ordering::Relaxed) {
return Err(timed_out());
}
let tokenized = tokenize_with_word_map(
&self.tokenizer,
normalized.normalized(),
n_words,
self.normalizer.use_word_delimiter(),
self.vocab_uppercase_only,
self.unk_token_id,
normalized.wildcard_boundary_per_word(),
&self.language,
oov_decisions,
)
.map_err(|e| e.into_work_failure(&self.language))?;
if tokenized.token_ids().is_empty() {
return Ok(PreparedChunk {
owner: self.id,
inner: None,
});
}
if abort_flag.load(Ordering::Relaxed) {
return Err(timed_out());
}
let normalized_samples: Vec<f32> = samples
.iter()
.zip(speech_mask.iter())
.map(|(&s, &is_speech)| if is_speech { s } else { 0.0_f32 })
.collect();
let encoder_input: Vec<f32> = if normalized_samples.len() < 400 {
let mut buf = Vec::with_capacity(400);
buf.extend_from_slice(&normalized_samples);
buf.resize(400, 0.0_f32);
buf
} else {
normalized_samples
};
Ok(PreparedChunk {
owner: self.id,
inner: Some(PreparedInner {
encoder_input,
real_samples: samples.len(),
speech: speech.clone(),
normalized,
tokenized,
}),
})
}
pub(crate) fn finish<F>(
&self,
prepared: PreparedChunk<'_>,
log_probs: &LogProbsTV,
chunk_first_sample_in_stream: u64,
samples_to_output_range: F,
abort_flag: &AtomicBool,
) -> Result<AlignmentResult, WorkFailure>
where
F: Fn(u64, u64) -> TimeRange,
{
if !self.owns(&prepared) {
return Err(WorkFailure::Alignment(AlignmentError::ModelInference(
AlignmentFailure::new(
format_smolstr!(
"PreparedChunk was produced by a different aligner (prepared by aligner \
{:?}, finished on aligner {:?}). A PreparedChunk carries token ids, a word map, and \
OOV decisions resolved against ITS aligner's tokenizer, blank id, and language; \
applying them to another aligner's emissions reads posteriors from columns that do \
not correspond to those tokens — a believable but incorrect alignment. Call `finish` \
on the same aligner that called `prepare`.",
prepared.owner,
self.id,
),
self.language.clone(),
),
)));
}
let Some(prepared) = prepared.inner else {
return Ok(AlignmentResult::new(Vec::new()));
};
let tokenized = &prepared.tokenized;
#[cfg(feature = "parity-dump-emission")]
{
use core::sync::atomic::AtomicUsize;
static SEG_COUNTER: AtomicUsize = AtomicUsize::new(0);
if let Ok(dir) = std::env::var("ASRY_PARITY_DUMP_TRELLIS") {
let n = SEG_COUNTER.fetch_add(1, Ordering::Relaxed);
let dir_path = std::path::PathBuf::from(dir);
let _ = std::fs::create_dir_all(&dir_path);
let em_path = dir_path.join(format!("wy_seg{n}.emission.bin"));
if let Ok(mut f) = std::fs::File::create(&em_path) {
use std::io::Write;
let _ = f.write_all(&(log_probs.t() as u32).to_le_bytes());
let _ = f.write_all(&(log_probs.v() as u32).to_le_bytes());
let mut buf: Vec<u8> = Vec::with_capacity(log_probs.data().len() * 4);
for v in log_probs.data() {
buf.extend_from_slice(&v.to_le_bytes());
}
let _ = f.write_all(&buf);
}
let tok_path = dir_path.join(format!("wy_seg{n}.tokens.json"));
if let Ok(mut f) = std::fs::File::create(&tok_path) {
use std::io::Write;
let mut payload = format!("{{\"blank_id\":{},\"tokens\":[", self.blank_token_id);
for (i, t) in tokenized.token_ids().iter().enumerate() {
if i > 0 {
payload.push(',');
}
payload.push_str(&format!("{t}"));
}
payload.push_str(&format!(
"],\"n_samples\":{},\"T\":{},\"V\":{}}}",
prepared.encoder_input.len(),
log_probs.t(),
log_probs.v()
));
let _ = f.write_all(payload.as_bytes());
}
}
}
validate_stride_extent(
log_probs.t(),
self.hop_samples.get(),
prepared.real_samples,
&self.language,
)?;
validate_vocab_dim(
log_probs.v(),
self.tokenizer_vocab_size.get(),
&self.language,
)?;
if abort_flag.load(Ordering::Relaxed) {
return Err(timed_out());
}
let word_segments = align_to_word_segments(
log_probs,
tokenized.token_ids(),
tokenized.word_idx_per_token(),
tokenized.separator_token_id(),
self.blank_token_id,
abort_flag,
&self.language,
)?;
#[cfg(feature = "parity-dump-emission")]
{
use core::sync::atomic::AtomicUsize;
static TRELLIS_COUNTER: AtomicUsize = AtomicUsize::new(0);
if let Ok(dir) = std::env::var("ASRY_PARITY_DUMP_TRELLIS") {
let n = TRELLIS_COUNTER.fetch_add(1, Ordering::Relaxed);
let dir_path = std::path::PathBuf::from(dir);
let trellis = crate::runner::aligner::algorithm::trellis_beam::get_trellis(
log_probs,
tokenized.token_ids(),
self.blank_token_id,
abort_flag,
&self.language,
);
if let Ok(trellis) = trellis {
let path = dir_path.join(format!("wy_seg{n}.trellis.bin"));
if let Ok(mut f) = std::fs::File::create(&path) {
use std::io::Write;
let _ = f.write_all(&(log_probs.t() as u32).to_le_bytes());
let _ = f.write_all(&(tokenized.token_ids().len() as u32).to_le_bytes());
let mut buf: Vec<u8> = Vec::with_capacity(trellis.len() * 4);
for v in &trellis {
buf.extend_from_slice(&v.to_le_bytes());
}
let _ = f.write_all(&buf);
}
}
}
}
if abort_flag.load(Ordering::Relaxed) {
return Err(timed_out());
}
let encoder_n_samples = prepared.encoder_input.len() as u64;
let samples_per_frame =
effective_samples_per_frame(encoder_n_samples, log_probs.t(), self.hop_samples.get());
let real_n_samples = prepared.real_samples as u64;
let speech_frames = build_speech_frames(
log_probs.t(),
samples_per_frame,
encoder_n_samples,
real_n_samples,
&prepared.speech,
);
Ok(compose_words(
&word_segments,
prepared.normalized.original_words(),
&speech_frames,
chunk_first_sample_in_stream,
self.hop_samples.get(),
encoder_n_samples,
real_n_samples,
log_probs.t(),
samples_to_output_range,
self.min_speech_coverage,
self.max_intra_silent_run,
))
}
}
fn timed_out() -> WorkFailure {
WorkFailure::WorkerHang(WorkerHangTimeout::new(
WorkerKind::Alignment,
Duration::ZERO,
))
}
pub(crate) fn build_speech_mask(n_samples: usize, speech: &SpeechSpans) -> Vec<bool> {
let mut mask = vec![false; n_samples];
let n_samples_u64 = n_samples as u64;
for span in speech.as_slice() {
let start = span.start().min(n_samples_u64) as usize;
let end = span.end().min(n_samples_u64) as usize;
if end > start {
for slot in &mut mask[start..end] {
*slot = true;
}
}
}
mask
}
#[cfg(test)]
mod tests {
use super::*;
use crate::time::SAMPLE_RATE_HZ;
#[test]
fn load_tokenizer_with_compat_handles_unpatched_hf_format() {
let raw = br#"{
"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": {
"vocab": {
"<pad>": 0, "<s>": 1, "</s>": 2, "<unk>": 3, "|": 4,
"A": 5, "B": 6, "C": 7
}
}
}"#;
assert!(
Tokenizer::from_bytes(raw).is_err(),
"tokenizers 0.20 unexpectedly accepted raw upstream HF format; \
the compat shim is no longer necessary"
);
let patched =
inject_wordlevel_model_type(raw).expect("inject_wordlevel_model_type must succeed");
let tok = Tokenizer::from_bytes(&patched).expect("patched JSON must parse");
assert_eq!(tok.token_to_id("A"), Some(5));
assert_eq!(tok.token_to_id("<unk>"), Some(3));
}
#[test]
fn load_tokenizer_with_compat_skips_already_patched_input() {
let already_typed = br#"{
"model": {
"type": "WordLevel",
"vocab": {"<unk>": 0, "A": 1},
"unk_token": "<unk>"
}
}"#;
assert!(inject_wordlevel_model_type(already_typed).is_none());
}
#[test]
fn inject_wordlevel_model_type_ignores_model_substring_inside_strings() {
let raw = br#"{
"decoy": "this string mentions \"model\" with escape-quoted braces",
"model": {
"vocab": {"<pad>": 0, "<unk>": 1, "|": 2, "A": 3}
}
}"#;
let patched = inject_wordlevel_model_type(raw)
.expect("patcher must locate the real top-level model key, not the decoy substring");
let s = core::str::from_utf8(&patched).expect("UTF-8");
let inj = s
.find("\"type\": \"WordLevel\"")
.expect("patched output must contain injected discriminator");
let real_model_key = s
.find("\n \"model\": {")
.expect("real model key must remain in output");
assert!(
inj > real_model_key,
"injection at offset {inj} must come AFTER real model key at offset {real_model_key}; \
the decoy substring would have placed it earlier"
);
}
#[test]
fn inject_wordlevel_model_type_ignores_braces_inside_strings() {
let raw = br#"{
"decoy": "value with { braces } and more { } inside",
"model": {
"vocab": {"<pad>": 0, "<unk>": 1, "|": 2, "B": 3}
}
}"#;
let patched = inject_wordlevel_model_type(raw)
.expect("patcher must skip braces inside string values when finding model body close");
let s = core::str::from_utf8(&patched).expect("UTF-8");
assert!(
s.contains("\"type\": \"WordLevel\""),
"patched output must contain injected discriminator"
);
assert!(
s.contains("\"decoy\": \"value with { braces } and more { } inside\""),
"decoy field must remain byte-identical"
);
}
#[test]
fn inject_wordlevel_model_type_does_not_treat_quoted_type_as_discriminator() {
let raw = br#"{
"model": {
"_note": "the type of model is wav2vec2",
"vocab": {"<pad>": 0, "<unk>": 1, "|": 2, "C": 3}
}
}"#;
let patched = inject_wordlevel_model_type(raw).expect(
"patcher must NOT short-circuit on a quoted `type` substring inside a string value; \
it must inject the real discriminator key",
);
let s = core::str::from_utf8(&patched).expect("UTF-8");
assert!(
s.contains("\"type\": \"WordLevel\""),
"patched output must contain the injected discriminator key"
);
}
#[test]
fn coerce_speech_coverage_passes_through_valid_values() {
assert_eq!(coerce_speech_coverage(0.0), 0.0);
assert_eq!(coerce_speech_coverage(0.25), 0.25);
assert_eq!(coerce_speech_coverage(0.5), 0.5);
assert_eq!(coerce_speech_coverage(0.99), 0.99);
assert_eq!(coerce_speech_coverage(1.0), 1.0);
}
#[test]
fn coerce_speech_coverage_clamps_above_one() {
assert_eq!(coerce_speech_coverage(1.5), 1.0);
assert_eq!(coerce_speech_coverage(100.0), 1.0);
assert_eq!(coerce_speech_coverage(f32::INFINITY), 1.0);
}
#[test]
fn coerce_speech_coverage_clamps_below_zero() {
assert_eq!(coerce_speech_coverage(-0.1), 0.0);
assert_eq!(coerce_speech_coverage(-100.0), 0.0);
assert_eq!(coerce_speech_coverage(f32::NEG_INFINITY), 0.0);
}
#[test]
fn coerce_speech_coverage_treats_nan_as_default_in_both_profiles() {
assert_eq!(
coerce_speech_coverage(f32::NAN),
crate::runner::aligner::algorithm::compose::DEFAULT_MIN_SPEECH_COVERAGE,
"NaN must coerce to the default without panicking, in debug and release alike"
);
}
#[test]
fn validate_direct_decision_languages_rejects_cross_language_payload() {
use crate::core::{OovDecision, OovEvent, OovKind, ResolvedOov};
let stale = vec![ResolvedOov::new(
OovEvent::new(OovKind::Symbol('&'), 2, 0, Lang::Ko),
OovDecision::Wildcard,
)];
let result = validate_decision_languages(&stale, &Lang::En);
match result {
Err(WorkFailure::Alignment(AlignmentError::Tokenization(payload))) => assert!(
payload
.message()
.contains("oov_decisions[0].event.language")
&& payload.message().contains("Ko")
&& payload.message().contains("En"),
"diagnostic should cite the offending index + the languages; got {message}",
message = payload.message(),
),
other => panic!("expected TokenizationFailed cross-language; got {other:?}"),
}
}
#[test]
fn validate_direct_decision_languages_accepts_matching_payload() {
use crate::core::{OovDecision, OovEvent, OovKind, ResolvedOov};
let ok = vec![ResolvedOov::new(
OovEvent::new(OovKind::Symbol('&'), 2, 0, Lang::En),
OovDecision::Wildcard,
)];
assert!(validate_decision_languages(&ok, &Lang::En).is_ok());
}
#[test]
fn validate_direct_decision_languages_accepts_empty() {
assert!(validate_decision_languages(&[], &Lang::En).is_ok());
}
#[test]
fn validate_decision_languages_accepts_a_fallback_payload_keyed_on_the_request() {
use crate::core::{OovDecision, OovEvent, OovKind, ResolvedOov};
let decisions = vec![ResolvedOov::new(
OovEvent::new(OovKind::Symbol('&'), 2, 0, Lang::Fr),
OovDecision::Wildcard,
)];
assert!(validate_decision_languages(&decisions, &Lang::Fr).is_ok());
}
fn tokenizer_with_pipe_delimiter() -> Tokenizer {
let json = 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},
"unk_token": "<unk>"
}
}"#;
Tokenizer::from_bytes(json.as_bytes()).expect("parse")
}
fn tokenizer_without_pipe_delimiter() -> Tokenizer {
let json = 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, "A": 2, "B": 3},
"unk_token": "<unk>"
}
}"#;
Tokenizer::from_bytes(json.as_bytes()).expect("parse")
}
#[test]
fn delimiter_check_passes_when_token_present_and_required() {
let tok = tokenizer_with_pipe_delimiter();
assert!(validate_word_delimiter_present(&tok, true).is_ok());
}
#[test]
fn delimiter_check_fails_when_required_but_missing() {
let tok = tokenizer_without_pipe_delimiter();
let err = validate_word_delimiter_present(&tok, true).unwrap_err();
let message = err.message();
assert!(
message.contains("`|` word-delimiter"),
"must call out the missing delimiter; got {message}"
);
}
#[test]
fn delimiter_check_passes_for_char_segmented_normalizers() {
let tok = tokenizer_without_pipe_delimiter();
assert!(validate_word_delimiter_present(&tok, false).is_ok());
}
fn tokenizer_kresnik_shape() -> Tokenizer {
let json = 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": {"안": 0, "녕": 1, "하": 2, "세": 3, "요": 4, "|": 859, "[UNK]": 1203, "[PAD]": 1204},
"unk_token": "[UNK]"
}
}"#;
Tokenizer::from_bytes(json.as_bytes()).expect("parse")
}
#[test]
fn detect_blank_token_id_resolves_bracket_pad_at_high_index() {
let tok = tokenizer_kresnik_shape();
assert_eq!(detect_blank_token_id(&tok), Some(1204));
}
#[test]
fn unk_fallback_resolves_bracket_unk() {
let tok = tokenizer_kresnik_shape();
let unk = tok
.token_to_id("<unk>")
.or_else(|| tok.token_to_id("[UNK]"));
assert_eq!(unk, Some(1203));
}
#[test]
fn detect_unk_token_id_resolves_bracket_unk() {
let tok = tokenizer_kresnik_shape();
assert_eq!(detect_unk_token_id(&tok), Some(1203));
}
#[test]
fn delimiter_check_for_korean_normalizer_passes_even_with_pipe_present() {
let tok = tokenizer_kresnik_shape();
assert!(validate_word_delimiter_present(&tok, false).is_ok());
}
#[test]
fn detect_vocab_uppercase_only_probes_the_case_convention() {
assert!(detect_vocab_uppercase_only(&tokenizer_with_pipe_delimiter()));
assert!(!detect_vocab_uppercase_only(&tokenizer_kresnik_shape()));
}
#[test]
fn capture_vocab_size_is_nonzero_for_a_real_vocab() {
let tok = tokenizer_with_pipe_delimiter();
let v = capture_vocab_size(&tok).expect("a vocab with 5 entries is non-zero");
assert_eq!(v.get(), tok.get_vocab_size(true));
assert_eq!(v.get(), 5);
}
#[test]
fn load_tokenizer_bytes_with_compat_patches_unpatched_hf_format() {
let raw = br#"{
"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": {
"vocab": {
"<pad>": 0, "<s>": 1, "</s>": 2, "<unk>": 3, "|": 4,
"A": 5, "B": 6, "C": 7
}
}
}"#;
let tok = load_tokenizer_bytes_with_compat(raw, "<test>").expect("compat shim must patch");
assert_eq!(tok.token_to_id("A"), Some(5));
assert_eq!(detect_blank_token_id(&tok), Some(0));
assert_eq!(detect_unk_token_id(&tok), Some(3));
}
#[test]
fn load_tokenizer_bytes_with_compat_rejects_garbage() {
let err = load_tokenizer_bytes_with_compat(b"not json at all", "tokenizer.json")
.expect_err("garbage must not parse");
assert!(
err.message().contains("tokenizer.json"),
"diagnostic must name the origin; got {}",
err.message()
);
}
fn analysis_tb() -> mediatime::Timebase {
mediatime::Timebase::new(1, core::num::NonZeroU32::new(SAMPLE_RATE_HZ).unwrap())
}
fn spans(ranges: &[TimeRange]) -> SpeechSpans {
SpeechSpans::from_time_ranges(ranges).expect("test ranges are in the analysis timebase")
}
#[test]
fn build_speech_mask_marks_inrange_segments() {
let segs = spans(&[TimeRange::new(2, 5, analysis_tb())]);
let mask = build_speech_mask(8, &segs);
assert_eq!(
mask,
vec![false, false, true, true, true, false, false, false]
);
}
#[test]
fn build_speech_mask_clamps_negative_overlap_to_zero() {
let segs = spans(&[TimeRange::new(-3, 4, analysis_tb())]);
let mask = build_speech_mask(8, &segs);
assert_eq!(
mask,
vec![true, true, true, true, false, false, false, false]
);
}
#[test]
fn build_speech_mask_clamps_overshoot_to_buffer_end() {
let segs = spans(&[TimeRange::new(5, 100, analysis_tb())]);
let mask = build_speech_mask(8, &segs);
assert_eq!(
mask,
vec![false, false, false, false, false, true, true, true]
);
}
#[test]
fn build_speech_mask_drops_fully_negative_range() {
let segs = spans(&[TimeRange::new(-10, -3, analysis_tb())]);
let mask = build_speech_mask(8, &segs);
assert_eq!(mask, vec![false; 8]);
}
#[test]
fn build_speech_mask_drops_fully_overshoot_range() {
let segs = spans(&[TimeRange::new(20, 30, analysis_tb())]);
let mask = build_speech_mask(8, &segs);
assert_eq!(mask, vec![false; 8]);
}
#[test]
fn build_speech_mask_zero_width_range_is_dropped() {
let segs = spans(&[TimeRange::new(5, 5, analysis_tb())]);
let mask = build_speech_mask(8, &segs);
assert_eq!(mask, vec![false; 8]);
}
#[test]
fn build_speech_mask_unions_overlapping_segments() {
let segs = spans(&[
TimeRange::new(1, 4, analysis_tb()),
TimeRange::new(3, 6, analysis_tb()),
]);
let mask = build_speech_mask(8, &segs);
assert_eq!(
mask,
vec![false, true, true, true, true, true, false, false]
);
}
#[test]
fn build_speech_mask_empty_buffer_returns_empty_mask() {
let segs = spans(&[TimeRange::new(0, 0, analysis_tb())]);
let mask = build_speech_mask(0, &segs);
assert!(mask.is_empty());
}
#[test]
fn build_speech_mask_distinguishes_no_vad_from_all_silence() {
let silence = build_speech_mask(8, &SpeechSpans::new([]));
assert_eq!(silence, vec![false; 8], "an empty span list IS all silence");
let all = build_speech_mask(8, &SpeechSpans::all_speech());
assert_eq!(all, vec![true; 8], "all_speech() covers the whole chunk");
}
}