#[cfg(feature = "alignment")]
use ort::{
session::{RunOptions, Session},
value::{Shape, Tensor},
};
use smol_str::{SmolStr, format_smolstr};
use super::errors::{EmissionsError, EmissionsFailure};
use crate::types::{AlignmentError, AlignmentFailure, Lang, WorkFailure};
#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)]
pub struct LogProbsShapeError {
t: usize,
v: usize,
data_len: usize,
}
impl core::fmt::Display for LogProbsShapeError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
if self.v == 0 {
write!(
f,
"LogProbsTV has a zero-length vocab dim: t={}, v=0, data.len()={} \
(a CTC vocabulary must contain at least the blank token)",
self.t, self.data_len
)
} else {
write!(
f,
"LogProbsTV shape mismatch: t={}, v={}, data.len()={} (expected \
data.len() == t * v)",
self.t, self.v, self.data_len
)
}
}
}
impl LogProbsShapeError {
const fn new(t: usize, v: usize, data_len: usize) -> Self {
Self { t, v, data_len }
}
#[must_use]
pub const fn t(&self) -> usize {
self.t
}
#[must_use]
pub const fn v(&self) -> usize {
self.v
}
#[must_use]
pub const fn data_len(&self) -> usize {
self.data_len
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum LogProbsValueClass {
Nan,
PosInf,
NegInf,
Positive,
}
impl core::fmt::Display for LogProbsValueClass {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.write_str(match self {
Self::Nan => "NaN",
Self::PosInf => "+Inf",
Self::NegInf => "-Inf",
Self::Positive => "positive",
})
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)]
#[error(
"log-probability out of domain (finite and ≤ 0) at frame {frame}, vocab {vocab_index}: {class}"
)]
pub struct LogProbsValueError {
frame: usize,
vocab_index: usize,
class: LogProbsValueClass,
}
impl LogProbsValueError {
const fn new(frame: usize, vocab_index: usize, class: LogProbsValueClass) -> Self {
Self {
frame,
vocab_index,
class,
}
}
#[must_use]
pub const fn frame(&self) -> usize {
self.frame
}
#[must_use]
pub const fn vocab_index(&self) -> usize {
self.vocab_index
}
#[must_use]
pub const fn class(&self) -> LogProbsValueClass {
self.class
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)]
pub enum LogProbsError {
#[error(transparent)]
Shape(LogProbsShapeError),
#[error(transparent)]
Value(LogProbsValueError),
}
pub struct LogProbsTV {
t: usize,
v: usize,
data: Vec<f32>,
}
impl LogProbsTV {
pub fn new(t: usize, v: usize, data: Vec<f32>) -> Result<Self, LogProbsError> {
if v == 0 {
return Err(LogProbsError::Shape(LogProbsShapeError::new(
t,
v,
data.len(),
)));
}
if t.checked_mul(v) != Some(data.len()) {
return Err(LogProbsError::Shape(LogProbsShapeError::new(
t,
v,
data.len(),
)));
}
if let Some(idx) = data.iter().position(|&x| !(x.is_finite() && x <= 0.0)) {
let bad = data[idx];
let class = if bad.is_nan() {
LogProbsValueClass::Nan
} else if bad.is_infinite() {
if bad > 0.0 {
LogProbsValueClass::PosInf
} else {
LogProbsValueClass::NegInf
}
} else {
LogProbsValueClass::Positive
};
return Err(LogProbsError::Value(LogProbsValueError::new(
idx / v,
idx % v,
class,
)));
}
Ok(Self { t, v, data })
}
pub(crate) const fn from_parts_unchecked(t: usize, v: usize, data: Vec<f32>) -> Self {
Self { t, v, data }
}
#[must_use]
pub const fn t(&self) -> usize {
self.t
}
#[must_use]
pub const fn v(&self) -> usize {
self.v
}
#[must_use]
pub fn data(&self) -> &[f32] {
&self.data
}
#[must_use]
pub fn get(&self, t_idx: usize, v_idx: usize) -> Option<f32> {
if v_idx >= self.v {
return None;
}
let idx = t_idx.checked_mul(self.v)?.checked_add(v_idx)?;
self.data.get(idx).copied()
}
#[must_use]
pub fn at(&self, t_idx: usize, v_idx: usize) -> f32 {
match self.get(t_idx, v_idx) {
Some(lp) => lp,
None => panic!(
"LogProbsTV index out of bounds: (t={t_idx}, v={v_idx}) is outside the (T={}, V={}) grid",
self.t, self.v
),
}
}
}
#[cfg(feature = "alignment")]
pub(crate) fn encode_log_softmax(
session: &mut Session,
samples_for_aligner: &[f32],
run_options: &RunOptions,
language: &Lang,
) -> Result<LogProbsTV, WorkFailure> {
let t_samples = samples_for_aligner.len();
if t_samples == 0 {
return Err(WorkFailure::Alignment(AlignmentError::ModelInference(
AlignmentFailure::new(
SmolStr::from("samples_for_aligner is empty"),
language.clone(),
),
)));
}
reject_non_finite_input(samples_for_aligner, language)?;
let input_shape: [i64; 2] = [1, t_samples as i64];
let input_tensor =
Tensor::from_array((input_shape, samples_for_aligner.to_vec())).map_err(|e| {
WorkFailure::Alignment(AlignmentError::ModelInference(AlignmentFailure::new(
format_smolstr!("Tensor::from_array failed: {e:?}"),
language.clone(),
)))
})?;
let outputs = session
.run_with_options(ort::inputs![input_tensor], run_options)
.map_err(|e| {
WorkFailure::Alignment(AlignmentError::ModelInference(AlignmentFailure::new(
format_smolstr!("Session::run_with_options failed: {e:?}"),
language.clone(),
)))
})?;
let mut iter = outputs.into_iter();
let (_, output_value) = iter.next().ok_or_else(|| {
WorkFailure::Alignment(AlignmentError::ModelInference(AlignmentFailure::new(
SmolStr::from("Session::run returned no outputs"),
language.clone(),
)))
})?;
let (shape, raw): (&Shape, &[f32]) = output_value.try_extract_tensor::<f32>().map_err(|e| {
WorkFailure::Alignment(AlignmentError::ModelInference(AlignmentFailure::new(
format_smolstr!("try_extract_tensor::<f32> failed: {e:?}"),
language.clone(),
)))
})?;
if shape.len() != 3 || shape[0] != 1 {
return Err(WorkFailure::Alignment(AlignmentError::ModelInference(
AlignmentFailure::new(
format_smolstr!("expected output shape (1, T, V); got {shape:?}"),
language.clone(),
),
)));
}
let (t, v) = validate_output_dims(shape[1], shape[2], raw.len(), language)?;
let data = log_softmax_with_finite_guard(raw, t, v).map_err(|e| e.into_work_failure(language))?;
Ok(LogProbsTV::from_parts_unchecked(t, v, data))
}
pub(crate) fn validate_output_dims(
raw_t: i64,
raw_v: i64,
raw_len: usize,
language: &Lang,
) -> Result<(usize, usize), WorkFailure> {
if raw_v <= 0 {
return Err(WorkFailure::Alignment(AlignmentError::ModelInference(
AlignmentFailure::new(
format_smolstr!("ORT output has non-positive vocab dim: V={raw_v}"),
language.clone(),
),
)));
}
if raw_t < 0 {
return Err(WorkFailure::Alignment(AlignmentError::ModelInference(
AlignmentFailure::new(
format_smolstr!("ORT output has negative time dim: T={raw_t}"),
language.clone(),
),
)));
}
if raw_t == 0 {
if raw_len != 0 {
return Err(WorkFailure::Alignment(AlignmentError::ModelInference(
AlignmentFailure::new(
format_smolstr!(
"ORT output declared T=0 but buffer has {raw_len} elements; shape/data mismatch"
),
language.clone(),
),
)));
}
return Err(WorkFailure::Alignment(AlignmentError::NoAlignmentPath(
AlignmentFailure::new(
SmolStr::from(
"ORT output has zero encoder frames (chunk too short to align); \
transcript will surface with words: []",
),
language.clone(),
),
)));
}
let t = match usize::try_from(raw_t) {
Ok(v) => v,
Err(_) => {
return Err(WorkFailure::Alignment(AlignmentError::ModelInference(
AlignmentFailure::new(
format_smolstr!("ORT output T={raw_t} doesn't fit in usize"),
language.clone(),
),
)));
}
};
let v = match usize::try_from(raw_v) {
Ok(v) => v,
Err(_) => {
return Err(WorkFailure::Alignment(AlignmentError::ModelInference(
AlignmentFailure::new(
format_smolstr!("ORT output V={raw_v} doesn't fit in usize"),
language.clone(),
),
)));
}
};
let total = match t.checked_mul(v) {
Some(p) => p,
None => {
return Err(WorkFailure::Alignment(AlignmentError::ModelInference(
AlignmentFailure::new(
format_smolstr!("ORT output dimensions overflow: T={t} * V={v} doesn't fit in usize"),
language.clone(),
),
)));
}
};
if total != raw_len {
return Err(WorkFailure::Alignment(AlignmentError::ModelInference(
AlignmentFailure::new(
format_smolstr!(
"ORT output buffer length {raw_len} doesn't match declared T={t} × V={v} = {total}"
),
language.clone(),
),
)));
}
Ok((t, v))
}
pub(crate) fn validate_stride_extent(
t: usize,
hop_samples: u32,
chunk_extent: usize,
language: &Lang,
) -> Result<(), WorkFailure> {
let frame_extent = (t as u64).saturating_mul(hop_samples as u64);
let chunk_extent_u64 = chunk_extent as u64;
let slack = 2u64.saturating_mul(hop_samples as u64);
let upper_bound = chunk_extent_u64.saturating_add(slack);
let lower_bound = chunk_extent_u64.saturating_sub(slack);
if frame_extent > upper_bound {
return Err(WorkFailure::Alignment(AlignmentError::ModelInference(
AlignmentFailure::new(
format_smolstr!(
"ORT output stride mismatch: T={t} × hop={hop_samples} = {frame_extent} \
sample-equivalents exceeds chunk ({chunk_extent} samples) + 2-frame slack \
({upper_bound}); model export uses a smaller stride than `hop_samples` \
or `hop_samples` is misconfigured"
),
language.clone(),
),
)));
}
if frame_extent < lower_bound {
return Err(WorkFailure::Alignment(AlignmentError::ModelInference(
AlignmentFailure::new(
format_smolstr!(
"ORT output stride mismatch: T={t} × hop={hop_samples} = {frame_extent} \
sample-equivalents below chunk ({chunk_extent} samples) − 2-frame slack \
({lower_bound}); model export uses a larger stride than `hop_samples` \
or `hop_samples` is misconfigured"
),
language.clone(),
),
)));
}
Ok(())
}
pub(crate) fn validate_vocab_dim(
v: usize,
expected_v: usize,
language: &Lang,
) -> Result<(), WorkFailure> {
if v != expected_v {
return Err(WorkFailure::Alignment(AlignmentError::ModelInference(
AlignmentFailure::new(
format_smolstr!(
"ORT output vocab dim V={v} doesn't match tokenizer vocab size {expected_v}; \
model and tokenizer are paired incorrectly — Viterbi would otherwise read \
posteriors from columns that don't correspond to the tokenizer's tokens"
),
language.clone(),
),
)));
}
Ok(())
}
pub fn log_softmax_with_finite_guard(
raw: &[f32],
t: usize,
v: usize,
) -> Result<Vec<f32>, EmissionsError> {
if v == 0 {
return Err(EmissionsError::Shape(LogProbsShapeError::new(
t,
v,
raw.len(),
)));
}
let Some(total) = t.checked_mul(v) else {
return Err(EmissionsError::Shape(LogProbsShapeError::new(
t,
v,
raw.len(),
)));
};
if total != raw.len() {
return Err(EmissionsError::Shape(LogProbsShapeError::new(
t,
v,
raw.len(),
)));
}
let mut data = Vec::with_capacity(total);
for t_idx in 0..t {
let row = &raw[t_idx * v..(t_idx + 1) * v];
if let Some(bad_v) = row.iter().position(|x| !x.is_finite()) {
return Err(EmissionsError::Numeric(EmissionsFailure::new(
format_smolstr!(
"encoder supplied non-finite logit at frame {t_idx}, vocab {bad_v}: {}",
row[bad_v]
),
)));
}
let max = row.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let max_f64 = max as f64;
let mut sum = 0.0_f64;
for &x in row {
sum += ((x as f64) - max_f64).exp();
}
let log_z_shifted = sum.ln();
if !log_z_shifted.is_finite() {
return Err(EmissionsError::Numeric(EmissionsFailure::new(
format_smolstr!(
"log-softmax shifted normaliser non-finite at frame {t_idx}: \
sum.ln()={log_z_shifted}, max={max}"
),
)));
}
for &x in row {
let lp_f64 = ((x as f64) - max_f64) - log_z_shifted;
let lp = lp_f64 as f32;
if !lp.is_finite() {
return Err(EmissionsError::Numeric(EmissionsFailure::new(
format_smolstr!(
"log-softmax output non-finite at frame {t_idx}: \
x={x}, max={max}, sum_ln={log_z_shifted}, lp={lp}"
),
)));
}
data.push(lp);
}
}
Ok(data)
}
pub(crate) fn reject_non_finite_input(samples: &[f32], language: &Lang) -> Result<(), WorkFailure> {
if let Some(bad_idx) = samples.iter().position(|s| !s.is_finite()) {
return Err(WorkFailure::Alignment(AlignmentError::ModelInference(
AlignmentFailure::new(
format_smolstr!(
"samples_for_aligner contains non-finite value at index {bad_idx}: {}",
samples[bad_idx]
),
language.clone(),
),
)));
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn log_softmax_sums_to_zero_in_log_space() {
let row = [1.0f32, 2.0, 3.0];
let max = row.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0_f64;
for &x in &row {
sum += ((x - max) as f64).exp();
}
let log_z = max + (sum.ln() as f32);
let lp: Vec<f32> = row.iter().map(|x| x - log_z).collect();
let exp_sum: f32 = lp.iter().map(|x| x.exp()).sum();
assert!((exp_sum - 1.0).abs() < 1e-5, "softmax must sum to 1");
for &v in &lp {
assert!(v <= 0.0, "log-prob must be <= 0");
}
}
#[test]
fn reject_non_finite_input_flags_nan() {
use crate::types::Lang;
let samples = vec![0.1_f32, 0.2, f32::NAN, 0.4];
let err = reject_non_finite_input(&samples, &Lang::En).unwrap_err();
match err {
WorkFailure::Alignment(AlignmentError::ModelInference(payload)) => {
assert!(
payload.message().contains("index 2"),
"message must name index; got {message}",
message = payload.message()
);
}
other => panic!("expected AlignmentFailed; got {other:?}"),
}
}
#[test]
fn reject_non_finite_input_flags_positive_infinity() {
use crate::types::Lang;
let samples = vec![0.0_f32, f32::INFINITY];
assert!(reject_non_finite_input(&samples, &Lang::En).is_err());
}
#[test]
fn reject_non_finite_input_flags_negative_infinity() {
use crate::types::Lang;
let samples = vec![f32::NEG_INFINITY, 0.0_f32];
assert!(reject_non_finite_input(&samples, &Lang::En).is_err());
}
#[test]
fn reject_non_finite_input_passes_finite_audio() {
use crate::types::Lang;
let samples = vec![-1.0_f32, 0.0, 1.0, 1e10, -1e10];
assert!(reject_non_finite_input(&samples, &Lang::En).is_ok());
}
#[test]
fn new_accepts_matching_shape() {
let lp = LogProbsTV::new(2, 3, vec![-1.0, -2.0, -3.0, -4.0, -5.0, -6.0]).expect("2 * 3 == 6");
assert_eq!(lp.t(), 2);
assert_eq!(lp.v(), 3);
assert_eq!(lp.data(), &[-1.0, -2.0, -3.0, -4.0, -5.0, -6.0]);
}
#[test]
fn new_rejects_undersized_buffer() {
let Err(LogProbsError::Shape(err)) = LogProbsTV::new(2, 3, vec![0.0_f32; 5]) else {
panic!("t=2, v=3, data.len()=5 must be rejected as a shape mismatch");
};
assert_eq!(err.t(), 2);
assert_eq!(err.v(), 3);
assert_eq!(err.data_len(), 5);
assert!(err.to_string().contains("t=2"));
assert!(err.to_string().contains("v=3"));
assert!(err.to_string().contains("data.len()=5"));
}
#[test]
fn new_rejects_oversized_buffer() {
let Err(LogProbsError::Shape(err)) = LogProbsTV::new(2, 3, vec![0.0_f32; 7]) else {
panic!("t=2, v=3, data.len()=7 must be rejected as a shape mismatch");
};
assert_eq!(err.data_len(), 7);
}
#[test]
fn new_rejects_t_v_overflow() {
let big = usize::MAX / 2 + 1;
let Err(LogProbsError::Shape(err)) = LogProbsTV::new(big, big, Vec::new()) else {
panic!("t * v overflowing usize must be rejected, not silently wrapped");
};
assert_eq!(err.t(), big);
assert_eq!(err.v(), big);
assert_eq!(err.data_len(), 0);
}
#[test]
fn new_rejects_zero_vocab_with_zero_t() {
let Err(LogProbsError::Shape(err)) = LogProbsTV::new(0, 0, Vec::new()) else {
panic!("t=0, v=0 must be rejected: a CTC vocabulary needs at least the blank token");
};
assert_eq!(err.t(), 0);
assert_eq!(err.v(), 0);
assert_eq!(err.data_len(), 0);
let message = err.to_string();
assert!(
message.contains("zero-length vocab"),
"message must call out the zero-length vocab dim, not a shape \
mismatch (t * v == data.len() actually holds here); got {message}"
);
assert!(
!message.contains("expected data.len()"),
"must not reuse the shape-mismatch wording, which would falsely \
claim t * v != data.len(); got {message}"
);
}
#[test]
fn new_rejects_zero_vocab_with_positive_t() {
let Err(LogProbsError::Shape(err)) = LogProbsTV::new(1, 0, Vec::new()) else {
panic!("t=1, v=0 must be rejected: a CTC vocabulary needs at least the blank token");
};
assert_eq!(err.t(), 1);
assert_eq!(err.v(), 0);
assert_eq!(err.data_len(), 0);
assert!(err.to_string().contains("zero-length vocab"));
}
#[test]
fn new_rejects_nan_from_codex_failing_history() {
let Err(LogProbsError::Value(err)) = LogProbsTV::new(1, 2, vec![f32::NAN, 0.0]) else {
panic!("a NaN emission must be rejected as a value-domain error, never accepted");
};
assert_eq!(err.frame(), 0);
assert_eq!(err.vocab_index(), 0);
assert_eq!(err.class(), LogProbsValueClass::Nan);
}
#[test]
fn new_rejects_positive_infinity_in_token_column() {
let mut data = vec![-1.0_f32, -2.0, -3.0, -4.0];
data[3] = f32::INFINITY;
let Err(LogProbsError::Value(err)) = LogProbsTV::new(2, 2, data) else {
panic!("a +inf emission must be rejected");
};
assert_eq!(err.frame(), 1);
assert_eq!(err.vocab_index(), 1);
assert_eq!(err.class(), LogProbsValueClass::PosInf);
}
#[test]
fn new_rejects_negative_infinity_hard_mask_value() {
let mut data = vec![-1.0_f32, -2.0, -3.0, -4.0];
data[2] = f32::NEG_INFINITY;
let Err(LogProbsError::Value(err)) = LogProbsTV::new(2, 2, data) else {
panic!("a -inf emission must be rejected");
};
assert_eq!(err.frame(), 1);
assert_eq!(err.vocab_index(), 0);
assert_eq!(err.class(), LogProbsValueClass::NegInf);
}
#[test]
fn new_rejects_nan_in_final_frame_blank_column() {
let mut data = vec![-1.0_f32; 6];
data[4] = f32::NAN;
let Err(LogProbsError::Value(err)) = LogProbsTV::new(3, 2, data) else {
panic!("a NaN in the final-frame blank column must be rejected");
};
assert_eq!(err.frame(), 2);
assert_eq!(err.vocab_index(), 0);
assert_eq!(err.class(), LogProbsValueClass::Nan);
}
#[test]
fn new_rejects_tiny_positive_value() {
let Err(LogProbsError::Value(err)) = LogProbsTV::new(1, 2, vec![1.0e-7_f32, -0.5]) else {
panic!("a finite positive value must be rejected as out-of-domain");
};
assert_eq!(err.frame(), 0);
assert_eq!(err.vocab_index(), 0);
assert_eq!(err.class(), LogProbsValueClass::Positive);
}
#[test]
fn new_rejects_f32_max_from_codex_failing_history() {
let Err(LogProbsError::Value(err)) = LogProbsTV::new(1, 2, vec![f32::MAX, -1.0]) else {
panic!("f32::MAX (finite but > 0) must be rejected as out-of-domain");
};
assert_eq!(err.frame(), 0);
assert_eq!(err.vocab_index(), 0);
assert_eq!(err.class(), LogProbsValueClass::Positive);
}
#[test]
fn new_accepts_zero_and_negative_zero() {
let lp =
LogProbsTV::new(1, 3, vec![0.0_f32, -0.0, -1.0]).expect("0.0 and -0.0 are ≤ 0, so accepted");
assert_eq!(lp.t(), 1);
assert_eq!(lp.v(), 3);
assert_eq!(lp.at(0, 0), 0.0);
assert_eq!(lp.at(0, 1), -0.0);
}
#[test]
fn log_probs_error_display_is_transparent() {
let Err(shape) = LogProbsTV::new(2, 3, vec![0.0_f32; 5]) else {
panic!("shape mismatch expected");
};
let LogProbsError::Shape(inner) = shape else {
panic!("expected the Shape arm");
};
assert_eq!(LogProbsError::Shape(inner).to_string(), inner.to_string());
let Err(value) = LogProbsTV::new(1, 1, vec![f32::NAN]) else {
panic!("value error expected");
};
let LogProbsError::Value(inner) = value else {
panic!("expected the Value arm");
};
assert_eq!(LogProbsError::Value(inner).to_string(), inner.to_string());
assert!(inner.to_string().contains("out of domain"));
}
#[test]
fn new_accepts_zero_t_with_positive_vocab() {
let lp = LogProbsTV::new(0, 5, Vec::new()).expect("t=0 with v=5 and an empty buffer is valid");
assert_eq!(lp.t(), 0);
assert_eq!(lp.v(), 5);
assert!(lp.data().is_empty());
}
#[test]
fn at_indexes_correctly() {
let lp = LogProbsTV {
t: 2,
v: 3,
data: vec![-1.0, -2.0, -3.0, -4.0, -5.0, -6.0],
};
assert_eq!(lp.at(0, 0), -1.0);
assert_eq!(lp.at(0, 2), -3.0);
assert_eq!(lp.at(1, 0), -4.0);
assert_eq!(lp.at(1, 2), -6.0);
}
fn two_by_three() -> LogProbsTV {
LogProbsTV::new(2, 3, vec![-1.0, -2.0, -3.0, -4.0, -5.0, -6.0])
.expect("2 * 3 == 6 and every value is a valid log-probability")
}
#[test]
fn get_is_total_over_its_argument_space() {
let lp = two_by_three();
assert_eq!(lp.get(0, 0), Some(-1.0));
assert_eq!(lp.get(1, 2), Some(-6.0));
assert_eq!(lp.get(0, 3), None);
assert_eq!(lp.get(2, 0), None);
assert_eq!(lp.get(usize::MAX, 0), None);
assert_eq!(lp.get(usize::MAX, usize::MAX), None);
}
#[test]
#[should_panic(expected = "(t=0, v=3) is outside the (T=2, V=3) grid")]
fn at_rejects_vocab_index_aliasing_into_the_next_frame() {
let lp = two_by_three();
let _ = lp.at(0, 3);
}
#[test]
#[should_panic(expected = "is outside the (T=2, V=3) grid")]
fn at_rejects_frame_index_that_overflows_the_flat_index() {
let lp = two_by_three();
let _ = lp.at((usize::MAX / 3) + 1, 0);
}
#[test]
fn at_agrees_with_get_across_the_whole_grid() {
let lp = two_by_three();
for t_idx in 0..lp.t() {
for v_idx in 0..lp.v() {
let expected = lp.data()[t_idx * lp.v() + v_idx];
assert_eq!(lp.at(t_idx, v_idx), expected);
assert_eq!(lp.get(t_idx, v_idx), Some(expected));
}
}
}
#[test]
fn log_softmax_rejects_nan_logits_with_numeric_failure() {
let raw = vec![0.0_f32, f32::NAN, 0.0]; let err = log_softmax_with_finite_guard(&raw, 1, 3).unwrap_err();
match err {
EmissionsError::Numeric(payload) => {
let message = payload.message();
assert!(
message.contains("non-finite logit"),
"message must call out the non-finite logit; got {message}"
);
assert!(message.contains("frame 0"));
assert!(message.contains("vocab 1"));
assert!(
!message.to_ascii_lowercase().contains("ort"),
"message must not attribute the failure to ORT specifically — \
`log_softmax_with_finite_guard` is also reachable ort-free under the \
`emissions` feature, where the non-finite value came from the caller's \
own encoder; got {message:?}"
);
}
other => panic!("expected EmissionsError::Numeric; got {other:?}"),
}
}
#[test]
fn log_softmax_rejects_positive_infinity_logits() {
let raw = vec![0.0_f32, f32::INFINITY, 0.0];
assert!(log_softmax_with_finite_guard(&raw, 1, 3).is_err());
}
#[test]
fn log_softmax_rejects_negative_infinity_logits() {
let raw = vec![f32::NEG_INFINITY, 0.0, 0.0];
assert!(log_softmax_with_finite_guard(&raw, 1, 3).is_err());
}
#[test]
fn log_softmax_rejects_all_neg_infinity_row() {
let raw = vec![f32::NEG_INFINITY; 3];
assert!(log_softmax_with_finite_guard(&raw, 1, 3).is_err());
}
#[test]
fn log_softmax_rejects_finite_extremes_that_overflow_lp() {
let raw = vec![f32::MAX, -f32::MAX];
let err = log_softmax_with_finite_guard(&raw, 1, 2).unwrap_err();
let EmissionsError::Numeric(payload) = err else {
panic!("expected EmissionsError::Numeric; got {err:?}");
};
let message = payload.message();
assert!(
message.contains("log-softmax output non-finite"),
"diagnostic must call out the per-element finite check; got {message}",
);
}
#[test]
fn log_softmax_finite_input_roundtrips() {
let raw = vec![1.0_f32, 2.0, 3.0];
let out = log_softmax_with_finite_guard(&raw, 1, 3).expect("ok");
assert_eq!(out.len(), 3);
assert!(out.iter().all(|x| x.is_finite()));
let sum: f32 = out.iter().map(|x| x.exp()).sum();
assert!((sum - 1.0).abs() < 1e-5);
}
#[test]
fn log_softmax_large_common_offset_normalises_to_unit_exp_sum() {
let raw = vec![1.0e20_f32, 1.0e20_f32];
let out = log_softmax_with_finite_guard(&raw, 1, 2).expect("ok");
assert_eq!(out.len(), 2);
assert!(out.iter().all(|x| x.is_finite()));
let exp_sum: f32 = out.iter().map(|x| x.exp()).sum();
assert!(
(exp_sum - 1.0).abs() < 1e-5,
"exp(lp) sum must equal 1, got {exp_sum}; lps = {out:?}"
);
for lp in &out {
assert!(
(lp - (-(2.0_f32).ln())).abs() < 1e-4,
"expected ~{}, got {lp}",
-(2.0_f32).ln()
);
}
}
#[test]
fn log_softmax_rejects_undersized_buffer() {
let err = log_softmax_with_finite_guard(&[0.0f32, 0.0], 2, 2).unwrap_err();
let EmissionsError::Shape(payload) = err else {
panic!("expected EmissionsError::Shape; got {err:?}");
};
assert!(
payload.to_string().contains("shape mismatch"),
"diagnostic must call out the shape mismatch; got {payload}",
);
}
#[test]
fn log_softmax_rejects_empty_buffer_with_nonzero_dims() {
let err = log_softmax_with_finite_guard(&[], 1, 1).unwrap_err();
let EmissionsError::Shape(payload) = err else {
panic!("expected EmissionsError::Shape; got {err:?}");
};
assert!(payload.to_string().contains("shape mismatch"));
}
#[test]
fn log_softmax_rejects_oversized_buffer_instead_of_silently_truncating() {
let err = log_softmax_with_finite_guard(&[0.0f32, 99.0], 1, 1).unwrap_err();
let EmissionsError::Shape(payload) = err else {
panic!("expected EmissionsError::Shape; got {err:?}");
};
assert!(
payload.to_string().contains("shape mismatch"),
"diagnostic must call out the shape mismatch; got {payload}",
);
}
#[test]
fn log_softmax_rejects_zero_t_with_nonempty_buffer_instead_of_silently_discarding() {
let err = log_softmax_with_finite_guard(&[1.0f32, 2.0, 3.0], 0, 5).unwrap_err();
let EmissionsError::Shape(payload) = err else {
panic!("expected EmissionsError::Shape; got {err:?}");
};
assert!(payload.to_string().contains("shape mismatch"));
}
#[test]
fn log_softmax_accepts_zero_t_with_empty_buffer() {
let out = log_softmax_with_finite_guard(&[], 0, 5).expect("ok");
assert!(out.is_empty());
}
#[test]
fn log_softmax_rejects_zero_vocab_with_zero_t() {
let err = log_softmax_with_finite_guard(&[], 0, 0).unwrap_err();
let EmissionsError::Shape(payload) = err else {
panic!("expected EmissionsError::Shape; got {err:?}");
};
assert!(
payload.to_string().contains("zero-length vocab"),
"diagnostic must call out the zero-length vocab dim; got {payload}",
);
}
#[test]
fn log_softmax_rejects_zero_vocab_with_positive_t() {
let err = log_softmax_with_finite_guard(&[], 1, 0).unwrap_err();
let EmissionsError::Shape(payload) = err else {
panic!("expected EmissionsError::Shape; got {err:?}");
};
assert!(
payload.to_string().contains("zero-length vocab"),
"diagnostic must call out the zero-length vocab dim, not the \
unrelated shifted-normaliser-non-finite path it used to fall \
through to; got {payload}",
);
}
#[test]
fn log_softmax_rejects_t_v_product_overflow() {
let err = log_softmax_with_finite_guard(&[], usize::MAX, 2).unwrap_err();
let EmissionsError::Shape(payload) = err else {
panic!("expected EmissionsError::Shape; got {err:?}");
};
assert!(
payload.to_string().contains("shape mismatch"),
"an unrepresentable T*V product is rejected as a shape error; got {payload}",
);
}
#[test]
fn validate_output_dims_rejects_negative_t() {
use crate::types::Lang;
let err = validate_output_dims(-1, 32, 32, &Lang::En).unwrap_err();
let WorkFailure::Alignment(AlignmentError::ModelInference(payload)) = err else {
panic!("expected AlignmentFailed");
};
let message = payload.message();
assert!(message.contains("negative time dim"));
}
#[test]
fn validate_output_dims_rejects_zero_v() {
use crate::types::Lang;
let err = validate_output_dims(100, 0, 0, &Lang::En).unwrap_err();
assert!(matches!(
err,
WorkFailure::Alignment(AlignmentError::ModelInference(_))
));
}
#[test]
fn validate_output_dims_zero_t_with_empty_buffer_is_recoverable_no_alignment_path() {
use crate::types::Lang;
let err = validate_output_dims(0, 32, 0, &Lang::En).unwrap_err();
let WorkFailure::Alignment(AlignmentError::NoAlignmentPath(payload)) = &err else {
panic!("expected NoAlignmentPath; got {err:?}");
};
let message = payload.message();
assert!(
message.contains("zero encoder frames"),
"diagnostic must explain the short-chunk cause; got {message}",
message = message
);
}
#[test]
fn validate_output_dims_zero_t_with_nonempty_buffer_stays_fatal() {
use crate::types::Lang;
let err = validate_output_dims(0, 32, 5, &Lang::En).unwrap_err();
let WorkFailure::Alignment(AlignmentError::ModelInference(payload)) = &err else {
panic!("T=0 with non-empty buffer must stay fatal; got {err:?}");
};
let message = payload.message();
assert!(
message.contains("shape/data mismatch") || message.contains("buffer has"),
"diagnostic must call out the shape/data inconsistency; got {message}",
message = message
);
}
#[test]
fn validate_stride_extent_accepts_typical_under_extent() {
use crate::types::Lang;
assert!(validate_stride_extent(49, 320, 16_000, &Lang::En).is_ok());
assert!(validate_stride_extent(50, 320, 16_000, &Lang::En).is_ok());
assert!(validate_stride_extent(51, 320, 16_000, &Lang::En).is_ok());
}
#[test]
fn validate_stride_extent_rejects_t_too_large() {
use crate::types::Lang;
let err = validate_stride_extent(100, 320, 16_000, &Lang::En).unwrap_err();
let WorkFailure::Alignment(AlignmentError::ModelInference(payload)) = err else {
panic!("expected AlignmentFailed");
};
let message = payload.message();
assert!(
message.contains("smaller stride"),
"diagnostic must call out the smaller-stride case; got {message}",
message = message
);
}
#[test]
fn validate_stride_extent_rejects_t_too_small() {
use crate::types::Lang;
let err = validate_stride_extent(25, 320, 16_000, &Lang::En).unwrap_err();
let WorkFailure::Alignment(AlignmentError::ModelInference(payload)) = err else {
panic!("expected AlignmentFailed");
};
let message = payload.message();
assert!(
message.contains("larger stride"),
"diagnostic must call out the larger-stride case; got {message}",
message = message
);
}
#[test]
fn validate_stride_extent_accepts_very_short_chunk_with_small_t() {
use crate::types::Lang;
assert!(validate_stride_extent(1, 320, 200, &Lang::En).is_ok());
}
#[test]
fn validate_vocab_dim_accepts_exact_match() {
use crate::types::Lang;
assert!(validate_vocab_dim(32, 32, &Lang::En).is_ok());
}
#[test]
fn validate_vocab_dim_rejects_oversized_model_output() {
use crate::types::Lang;
let err = validate_vocab_dim(1024, 32, &Lang::En).unwrap_err();
let WorkFailure::Alignment(AlignmentError::ModelInference(payload)) = err else {
panic!("expected AlignmentFailed");
};
let message = payload.message();
assert!(
message.contains("doesn't match tokenizer vocab"),
"diagnostic must call out the vocab mismatch; got {message}",
message = message
);
}
#[test]
fn validate_vocab_dim_rejects_undersized_model_output() {
use crate::types::Lang;
let err = validate_vocab_dim(16, 32, &Lang::En).unwrap_err();
assert!(matches!(
err,
WorkFailure::Alignment(AlignmentError::ModelInference(_))
));
}
#[test]
fn validate_output_dims_rejects_buffer_length_mismatch() {
use crate::types::Lang;
let err = validate_output_dims(10, 4, 39, &Lang::En).unwrap_err();
let WorkFailure::Alignment(payload) = err else {
panic!("expected AlignmentFailed");
};
let message = payload.to_string();
assert!(
message.contains("doesn't match"),
"must call out length mismatch; got {message}",
message = message
);
}
#[test]
fn validate_output_dims_rejects_t_v_product_overflow() {
use crate::types::Lang;
let big = i64::from(u32::MAX) + 1; let err = validate_output_dims(big, big, 0, &Lang::En).unwrap_err();
let WorkFailure::Alignment(payload) = err else {
panic!("expected AlignmentFailed");
};
let message = payload.to_string();
assert!(
message.contains("overflow") || message.contains("doesn't fit"),
"must call out overflow; got {message}",
message = message
);
}
#[test]
fn validate_output_dims_accepts_well_formed_shape() {
use crate::types::Lang;
let (t, v) = validate_output_dims(1500, 32, 1500 * 32, &Lang::En).expect("ok");
assert_eq!(t, 1500);
assert_eq!(v, 32);
}
#[test]
fn log_softmax_locates_nan_to_specific_frame() {
let raw = vec![0.0_f32, 0.1, 0.0, 0.1, f32::NAN, 0.1];
let err = log_softmax_with_finite_guard(&raw, 3, 2).unwrap_err();
let EmissionsError::Numeric(payload) = err else {
panic!("expected EmissionsError::Numeric; got {err:?}");
};
let message = payload.to_string();
assert!(
message.contains("frame 2"),
"must locate the bad frame; got {message}",
message = message
);
}
}