use core::sync::atomic::{AtomicBool, Ordering};
use smol_str::{SmolStr, format_smolstr};
use crate::{
runner::aligner::algorithm::{
encode::LogProbsTV,
errors::{EmissionsError, EmissionsFailure},
tokenize::TokenizedText,
},
types::{AlignmentError, AlignmentFailure, Lang, WorkFailure, WorkerHangTimeout, WorkerKind},
};
pub const WILDCARD_TOKEN_ID: i32 = -1;
pub const ALIGN_BEAM_WIDTH: usize = 2;
const TRELLIS_CELL_BUDGET: usize = 32_000_000;
const BEAM_NODE_BUDGET: usize = 2_000_000;
pub(crate) const SEAM_PATH_FRAME_BUDGET: usize = 2_000_000;
#[derive(Debug, Clone)]
pub(crate) struct CharSegment {
pub token_index: usize,
pub start_frame: usize,
pub end_frame: usize,
pub score: f32,
}
impl CharSegment {
pub(crate) const fn length(&self) -> usize {
self.end_frame - self.start_frame
}
}
#[derive(Debug, Clone)]
pub struct WordSegment {
word_index: usize,
start_frame: usize,
end_frame: usize,
score: f32,
}
impl WordSegment {
#[must_use]
pub const fn new(word_index: usize, start_frame: usize, end_frame: usize, score: f32) -> Self {
Self {
word_index,
start_frame,
end_frame,
score,
}
}
#[must_use]
pub const fn word_index(&self) -> usize {
self.word_index
}
#[must_use]
pub const fn start_frame(&self) -> usize {
self.start_frame
}
#[must_use]
pub const fn end_frame(&self) -> usize {
self.end_frame
}
#[must_use]
pub const fn score(&self) -> f32 {
self.score
}
}
pub fn get_trellis(
log_probs: &LogProbsTV,
tokens: &[i32],
blank_id: u32,
abort_flag: &AtomicBool,
language: &Lang,
) -> Result<Vec<f32>, WorkFailure> {
let t = log_probs.t();
let num_tokens = tokens.len();
if num_tokens == 0 {
return Err(WorkFailure::Alignment(AlignmentError::NoAlignmentPath(
AlignmentFailure::new(SmolStr::from("token sequence is empty"), language.clone()),
)));
}
let v = log_probs.v();
if (blank_id as usize) >= v {
return Err(WorkFailure::Alignment(AlignmentError::ModelInference(
AlignmentFailure::new(
format_smolstr!(
"blank token id {blank_id} >= model output vocab dim {v}; tokenizer/model mismatch?"
),
language.clone(),
),
)));
}
for (i, &tok) in tokens.iter().enumerate() {
if tok == WILDCARD_TOKEN_ID {
continue;
}
if tok < 0 {
return Err(WorkFailure::Alignment(AlignmentError::Tokenization(
AlignmentFailure::new(
format_smolstr!(
"token id {tok} at position {i} is negative (only the wildcard \
sentinel {WILDCARD_TOKEN_ID} is allowed); tokenizer bug?"
),
language.clone(),
),
)));
}
if (tok as usize) >= v {
return Err(WorkFailure::Alignment(AlignmentError::Tokenization(
AlignmentFailure::new(
format_smolstr!(
"token id {tok} at position {i} >= model output vocab dim {v}; \
tokenizer/model mismatch?"
),
language.clone(),
),
)));
}
}
if t < num_tokens {
return Err(WorkFailure::Alignment(AlignmentError::NoAlignmentPath(
AlignmentFailure::new(
format_smolstr!(
"audio too short: T={} frames < {} chars; trellis is degenerate",
t,
num_tokens
),
language.clone(),
),
)));
}
let cells = match t.checked_mul(num_tokens) {
Some(v) => v,
None => {
return Err(WorkFailure::Alignment(AlignmentError::NoAlignmentPath(
AlignmentFailure::new(
format_smolstr!("trellis size overflows usize: T={t} * num_tokens={num_tokens}"),
language.clone(),
),
)));
}
};
if cells > TRELLIS_CELL_BUDGET {
return Err(WorkFailure::Alignment(AlignmentError::NoAlignmentPath(
AlignmentFailure::new(
format_smolstr!(
"trellis exceeds {} cells (T={} × num_tokens={} = {})",
TRELLIS_CELL_BUDGET,
t,
num_tokens,
cells
),
language.clone(),
),
)));
}
if abort_flag.load(Ordering::Relaxed) {
return Err(WorkFailure::WorkerHang(WorkerHangTimeout::new(
WorkerKind::Alignment,
core::time::Duration::ZERO,
)));
}
let mut trellis = vec![0.0_f32; cells];
for j in 1..num_tokens {
trellis[j] = f32::NEG_INFINITY;
}
let mut acc = 0.0_f32;
for ti in 1..t {
acc += log_probs.at(ti, blank_id as usize);
trellis[ti * num_tokens] = acc;
}
if num_tokens >= 2 {
let row_start = t.saturating_sub(num_tokens - 1);
for ti in row_start..t {
trellis[ti * num_tokens] = f32::INFINITY;
}
}
for t_idx in 0..t.saturating_sub(1) {
if t_idx % 64 == 0 && abort_flag.load(Ordering::Relaxed) {
return Err(WorkFailure::WorkerHang(WorkerHangTimeout::new(
WorkerKind::Alignment,
core::time::Duration::ZERO,
)));
}
let blank_emit = log_probs.at(t_idx, blank_id as usize);
let wildcard_emit_for_frame = max_non_blank_logprob(log_probs, t_idx, blank_id as usize);
for j in 1..num_tokens {
let stay = trellis[t_idx * num_tokens + j] + blank_emit;
let prev = trellis[t_idx * num_tokens + (j - 1)];
let change_emit = match tokens[j] {
id if id == WILDCARD_TOKEN_ID => wildcard_emit_for_frame,
id => log_probs.at(t_idx, id as usize),
};
let change = prev + change_emit;
trellis[(t_idx + 1) * num_tokens + j] = if stay >= change { stay } else { change };
}
}
Ok(trellis)
}
fn max_non_blank_logprob(log_probs: &LogProbsTV, t_idx: usize, blank_v: usize) -> f32 {
let row_start = t_idx * log_probs.v();
let mut best = f32::NEG_INFINITY;
for v in 0..log_probs.v() {
if v == blank_v {
continue;
}
let lp = log_probs.data()[row_start + v];
if lp > best {
best = lp;
}
}
best
}
#[derive(Debug, Clone)]
struct BeamNode {
token_index: usize,
time_index: usize,
score: f32,
point_score: f32,
prev: Option<u32>,
}
pub fn backtrack_beam(
trellis: &[f32],
log_probs: &LogProbsTV,
tokens: &[i32],
blank_id: u32,
beam_width: usize,
abort_flag: &AtomicBool,
language: &Lang,
) -> Result<Vec<PathPointPublic>, WorkFailure> {
let t = log_probs.t();
let num_tokens = tokens.len();
if num_tokens == 0 {
return Err(WorkFailure::Alignment(AlignmentError::NoAlignmentPath(
AlignmentFailure::new(SmolStr::from("token sequence is empty"), language.clone()),
)));
}
if t == 0 {
return Err(WorkFailure::Alignment(AlignmentError::NoAlignmentPath(
AlignmentFailure::new(SmolStr::from("emission has zero frames"), language.clone()),
)));
}
let final_t = t - 1;
let final_j = num_tokens - 1;
let final_score = trellis[final_t * num_tokens + final_j];
if !final_score.is_finite() {
return Err(WorkFailure::Alignment(AlignmentError::NoAlignmentPath(
AlignmentFailure::new(
format_smolstr!(
"trellis end cell at (t={}, j={}) is non-finite ({}); no path to backtrack",
final_t,
final_j,
final_score
),
language.clone(),
),
)));
}
let mut arena: Vec<BeamNode> = Vec::new();
arena.push(BeamNode {
token_index: final_j,
time_index: final_t,
score: final_score,
point_score: log_probs.at(final_t, blank_id as usize).exp(),
prev: None,
});
let mut active: Vec<u32> = vec![0_u32];
let mut next_active: Vec<u32> = Vec::with_capacity(beam_width * 2);
let mut iters = 0_usize;
while !active.is_empty() && arena[active[0] as usize].token_index > 0 {
iters += 1;
if iters.is_multiple_of(64) && abort_flag.load(Ordering::Relaxed) {
return Err(WorkFailure::WorkerHang(WorkerHangTimeout::new(
WorkerKind::Alignment,
core::time::Duration::ZERO,
)));
}
next_active.clear();
for &beam_idx in &active {
let (t_curr, j_curr) = {
let beam = &arena[beam_idx as usize];
(beam.time_index, beam.token_index)
};
if t_curr == 0 {
continue;
}
let p_stay_lp = log_probs.at(t_curr - 1, blank_id as usize);
let p_change_lp = match tokens[j_curr] {
id if id == WILDCARD_TOKEN_ID => {
max_non_blank_logprob(log_probs, t_curr - 1, blank_id as usize)
}
id => log_probs.at(t_curr - 1, id as usize),
};
let stay_score = trellis[(t_curr - 1) * num_tokens + j_curr];
let change_score = if j_curr > 0 {
trellis[(t_curr - 1) * num_tokens + (j_curr - 1)]
} else {
f32::NEG_INFINITY
};
if stay_score.is_finite() {
if arena.len() >= BEAM_NODE_BUDGET {
return Err(WorkFailure::Alignment(AlignmentError::NoAlignmentPath(
AlignmentFailure::new(
format_smolstr!(
"beam arena exceeded {BEAM_NODE_BUDGET} nodes; lattice likely degenerate \
(high T, very few tokens). Aborting backtrack to bound memory."
),
language.clone(),
),
)));
}
let new_idx = arena.len() as u32;
arena.push(BeamNode {
token_index: j_curr,
time_index: t_curr - 1,
score: stay_score,
point_score: p_stay_lp.exp(),
prev: Some(beam_idx),
});
next_active.push(new_idx);
}
if j_curr > 0 && change_score.is_finite() {
if arena.len() >= BEAM_NODE_BUDGET {
return Err(WorkFailure::Alignment(AlignmentError::NoAlignmentPath(
AlignmentFailure::new(
format_smolstr!(
"beam arena exceeded {BEAM_NODE_BUDGET} nodes (change branch); lattice \
likely degenerate. Aborting backtrack to bound memory."
),
language.clone(),
),
)));
}
let new_idx = arena.len() as u32;
arena.push(BeamNode {
token_index: j_curr - 1,
time_index: t_curr - 1,
score: change_score,
point_score: p_change_lp.exp(),
prev: Some(beam_idx),
});
next_active.push(new_idx);
}
}
next_active.sort_by(|&a, &b| arena[b as usize].score.total_cmp(&arena[a as usize].score));
if next_active.len() > beam_width {
next_active.truncate(beam_width);
}
core::mem::swap(&mut active, &mut next_active);
}
if active.is_empty() {
return Err(WorkFailure::Alignment(AlignmentError::NoAlignmentPath(
AlignmentFailure::new(
SmolStr::from("beam search emptied before reaching token 0"),
language.clone(),
),
)));
}
let winner_idx = active[0] as usize;
let winner_t = arena[winner_idx].time_index;
let winner_token = arena[winner_idx].token_index;
let mut path: Vec<PathPointPublic> = Vec::with_capacity(t);
for ti in 0..winner_t {
let prob = log_probs.at(ti, blank_id as usize).exp();
path.push(PathPointPublic {
token_index: winner_token,
time_index: ti,
score: prob,
});
}
let mut cur: Option<u32> = Some(active[0]);
while let Some(idx) = cur {
let node = &arena[idx as usize];
path.push(PathPointPublic {
token_index: node.token_index,
time_index: node.time_index,
score: node.point_score,
});
cur = node.prev;
}
Ok(path)
}
#[derive(Debug, Clone, PartialEq)]
pub struct PathPointPublic {
token_index: usize,
time_index: usize,
score: f32,
}
impl PathPointPublic {
#[must_use]
pub const fn new(token_index: usize, time_index: usize, score: f32) -> Self {
Self {
token_index,
time_index,
score,
}
}
#[must_use]
pub const fn token_index(&self) -> usize {
self.token_index
}
#[must_use]
pub const fn time_index(&self) -> usize {
self.time_index
}
#[must_use]
pub const fn score(&self) -> f32 {
self.score
}
}
pub(crate) fn merge_repeats(path: &[PathPointPublic]) -> Vec<CharSegment> {
let mut segments: Vec<CharSegment> = Vec::new();
if path.is_empty() {
return segments;
}
let mut i1 = 0;
while i1 < path.len() {
let mut i2 = i1;
while i2 < path.len() && path[i1].token_index == path[i2].token_index {
i2 += 1;
}
let n = (i2 - i1) as f32;
let mut score_sum = 0.0_f32;
for k in i1..i2 {
score_sum += path[k].score;
}
let score = if n > 0.0 { score_sum / n } else { 0.0 };
segments.push(CharSegment {
token_index: path[i1].token_index,
start_frame: path[i1].time_index,
end_frame: path[i2 - 1].time_index + 1,
score,
});
i1 = i2;
}
segments
}
pub(crate) fn merge_words<F, G>(
char_segments: &[CharSegment],
is_separator: F,
word_idx_for_token: G,
) -> Vec<WordSegment>
where
F: Fn(usize) -> bool,
G: Fn(usize) -> Option<usize>,
{
let mut words: Vec<WordSegment> = Vec::new();
let n = char_segments.len();
let mut i1 = 0_usize;
let mut i2 = 0_usize;
while i1 < n {
let at_boundary = i2 >= n
|| is_separator(char_segments[i2].token_index)
|| (i2 > i1
&& word_idx_for_token(char_segments[i2].token_index)
!= word_idx_for_token(char_segments[i1].token_index));
if at_boundary {
if i1 != i2 {
let segs = &char_segments[i1..i2];
let mut total_len = 0_usize;
let mut weighted = 0.0_f32;
for seg in segs {
let len = seg.length();
total_len += len;
weighted += seg.score * (len as f32);
}
let score = if total_len == 0 {
0.0
} else {
weighted / (total_len as f32)
};
let word_index = word_idx_for_token(segs[0].token_index).unwrap_or(usize::MAX);
if word_index != usize::MAX {
words.push(WordSegment {
word_index,
start_frame: segs[0].start_frame,
end_frame: segs[segs.len() - 1].end_frame,
score,
});
}
}
if i2 < n && !is_separator(char_segments[i2].token_index) {
i1 = i2;
} else {
i1 = i2 + 1;
i2 = i1;
}
} else {
i2 += 1;
}
}
words
}
pub fn align_to_word_segments(
log_probs: &LogProbsTV,
tokens: &[i32],
word_idx_per_token: &[Option<usize>],
separator_token_id: Option<u32>,
blank_id: u32,
abort_flag: &AtomicBool,
language: &Lang,
) -> Result<Vec<WordSegment>, WorkFailure> {
if tokens.len() != word_idx_per_token.len() {
return Err(WorkFailure::Alignment(AlignmentError::Tokenization(
AlignmentFailure::new(
format_smolstr!(
"tokens.len() = {} != word_idx_per_token.len() = {}; tokenizer bug?",
tokens.len(),
word_idx_per_token.len()
),
language.clone(),
),
)));
}
let trellis = get_trellis(log_probs, tokens, blank_id, abort_flag, language)?;
let path = backtrack_beam(
&trellis,
log_probs,
tokens,
blank_id,
ALIGN_BEAM_WIDTH,
abort_flag,
language,
)?;
let char_segments = merge_repeats(&path);
let is_separator = |tok_idx: usize| -> bool {
if word_idx_per_token.get(tok_idx).copied().flatten().is_none() {
return true;
}
if let Some(sep_id) = separator_token_id {
let token_id = tokens[tok_idx];
if token_id >= 0 && (token_id as u32) == sep_id {
return true;
}
}
false
};
let word_idx =
|tok_idx: usize| -> Option<usize> { word_idx_per_token.get(tok_idx).copied().flatten() };
Ok(merge_words(&char_segments, is_separator, word_idx))
}
#[derive(Debug, Clone)]
pub struct AlignEmissionsConfig {
blank_token_id: u32,
language: Lang,
}
impl AlignEmissionsConfig {
#[must_use]
pub const fn new(blank_token_id: u32, language: Lang) -> Self {
Self {
blank_token_id,
language,
}
}
#[must_use]
pub const fn blank_token_id(&self) -> u32 {
self.blank_token_id
}
#[must_use]
pub const fn language(&self) -> &Lang {
&self.language
}
}
pub fn align_emissions(
log_probs: &LogProbsTV,
tokenized: &TokenizedText,
abort_flag: &AtomicBool,
config: &AlignEmissionsConfig,
) -> Result<Vec<WordSegment>, EmissionsError> {
let t = log_probs.t();
if t > SEAM_PATH_FRAME_BUDGET {
return Err(EmissionsError::PathBudget(EmissionsFailure::new(
format_smolstr!(
"emissions frame count T={t} exceeds the seam path-reconstruction budget of \
{SEAM_PATH_FRAME_BUDGET} frames; the CTC path holds one point per frame, so aligning \
at this T would reserve ~{} MiB up front. Supply emissions with a realistic frame \
count (frames \u{2248} audio_samples / encoder_hop).",
t.saturating_mul(core::mem::size_of::<PathPointPublic>()) >> 20
),
)));
}
align_to_word_segments(
log_probs,
tokenized.token_ids(),
tokenized.word_idx_per_token(),
tokenized.separator_token_id(),
config.blank_token_id(),
abort_flag,
config.language(),
)
.map_err(into_emissions_error)
}
fn into_emissions_error(err: WorkFailure) -> EmissionsError {
let neutral = |f: AlignmentFailure| EmissionsFailure::new(f.message().clone());
match err {
WorkFailure::Alignment(inner) => match inner {
AlignmentError::ModelInference(f) => EmissionsError::Config(neutral(f)),
AlignmentError::Tokenization(f) => EmissionsError::Tokenization(neutral(f)),
AlignmentError::NoAlignmentPath(f) => EmissionsError::NoAlignmentPath(neutral(f)),
AlignmentError::SemanticOutOfVocab(f) => EmissionsError::SemanticOutOfVocab(neutral(f)),
AlignmentError::Aborted(f) => EmissionsError::Aborted(neutral(f)),
AlignmentError::Normalization(f) | AlignmentError::EmptyText(f) => {
EmissionsError::Tokenization(neutral(f))
}
},
WorkFailure::WorkerHang(_timeout) => EmissionsError::Aborted(EmissionsFailure::new(
format_smolstr!("align_emissions aborted via abort_flag before completing"),
)),
other @ (WorkFailure::Asr(_) | WorkFailure::LanguageUnsupported(_)) => {
EmissionsError::Config(EmissionsFailure::new(format_smolstr!(
"align_emissions: internal call chain produced an unexpected WorkFailure \
variant ({other:?}); this indicates a bug in the relocation, not the algorithm"
)))
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::Lang;
fn lp(t: usize, v: usize, vals: Vec<f32>) -> LogProbsTV {
assert_eq!(vals.len(), t * v);
LogProbsTV::new(t, v, vals).expect("t * v == vals.len(), checked above")
}
fn never() -> &'static AtomicBool {
static NEVER: AtomicBool = AtomicBool::new(false);
&NEVER
}
#[test]
fn trellis_single_token_initial_blank_column() {
let v = 3;
let t = 4;
let mut data = vec![0.0_f32; t * v];
for ti in 0..t {
data[ti * v] = -0.5; data[ti * v + 1] = -10.0;
data[ti * v + 2] = -10.0;
}
let log_probs = lp(t, v, data);
let trellis = get_trellis(&log_probs, &[1], 0, never(), &Lang::En).expect("trellis");
assert_eq!(trellis.len(), t);
assert_eq!(trellis[0], 0.0);
assert_eq!(trellis[1], -0.5);
assert_eq!(trellis[2], -1.0);
assert_eq!(trellis[3], -1.5);
}
#[test]
fn trellis_initial_row_pegs_to_neg_inf() {
let v = 3;
let t = 3;
let log_probs = lp(t, v, vec![-1.0_f32; t * v]);
let trellis = get_trellis(&log_probs, &[1, 2], 0, never(), &Lang::En).expect("trellis");
assert!(trellis[1].is_infinite());
assert!(trellis[1] < 0.0);
}
#[test]
fn trellis_final_rows_force_inf_on_column_zero() {
let v = 3;
let t = 5;
let log_probs = lp(t, v, vec![-1.0_f32; t * v]);
let trellis = get_trellis(&log_probs, &[1, 2, 1], 0, never(), &Lang::En).expect("trellis");
assert!(trellis[3 * 3].is_infinite() && trellis[3 * 3] > 0.0);
assert!(trellis[4 * 3].is_infinite() && trellis[4 * 3] > 0.0);
}
#[test]
fn trellis_recurrence_picks_max_of_stay_and_change() {
let v = 3;
let t = 3;
let mut data = vec![-100.0_f32; t * v];
for ti in 0..t {
data[ti * v] = -1.0; data[ti * v + 1] = -2.0; data[ti * v + 2] = -2.0; }
let log_probs = lp(t, v, data);
let trellis = get_trellis(&log_probs, &[1, 2], 0, never(), &Lang::En).expect("trellis");
let last_cell = trellis[2 * 2 + 1];
assert!(
last_cell.is_finite(),
"trellis end cell must be finite for a viable lattice; got {last_cell}"
);
}
#[test]
fn tokens_zeroth_emission_does_not_affect_trellis() {
let v = 4;
let t = 4;
let blank = 0;
let tokens = [1_i32, 2];
let mut base = vec![-2.0_f32; t * v];
for ti in 0..t {
base[ti * v + blank] = -0.5;
}
let mut a = base.clone();
let mut b = base.clone();
for ti in 0..t {
a[ti * v + tokens[0] as usize] = -10.0; b[ti * v + tokens[0] as usize] = -0.1; }
for ti in 0..t {
a[ti * v + tokens[1] as usize] = -1.0;
b[ti * v + tokens[1] as usize] = -1.0;
}
let lp_a = lp(t, v, a);
let lp_b = lp(t, v, b);
let trellis_a = get_trellis(&lp_a, &tokens, blank as u32, never(), &Lang::En).expect("a");
let trellis_b = get_trellis(&lp_b, &tokens, blank as u32, never(), &Lang::En).expect("b");
assert_eq!(
trellis_a, trellis_b,
"Changing `tokens[0]`'s emission posterior must NOT change the trellis. \
If this fires the WhisperX-parity quirk has been broken — see the long \
comment above the forward DP loop in get_trellis."
);
}
#[test]
fn wildcard_emission_uses_max_non_blank() {
let v = 4;
let log_probs = lp(1, v, vec![0.0, -2.0, -1.0, -3.0]);
let m = max_non_blank_logprob(&log_probs, 0, 0);
assert!((m - (-1.0)).abs() < 1e-6);
}
#[test]
fn beam_step_uses_predecessor_only_score() {
let v = 3;
let t = 3;
let mut data = vec![-100.0_f32; t * v];
data[0] = -0.5; data[1] = -0.4; data[3] = -0.5; data[4] = -0.3; data[5] = -0.4; data[6] = -0.5; data[7] = -0.2; let log_probs = lp(t, v, data);
let trellis = get_trellis(&log_probs, &[1, 2], 0, never(), &Lang::En).expect("trellis");
let path = backtrack_beam(
&trellis,
&log_probs,
&[1, 2],
0,
ALIGN_BEAM_WIDTH,
never(),
&Lang::En,
)
.expect("path");
assert_eq!(path.len(), t);
let token_seq: Vec<usize> = path.iter().map(|p| p.token_index).collect();
assert_eq!(token_seq[0], 0, "leading blank invariant");
assert_eq!(
*token_seq.last().expect("non-empty"),
1,
"must reach final token"
);
}
#[test]
fn backtrack_beam_simple_two_token_path() {
let v = 3;
let t = 3;
let mut data = vec![-100.0_f32; t * v];
data[1] = -0.1; data[3] = -0.1; data[8] = -0.1; data[0] = -0.5;
data[1] = -1.0;
data[2] = -1.0;
data[6] = -0.5;
let log_probs = lp(t, v, data);
let trellis = get_trellis(&log_probs, &[1, 2], 0, never(), &Lang::En).expect("trellis");
let path = backtrack_beam(
&trellis,
&log_probs,
&[1, 2],
0,
ALIGN_BEAM_WIDTH,
never(),
&Lang::En,
)
.expect("path");
assert_eq!(path.len(), t);
assert_eq!(path[0].time_index, 0);
assert_eq!(path[t - 1].time_index, t - 1);
}
#[test]
fn merge_repeats_groups_by_token_index() {
let path = vec![
PathPointPublic {
token_index: 0,
time_index: 0,
score: 0.5,
},
PathPointPublic {
token_index: 0,
time_index: 1,
score: 0.7,
},
PathPointPublic {
token_index: 1,
time_index: 2,
score: 0.9,
},
PathPointPublic {
token_index: 1,
time_index: 3,
score: 0.5,
},
PathPointPublic {
token_index: 2,
time_index: 4,
score: 0.5,
},
];
let segs = merge_repeats(&path);
assert_eq!(segs.len(), 3);
assert_eq!(segs[0].token_index, 0);
assert_eq!(segs[0].start_frame, 0);
assert_eq!(segs[0].end_frame, 2);
assert!((segs[0].score - 0.6).abs() < 1e-6);
assert_eq!(segs[1].start_frame, 2);
assert_eq!(segs[1].end_frame, 4);
assert_eq!(segs[2].start_frame, 4);
assert_eq!(segs[2].end_frame, 5);
let two_frame = vec![
PathPointPublic {
token_index: 7,
time_index: 10,
score: (0.0_f32).exp(),
},
PathPointPublic {
token_index: 7,
time_index: 11,
score: (-2.0_f32).exp(),
},
];
let two_seg = merge_repeats(&two_frame);
assert_eq!(two_seg.len(), 1);
let mean_of_exp = ((0.0_f32).exp() + (-2.0_f32).exp()) / 2.0; let exp_of_mean = (-1.0_f32).exp(); assert!(
(two_seg[0].score - mean_of_exp).abs() < 1e-6,
"score must be mean(exp(...)) ≈ 0.5677; got {}",
two_seg[0].score
);
assert!(
(two_seg[0].score - exp_of_mean).abs() > 0.1,
"score must NOT be exp(mean(...)) ≈ 0.3679; got {}",
two_seg[0].score
);
}
#[test]
fn merge_words_groups_chars_by_separator() {
let mut segs: Vec<CharSegment> = Vec::new();
for i in 0..11 {
segs.push(CharSegment {
token_index: i,
start_frame: i * 2,
end_frame: i * 2 + 2,
score: 0.5,
});
}
let is_sep = |t: usize| t == 5;
let word_idx = |t: usize| -> Option<usize> {
if t == 5 {
None
} else if t < 5 {
Some(0)
} else {
Some(1)
}
};
let words = merge_words(&segs, is_sep, word_idx);
assert_eq!(words.len(), 2);
assert_eq!(words[0].word_index, 0);
assert_eq!(words[0].start_frame, 0);
assert_eq!(words[0].end_frame, 10);
assert_eq!(words[1].word_index, 1);
assert_eq!(words[1].start_frame, 12);
assert_eq!(words[1].end_frame, 22);
}
#[test]
fn merge_words_score_is_duration_weighted() {
let segs = vec![
CharSegment {
token_index: 0,
start_frame: 0,
end_frame: 1,
score: 0.5,
},
CharSegment {
token_index: 1,
start_frame: 1,
end_frame: 4,
score: 1.0,
},
];
let is_sep = |_| false;
let word_idx = |_| Some(0_usize);
let words = merge_words(&segs, is_sep, word_idx);
assert_eq!(words.len(), 1);
assert!(
(words[0].score - 0.875).abs() < 1e-6,
"duration-weighted score wrong: {}",
words[0].score
);
}
#[test]
fn merge_words_no_separator_splits_by_word_idx() {
let segs = vec![
CharSegment {
token_index: 0,
start_frame: 0,
end_frame: 2,
score: 0.5,
},
CharSegment {
token_index: 1,
start_frame: 2,
end_frame: 4,
score: 0.5,
},
CharSegment {
token_index: 2,
start_frame: 4,
end_frame: 6,
score: 0.5,
},
];
let is_sep = |_| false;
let word_idx = |t: usize| -> Option<usize> {
match t {
0 => Some(0),
1 => Some(1),
2 => Some(2),
_ => None,
}
};
let words = merge_words(&segs, is_sep, word_idx);
assert_eq!(words.len(), 3, "each glyph must become its own word");
assert_eq!(words[0].word_index, 0);
assert_eq!(words[0].start_frame, 0);
assert_eq!(words[0].end_frame, 2);
assert_eq!(words[1].word_index, 1);
assert_eq!(words[1].start_frame, 2);
assert_eq!(words[1].end_frame, 4);
assert_eq!(words[2].word_index, 2);
assert_eq!(words[2].start_frame, 4);
assert_eq!(words[2].end_frame, 6);
}
#[test]
fn merge_words_no_separator_groups_same_word_idx_across_chars() {
let segs = vec![
CharSegment {
token_index: 0,
start_frame: 0,
end_frame: 2,
score: 0.5,
},
CharSegment {
token_index: 1,
start_frame: 2,
end_frame: 4,
score: 0.5,
},
CharSegment {
token_index: 2,
start_frame: 4,
end_frame: 6,
score: 0.5,
},
];
let is_sep = |_| false;
let word_idx = |t: usize| -> Option<usize> {
match t {
0 => Some(0),
1 => Some(0),
2 => Some(1),
_ => None,
}
};
let words = merge_words(&segs, is_sep, word_idx);
assert_eq!(words.len(), 2);
assert_eq!(words[0].word_index, 0);
assert_eq!(words[0].start_frame, 0);
assert_eq!(words[0].end_frame, 4); assert_eq!(words[1].word_index, 1);
assert_eq!(words[1].start_frame, 4);
assert_eq!(words[1].end_frame, 6); }
#[test]
fn merge_words_separator_still_works() {
let segs = vec![
CharSegment {
token_index: 0,
start_frame: 0,
end_frame: 2,
score: 0.5,
},
CharSegment {
token_index: 1, start_frame: 2,
end_frame: 3,
score: 0.5,
},
CharSegment {
token_index: 2,
start_frame: 3,
end_frame: 5,
score: 0.5,
},
];
let is_sep = |t: usize| t == 1;
let word_idx = |t: usize| -> Option<usize> {
match t {
0 => Some(0),
1 => None,
2 => Some(1),
_ => None,
}
};
let words = merge_words(&segs, is_sep, word_idx);
assert_eq!(words.len(), 2);
assert_eq!(words[0].word_index, 0);
assert_eq!(words[0].start_frame, 0);
assert_eq!(words[0].end_frame, 2);
assert_eq!(words[1].word_index, 1);
assert_eq!(words[1].start_frame, 3);
assert_eq!(words[1].end_frame, 5);
}
#[test]
fn backtrack_beam_visits_every_token_on_codex_counterexample() {
let v = 4;
let t = 4;
let mut data = vec![-100.0_f32; t * v];
data[0] = -10.0; data[1] = -0.1; data[4] = -10.0; data[6] = -0.1; data[8] = -0.1; data[10] = -2.0; data[11] = -2.0; data[12] = -10.0; data[15] = -0.1;
let log_probs = lp(t, v, data);
let tokens = vec![1_i32, 2_i32, 3_i32];
let abort = AtomicBool::new(false);
let trellis = get_trellis(&log_probs, &tokens, 0, &abort, &Lang::En).expect("trellis builds");
let path = backtrack_beam(
&trellis,
&log_probs,
&tokens,
0,
ALIGN_BEAM_WIDTH,
&abort,
&Lang::En,
)
.expect("beam backtracks");
let coords: Vec<(usize, usize)> = path.iter().map(|p| (p.token_index, p.time_index)).collect();
let visited: std::collections::BTreeSet<usize> = coords.iter().map(|(j, _)| *j).collect();
assert!(
visited.contains(&0) && visited.contains(&1) && visited.contains(&2),
"transition-scored backtrack must visit every token; got {:?}",
coords
);
}
#[test]
fn align_to_word_segments_simple_smoke() {
let v = 3;
let t = 4;
let mut data = vec![-100.0_f32; t * v];
data[1] = -0.1;
data[3] = -0.1;
data[8] = -0.1;
data[9] = -0.1;
let log_probs = lp(t, v, data);
let words = align_to_word_segments(
&log_probs,
&[1, 2],
&[Some(0), Some(0)],
None,
0,
never(),
&Lang::En,
)
.expect("words");
assert_eq!(words.len(), 1);
assert_eq!(words[0].word_index, 0);
}
#[test]
fn align_emissions_known_words_from_golden_emission() {
let v = 3;
let t = 4;
let mut data = vec![-100.0_f32; t * v];
data[1] = -0.1;
data[3] = -0.1;
data[8] = -0.1;
data[9] = -0.1;
let log_probs = lp(t, v, data);
let tokenized = TokenizedText::new(vec![1, 2], vec![Some(0), Some(0)], None);
let config = AlignEmissionsConfig::new(0, Lang::En);
assert_eq!(config.blank_token_id(), 0);
assert_eq!(config.language(), &Lang::En);
let words = align_emissions(&log_probs, &tokenized, never(), &config).expect("words");
assert_eq!(words.len(), 1);
assert_eq!(words[0].word_index(), 0);
let direct = align_to_word_segments(
&log_probs,
tokenized.token_ids(),
tokenized.word_idx_per_token(),
tokenized.separator_token_id(),
config.blank_token_id(),
never(),
config.language(),
)
.expect("words via the internal call chain align_emissions wraps");
assert_eq!(words.len(), direct.len());
for (via_emissions, via_internal) in words.iter().zip(direct.iter()) {
assert_eq!(via_emissions.word_index(), via_internal.word_index());
assert_eq!(via_emissions.start_frame(), via_internal.start_frame());
assert_eq!(via_emissions.end_frame(), via_internal.end_frame());
assert_eq!(via_emissions.score(), via_internal.score());
}
}
#[test]
fn align_emissions_reports_abort_flag_as_aborted() {
let v = 3;
let t = 4;
let log_probs = lp(t, v, vec![-1.0_f32; t * v]);
let tokenized = TokenizedText::new(vec![1, 2], vec![Some(0), Some(0)], None);
let config = AlignEmissionsConfig::new(0, Lang::En);
let abort = AtomicBool::new(true);
let err = align_emissions(&log_probs, &tokenized, &abort, &config).unwrap_err();
let EmissionsError::Aborted(payload) = &err else {
panic!("expected EmissionsError::Aborted; got {err:?}");
};
let message = payload.message().to_ascii_lowercase();
assert!(
message.contains("abort") || message.contains("cancel"),
"Aborted payload should name the abort/cancellation path; got {message:?}"
);
for banned in ["worker", "hung", "elapsed"] {
assert!(
!message.contains(banned),
"Aborted payload leaked pool/worker vocabulary ({banned:?}) that doesn't apply \
to a bare align_emissions call; got {message:?}"
);
}
}
#[test]
fn align_emissions_surfaces_tokenization_errors_unwrapped() {
let v = 3;
let t = 4;
let log_probs = lp(t, v, vec![-1.0_f32; t * v]);
let tokenized = TokenizedText::new(vec![1, 99], vec![Some(0), Some(0)], None);
let config = AlignEmissionsConfig::new(0, Lang::En);
let err = align_emissions(&log_probs, &tokenized, never(), &config).unwrap_err();
assert!(
matches!(err, EmissionsError::Tokenization(_)),
"expected EmissionsError::Tokenization; got {err:?}"
);
}
#[test]
fn align_emissions_reports_bad_blank_id_as_config() {
let v = 3;
let t = 4;
let log_probs = lp(t, v, vec![-1.0_f32; t * v]);
let tokenized = TokenizedText::new(vec![1, 2], vec![Some(0), Some(0)], None);
let config = AlignEmissionsConfig::new(99, Lang::En);
let err = align_emissions(&log_probs, &tokenized, never(), &config).unwrap_err();
assert!(
matches!(err, EmissionsError::Config(_)),
"an out-of-range blank id must surface as Config; got {err:?}"
);
let s = err.to_string();
for banned in [
"ORT",
"worker",
"pool",
"Event::Error",
"ASR text preserved",
] {
assert!(!s.contains(banned), "Config Display leaked {banned:?}: {s}");
}
}
#[test]
fn align_emissions_rejects_oversized_frame_count_before_allocating() {
let t = SEAM_PATH_FRAME_BUDGET + 1;
let log_probs = lp(t, 1, vec![-1.0_f32; t]);
let tokenized = TokenizedText::new(vec![0], vec![Some(0)], None);
let config = AlignEmissionsConfig::new(0, Lang::En);
let err = align_emissions(&log_probs, &tokenized, never(), &config).unwrap_err();
assert!(
matches!(err, EmissionsError::PathBudget(_)),
"an oversized frame count must fail fast as PathBudget; got {err:?}"
);
let s = err.to_string();
assert!(
s.contains("path budget"),
"Display should name the exceeded budget; got {s}"
);
for banned in ["ORT", "worker", "pool", "Event::Error"] {
assert!(
!s.contains(banned),
"PathBudget Display leaked {banned:?}: {s}"
);
}
let mut ok_data = vec![-100.0_f32; 12];
ok_data[1] = -0.1;
ok_data[3] = -0.1;
ok_data[8] = -0.1;
ok_data[9] = -0.1;
let ok_log_probs = lp(4, 3, ok_data);
let ok_tokens = TokenizedText::new(vec![1, 2], vec![Some(0), Some(0)], None);
let words = align_emissions(&ok_log_probs, &ok_tokens, never(), &config)
.expect("a normal-size lattice must still align");
assert!(!words.is_empty(), "normal lattice must produce a word");
}
#[test]
fn align_emissions_valid_lattices_produce_in_range_scores() {
let config = AlignEmissionsConfig::new(0, Lang::En);
let assert_in_range = |label: &str, words: &[WordSegment]| {
assert!(!words.is_empty(), "{label}: expected at least one word");
for w in words {
let s = w.score();
assert!(
s.is_finite() && (0.0..=1.0).contains(&s),
"{label}: score {s} must be finite and in [0, 1]"
);
}
};
{
let v = 3;
let t = 4;
let mut data = vec![-100.0_f32; t * v];
data[1] = 0.0; data[3] = -0.1; data[8] = -0.1; data[9] = 0.0; let log_probs = lp(t, v, data);
let tokenized = TokenizedText::new(vec![1, 2], vec![Some(0), Some(0)], None);
let words = align_emissions(&log_probs, &tokenized, never(), &config).expect("case A aligns");
assert_in_range("A/real-token+final-blank=log(1)", &words);
}
{
let v = 4;
let t = 4;
let mut data = vec![-100.0_f32; t * v];
data[1] = -0.1; data[4] = -0.1; data[2 * v + 2] = -0.1; data[3 * v] = -0.1; let log_probs = lp(t, v, data);
let tokenized = TokenizedText::new(vec![1, WILDCARD_TOKEN_ID], vec![Some(0), Some(0)], None);
let words = align_emissions(&log_probs, &tokenized, never(), &config).expect("case B aligns");
assert_in_range("B/wildcard", &words);
}
{
let v = 4;
let t = 6;
let sep = 3_u32;
let mut data = vec![-100.0_f32; t * v];
data[1] = -0.1; data[4] = -0.1; data[2 * v + sep as usize] = -0.1; data[3 * v] = -0.1; data[4 * v + 2] = -0.1; data[5 * v] = -0.1; let log_probs = lp(t, v, data);
let tokenized = TokenizedText::new(
vec![1, sep as i32, 2],
vec![Some(0), None, Some(1)],
Some(sep),
);
let words = align_emissions(&log_probs, &tokenized, never(), &config).expect("case C aligns");
assert_in_range("C/separator-two-words", &words);
assert_eq!(words.len(), 2, "case C must split into two words");
}
}
#[test]
fn beam_picks_globally_best_when_local_tie_exists() {
let v = 3;
let t = 4;
let mut data = vec![-1.0_f32; t * v];
data[0] = -0.1;
data[1] = -1.0;
data[2] = -1.0;
data[3] = -1.0;
data[4] = -0.1;
data[5] = -1.0;
data[6] = -1.0;
data[7] = -1.0;
data[8] = -0.1;
data[9] = -0.1;
data[10] = -1.0;
data[11] = -1.0;
let log_probs = lp(t, v, data);
let trellis = get_trellis(&log_probs, &[1, 2], 0, never(), &Lang::En).expect("trellis");
let path =
backtrack_beam(&trellis, &log_probs, &[1, 2], 0, 2, never(), &Lang::En).expect("path");
assert_eq!(path.len(), t);
let tokens: Vec<usize> = path.iter().map(|p| p.token_index).collect();
assert!(tokens.contains(&0));
assert!(tokens.contains(&1));
}
#[test]
fn empty_token_sequence_returns_no_alignment_path() {
let log_probs = lp(3, 3, vec![0.0_f32; 9]);
let err = get_trellis(&log_probs, &[], 0, never(), &Lang::En).unwrap_err();
assert!(matches!(
err,
WorkFailure::Alignment(AlignmentError::NoAlignmentPath(_))
));
}
#[test]
fn audio_too_short_t_lt_num_tokens_errors() {
let log_probs = lp(2, 4, vec![0.0_f32; 8]);
let err = get_trellis(&log_probs, &[1, 2, 3], 0, never(), &Lang::En).unwrap_err();
assert!(matches!(
err,
WorkFailure::Alignment(AlignmentError::NoAlignmentPath(_))
));
}
#[test]
fn out_of_vocab_real_token_id_errors() {
let log_probs = lp(3, 3, vec![0.0_f32; 9]);
let err = get_trellis(&log_probs, &[1, 99], 0, never(), &Lang::En).unwrap_err();
assert!(matches!(
err,
WorkFailure::Alignment(AlignmentError::Tokenization(_))
));
}
#[test]
fn wildcard_token_id_minus_one_passes_validation() {
let log_probs = lp(3, 4, vec![-0.5_f32; 12]);
let trellis = get_trellis(&log_probs, &[1, WILDCARD_TOKEN_ID], 0, never(), &Lang::En);
assert!(trellis.is_ok(), "wildcard tokens must pass validation");
}
#[test]
fn negative_real_token_id_other_than_wildcard_errors() {
let log_probs = lp(3, 3, vec![0.0_f32; 9]);
let err = get_trellis(&log_probs, &[1, -2], 0, never(), &Lang::En).unwrap_err();
assert!(matches!(
err,
WorkFailure::Alignment(AlignmentError::Tokenization(_))
));
}
#[test]
fn aborted_trellis_returns_worker_hang_timeout() {
let log_probs = lp(2_000, 4, vec![-0.1_f32; 2_000 * 4]);
let tokens: Vec<i32> = (0..200).map(|i| 1 + (i % 3)).collect();
let abort = AtomicBool::new(true);
let err = get_trellis(&log_probs, &tokens, 0, &abort, &Lang::En).unwrap_err();
assert!(matches!(
err,
WorkFailure::WorkerHang(ref t) if t.kind() == WorkerKind::Alignment
));
}
#[test]
fn budget_exceeded_returns_no_alignment_path() {
let log_probs =
LogProbsTV::new(8_000, 8, vec![0.0_f32; 8_000 * 8]).expect("t * v == vals.len()");
let tokens: Vec<i32> = (0..5_000).map(|i| 1 + (i % 4)).collect();
let err = get_trellis(&log_probs, &tokens, 0, never(), &Lang::En).unwrap_err();
let WorkFailure::Alignment(AlignmentError::NoAlignmentPath(payload)) = err else {
panic!("expected AlignmentFailed");
};
let message = payload.message();
assert!(
message.contains("trellis exceeds"),
"message must call out the budget; got {message}",
message = message
);
}
}