use ndarray::{Array1, Array2, Array3, Array4, ArrayView2, Axis};
use ort::session::Session;
use ort::value::Value;
use std::path::Path;
use std::sync::Mutex;
use crate::traits::AsrError;
use super::vocab::Vocab;
pub const DECODER_INPUT_IDS: &str = "input_ids";
pub const DECODER_INPUT_ENCODER_EMBEDDINGS: &str = "encoder_embeddings";
pub const DECODER_INPUT_ENCODER_MASK: &str = "encoder_mask";
pub const DECODER_INPUT_DECODER_MEMS: &str = "decoder_mems";
pub const DECODER_OUTPUT_LOGITS: &str = "logits";
pub const DECODER_OUTPUT_HIDDEN_STATES: &str = "decoder_hidden_states";
pub const PREFIX_LEN: usize = 10;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PrefixFormat {
OnnxAsr,
NemoCanary2,
}
impl PrefixFormat {
pub fn token_count(&self) -> usize {
match self {
PrefixFormat::OnnxAsr => 10,
PrefixFormat::NemoCanary2 => 9,
}
}
}
pub const DEFAULT_PREFIX_FORMAT: PrefixFormat = PrefixFormat::NemoCanary2;
pub const DEFAULT_MAX_SEQUENCE_LENGTH: usize = 1024;
pub const DEFAULT_REPETITION_PENALTY: f32 = 2.0;
pub const DEFAULT_MIN_TOKEN_TO_FRAME_RATIO: f32 = 0.2;
pub const DEFAULT_EOS_CONFIDENCE_MARGIN: f32 = 2.0;
pub const DEFAULT_BEAM_SIZE: usize = 1;
pub const DEFAULT_LENGTH_PENALTY: f32 = 0.6;
#[derive(Debug, Clone)]
pub struct DecodeOptions {
pub source_language: String,
pub target_language: String,
pub pnc: bool,
pub max_sequence_length: usize,
pub repetition_penalty: f32,
pub min_token_to_frame_ratio: f32,
pub eos_confidence_margin: f32,
pub beam_size: usize,
pub length_penalty: f32,
pub prefix_format: PrefixFormat,
}
impl DecodeOptions {
pub fn for_asr(language: impl Into<String>) -> Self {
let lang = language.into();
Self {
source_language: lang.clone(),
target_language: lang,
pnc: true,
max_sequence_length: DEFAULT_MAX_SEQUENCE_LENGTH,
repetition_penalty: DEFAULT_REPETITION_PENALTY,
min_token_to_frame_ratio: DEFAULT_MIN_TOKEN_TO_FRAME_RATIO,
eos_confidence_margin: DEFAULT_EOS_CONFIDENCE_MARGIN,
beam_size: DEFAULT_BEAM_SIZE,
length_penalty: DEFAULT_LENGTH_PENALTY,
prefix_format: DEFAULT_PREFIX_FORMAT,
}
}
}
#[derive(Debug, Clone)]
pub struct DecodeOutput {
pub tokens: Vec<u32>,
pub logprobs: Vec<f32>,
}
pub fn build_decoder_prefix(vocab: &Vocab, opts: &DecodeOptions) -> Result<Vec<i64>, AsrError> {
let space = vocab.last_id("\u{2581}").ok_or_else(|| AsrError::ModelLoad("vocab missing ▁ (used for the literal-space slot in the \
decoder prefix; onnx-asr maps `_tokens[\" \"]` to the \
last-occurring ▁ token id)"
.into()))?;
let soc = vocab.soc()?;
let sot = vocab.sot()?;
let emo = vocab.id("<|emo:undefined|>").ok_or_else(|| AsrError::ModelLoad("vocab missing <|emo:undefined|>".into()))?;
let src_lang = vocab
.language_token(&opts.source_language)
.ok_or_else(|| AsrError::ModelLoad(format!(
"vocab has no language token for source_language={:?}",
opts.source_language
)))?;
let tgt_lang = vocab
.language_token(&opts.target_language)
.ok_or_else(|| AsrError::ModelLoad(format!(
"vocab has no language token for target_language={:?}",
opts.target_language
)))?;
let pnc = if opts.pnc {
vocab.pnc()?
} else {
vocab.nopnc()?
};
let noitn = vocab.id("<|noitn|>").ok_or_else(|| AsrError::ModelLoad("vocab missing <|noitn|>".into()))?;
let notimestamp = vocab.id("<|notimestamp|>").ok_or_else(|| AsrError::ModelLoad("vocab missing <|notimestamp|>".into()))?;
let nodiarize = vocab.id("<|nodiarize|>").ok_or_else(|| AsrError::ModelLoad("vocab missing <|nodiarize|>".into()))?;
let prefix: Vec<u32> = match opts.prefix_format {
PrefixFormat::OnnxAsr => vec![
space,
soc,
sot,
emo,
src_lang,
tgt_lang,
pnc,
noitn,
notimestamp,
nodiarize,
],
PrefixFormat::NemoCanary2 => vec![
soc,
sot,
emo,
src_lang,
tgt_lang,
pnc,
noitn,
notimestamp,
nodiarize,
],
};
debug_assert_eq!(prefix.len(), opts.prefix_format.token_count());
Ok(prefix.into_iter().map(|x| x as i64).collect())
}
pub fn argmax_last_position(logits: &Array3<f32>) -> Vec<u32> {
let (batch, time, vocab_size) = logits.dim();
let mut out = Vec::with_capacity(batch);
for b in 0..batch {
let mut best_v = u32::MAX;
let mut best_score = f32::NEG_INFINITY;
for v in 0..vocab_size {
let score = logits[[b, time - 1, v]];
if score > best_score {
best_score = score;
best_v = v as u32;
}
}
out.push(best_v);
}
out
}
pub fn suppress_eos_until_min_length(
logits: &mut Array3<f32>,
eos_token_id: u32,
suffix_len: usize,
n_encoder_frames: usize,
ratio: f32,
) {
if ratio <= 0.0 || n_encoder_frames == 0 {
return;
}
let min_len = (ratio * n_encoder_frames as f32).ceil() as usize;
if suffix_len >= min_len {
return;
}
let (batch, time, vocab_size) = logits.dim();
if batch == 0 || time == 0 {
return;
}
let v = eos_token_id as usize;
if v >= vocab_size {
return;
}
let last_t = time - 1;
for b in 0..batch {
logits[[b, last_t, v]] = f32::NEG_INFINITY;
}
}
pub fn valid_frame_count(mask: &Array2<i64>, batch_idx: usize) -> usize {
if batch_idx >= mask.shape()[0] {
return 0;
}
mask.row(batch_idx).iter().filter(|&&v| v != 0).count()
}
pub fn enforce_eos_confidence_margin(logits: &mut Array3<f32>, eos_token_id: u32, margin: f32) {
if margin <= 0.0 {
return;
}
let (batch, time, vocab_size) = logits.dim();
if batch == 0 || time == 0 {
return;
}
let v = eos_token_id as usize;
if v >= vocab_size {
return;
}
let last_t = time - 1;
for b in 0..batch {
let eos_logit = logits[[b, last_t, v]];
if !eos_logit.is_finite() {
continue;
}
let mut max_other = f32::NEG_INFINITY;
for vi in 0..vocab_size {
if vi == v {
continue;
}
let l = logits[[b, last_t, vi]];
if l > max_other {
max_other = l;
}
}
if !max_other.is_finite() {
continue;
}
if eos_logit - max_other < margin {
logits[[b, last_t, v]] = f32::NEG_INFINITY;
}
}
}
pub fn apply_repetition_penalty(logits: &mut Array3<f32>, history: &[i64], penalty: f32) {
if penalty == 1.0 || history.is_empty() {
return;
}
let (batch, time, vocab_size) = logits.dim();
if batch == 0 || time == 0 {
return;
}
let last_t = time - 1;
let mut seen = std::collections::HashSet::<i64>::new();
for &tok in history {
if tok < 0 || !seen.insert(tok) {
continue;
}
let v = tok as usize;
if v >= vocab_size {
continue;
}
for b in 0..batch {
let score = logits[[b, last_t, v]];
let new_score = if score < 0.0 {
score * penalty
} else {
score / penalty
};
logits[[b, last_t, v]] = new_score;
}
}
}
pub fn strip_prefix_and_specials(all_tokens: &[u32], prefix_len: usize, vocab: &Vocab) -> Vec<u32> {
if all_tokens.len() <= prefix_len {
return Vec::new();
}
all_tokens[prefix_len..]
.iter()
.filter(|&&id| match vocab.piece(id) {
Some(p) => !(p.starts_with("<|") && p.ends_with("|>")),
None => false,
})
.copied()
.collect()
}
pub struct CanaryDecoder {
session: Mutex<Session>,
mems_layers: usize,
mems_hidden: usize,
profiling: bool,
}
impl CanaryDecoder {
pub fn load(path: impl AsRef<Path>) -> Result<Self, AsrError> {
let path = path.as_ref();
let builder = Session::builder().map_err(|e| AsrError::ModelLoad(format!("Canary decoder builder {}: {e}", path.display())))?;
let (mut builder, profiling) = crate::canary::profiling::apply(builder, "decoder")
.map_err(|e| AsrError::Inference(format!("Canary decoder profiling {}: {e}", path.display())))?;
let session = builder.commit_from_file(path).map_err(|e| AsrError::ModelLoad(format!("load Canary decoder {}: {e}", path.display())))?;
validate_decoder_io(&session, path)?;
let (mems_layers, mems_hidden) = read_mems_shape(&session, path)?;
Ok(Self {
session: Mutex::new(session),
mems_layers,
mems_hidden,
profiling,
})
}
pub fn decode(
&self,
encoder_embeddings: &Array3<f32>,
encoder_mask: &Array2<i64>,
vocab: &Vocab,
opts: &DecodeOptions,
) -> Result<DecodeOutput, AsrError> {
if encoder_embeddings.shape()[0] != 1 {
return Err(AsrError::Inference(format!(
"decoder currently supports batch_size=1 (got {})",
encoder_embeddings.shape()[0]
)));
}
if encoder_mask.shape()[0] != 1 {
return Err(AsrError::Inference(format!(
"encoder_mask batch must match encoder_embeddings (got {})",
encoder_mask.shape()[0]
)));
}
if opts.beam_size > 1 {
return self.decode_beam(encoder_embeddings, encoder_mask, vocab, opts);
}
let prefix = build_decoder_prefix(vocab, opts)?;
let prefix_len = prefix.len();
let max_len = opts.max_sequence_length.max(prefix_len + 1);
let eos = vocab.eos()?;
let mut batch_tokens: Vec<i64> = prefix;
let mut logprobs: Vec<f32> = Vec::new();
let mut decoder_mems: Array4<f32> =
Array4::zeros((self.mems_layers, 1, 0, self.mems_hidden));
let mut session = self.session.lock().map_err(|e| AsrError::Inference(format!("decoder session lock poisoned: {e}")))?;
while batch_tokens.len() < max_len {
let input_ids: Array2<i64> = if decoder_mems.shape()[2] == 0 {
Array2::from_shape_vec((1, batch_tokens.len()), batch_tokens.clone()).map_err(
|e| AsrError::Inference(format!("input_ids reshape (initial): {e}")),
)?
} else {
let last = *batch_tokens.last().unwrap();
Array2::from_shape_vec((1, 1), vec![last]).map_err(|e| AsrError::Inference(format!("input_ids reshape (step): {e}")))?
};
let input_ids_v = Value::from_array(input_ids).map_err(|e| AsrError::Inference(format!("input_ids Value: {e}")))?;
let enc_emb_v =
Value::from_array(encoder_embeddings.clone()).map_err(|e| AsrError::Inference(format!("encoder_embeddings Value: {e}")))?;
let enc_mask_v = Value::from_array(encoder_mask.clone()).map_err(|e| AsrError::Inference(format!("encoder_mask Value: {e}")))?;
let mems_v = Value::from_array(decoder_mems.clone()).map_err(|e| AsrError::Inference(format!("decoder_mems Value: {e}")))?;
let outputs = session
.run(vec![
(DECODER_INPUT_IDS, input_ids_v.into_dyn()),
(DECODER_INPUT_ENCODER_EMBEDDINGS, enc_emb_v.into_dyn()),
(DECODER_INPUT_ENCODER_MASK, enc_mask_v.into_dyn()),
(DECODER_INPUT_DECODER_MEMS, mems_v.into_dyn()),
])
.map_err(|e| AsrError::Inference(format!("Canary decoder run: {e}")))?;
let logits_idx =
output_index(&outputs, DECODER_OUTPUT_LOGITS).ok_or_else(|| AsrError::Inference(format!("decoder missing output {DECODER_OUTPUT_LOGITS}")))?;
let mems_idx =
output_index(&outputs, DECODER_OUTPUT_HIDDEN_STATES).ok_or_else(|| AsrError::Inference(format!("decoder missing output {DECODER_OUTPUT_HIDDEN_STATES}")))?;
let mut logits: Array3<f32> = outputs[logits_idx]
.try_extract_array::<f32>()
.map_err(|e| AsrError::Inference(format!("extract {DECODER_OUTPUT_LOGITS}: {e}")))?
.to_owned()
.into_dimensionality::<ndarray::Ix3>()
.map_err(|e| AsrError::Inference(format!("{DECODER_OUTPUT_LOGITS} rank: {e}")))?;
let new_mems: Array4<f32> = outputs[mems_idx]
.try_extract_array::<f32>()
.map_err(|e| AsrError::Inference(format!("extract {DECODER_OUTPUT_HIDDEN_STATES}: {e}")))?
.to_owned()
.into_dimensionality::<ndarray::Ix4>()
.map_err(|e| AsrError::Inference(format!("{DECODER_OUTPUT_HIDDEN_STATES} rank: {e}")))?;
decoder_mems = new_mems;
let suffix_len = batch_tokens.len().saturating_sub(prefix_len);
if opts.repetition_penalty != 1.0 && suffix_len > 0 {
let suffix = &batch_tokens[prefix_len..];
apply_repetition_penalty(&mut logits, suffix, opts.repetition_penalty);
}
if opts.min_token_to_frame_ratio > 0.0 {
let n_enc_frames = valid_frame_count(encoder_mask, 0);
suppress_eos_until_min_length(
&mut logits,
eos,
suffix_len,
n_enc_frames,
opts.min_token_to_frame_ratio,
);
}
if opts.eos_confidence_margin > 0.0 {
enforce_eos_confidence_margin(&mut logits, eos, opts.eos_confidence_margin);
}
let next_tokens = argmax_last_position(&logits);
let next = next_tokens[0];
if next == eos {
break;
}
let next_logprob = logprob_of_token(&logits, 0, next);
batch_tokens.push(next as i64);
logprobs.push(next_logprob);
}
let all_ids: Vec<u32> = batch_tokens.iter().map(|&x| x as u32).collect();
let kept = strip_prefix_and_specials(&all_ids, prefix_len, vocab);
let kept_logprobs: Vec<f32> = all_ids[prefix_len..]
.iter()
.zip(logprobs.iter())
.filter(|(id, _)| match vocab.piece(**id) {
Some(p) => !(p.starts_with("<|") && p.ends_with("|>")),
None => false,
})
.map(|(_, lp)| *lp)
.collect();
Ok(DecodeOutput {
tokens: kept,
logprobs: kept_logprobs,
})
}
}
impl Drop for CanaryDecoder {
fn drop(&mut self) {
if self.profiling {
crate::canary::profiling::flush(&self.session, "decoder");
}
}
}
fn logprob_of_token(logits: &Array3<f32>, batch: usize, token: u32) -> f32 {
let (_, time, vocab_size) = logits.dim();
let row: ArrayView2<f32> = logits.index_axis(Axis(0), batch);
let last = row.index_axis(Axis(0), time - 1);
let max = last.iter().copied().fold(f32::NEG_INFINITY, f32::max);
if !max.is_finite() {
return f32::NEG_INFINITY;
}
let sum_exp: f32 = (0..vocab_size).map(|v| (last[v] - max).exp()).sum();
let logsumexp = max + sum_exp.ln();
last[token as usize] - logsumexp
}
fn log_softmax_row(logits: &[f32]) -> Vec<f32> {
let mut max = f32::NEG_INFINITY;
for &v in logits {
if v > max {
max = v;
}
}
if !max.is_finite() {
return vec![f32::NEG_INFINITY; logits.len()];
}
let mut sum_exp = 0.0_f32;
for &v in logits {
sum_exp += (v - max).exp();
}
let logsumexp = max + sum_exp.ln();
logits.iter().map(|&v| v - logsumexp).collect()
}
fn topk_indices(xs: &[f32], k: usize) -> Vec<(usize, f32)> {
let mut indexed: Vec<(usize, f32)> = xs.iter().enumerate().map(|(i, &v)| (i, v)).collect();
indexed.sort_by(|a, b| {
b.1.partial_cmp(&a.1)
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.0.cmp(&b.0))
});
indexed.truncate(k);
indexed
}
fn length_normalised_score(cum_logprob: f32, length: usize, alpha: f32) -> f32 {
if length == 0 {
return cum_logprob;
}
if alpha == 0.0 {
return cum_logprob;
}
cum_logprob / (length as f32).powf(alpha)
}
#[derive(Clone)]
struct ActiveBeam {
tokens: Vec<i64>,
cum_logprob: f32,
}
#[derive(Clone)]
struct FinishedBeam {
tokens: Vec<i64>,
cum_logprob: f32,
length: usize,
}
impl CanaryDecoder {
fn decode_beam(
&self,
encoder_embeddings: &Array3<f32>,
encoder_mask: &Array2<i64>,
vocab: &Vocab,
opts: &DecodeOptions,
) -> Result<DecodeOutput, AsrError> {
let beam_size = opts.beam_size;
if beam_size < 2 {
return Err(AsrError::Config(format!("decode_beam requires beam_size ≥ 2 (got {beam_size})")));
}
let prefix = build_decoder_prefix(vocab, opts)?;
let prefix_len = prefix.len();
let max_len = opts.max_sequence_length.max(prefix_len + 1);
let eos = vocab.eos()?;
let mut session = self.session.lock().map_err(|e| AsrError::Inference(format!("decoder session lock poisoned: {e}")))?;
let initial_input: Array2<i64> = Array2::from_shape_vec((1, prefix_len), prefix.clone())
.map_err(|e| AsrError::Inference(format!("input_ids reshape (initial): {e}")))?;
let initial_mems: Array4<f32> = Array4::zeros((self.mems_layers, 1, 0, self.mems_hidden));
let (init_logits, init_hidden) = run_decoder_step(
&mut session,
initial_input,
encoder_embeddings.clone(),
encoder_mask.clone(),
initial_mems,
)?;
let mut step0_logits = init_logits;
if opts.min_token_to_frame_ratio > 0.0 {
let n_enc_frames = valid_frame_count(encoder_mask, 0);
suppress_eos_until_min_length(
&mut step0_logits,
eos,
0,
n_enc_frames,
opts.min_token_to_frame_ratio,
);
}
if opts.eos_confidence_margin > 0.0 {
enforce_eos_confidence_margin(&mut step0_logits, eos, opts.eos_confidence_margin);
}
let last_t = step0_logits.shape()[1] - 1;
let vocab_size = step0_logits.shape()[2];
let last_row: Vec<f32> = (0..vocab_size)
.map(|v| step0_logits[[0, last_t, v]])
.collect();
let log_probs0 = log_softmax_row(&last_row);
let candidates0 = topk_indices(&log_probs0, beam_size);
let mut active: Vec<ActiveBeam> = Vec::with_capacity(beam_size);
for (vid, lp) in &candidates0 {
let mut tokens = prefix.clone();
tokens.push(*vid as i64);
active.push(ActiveBeam {
tokens,
cum_logprob: *lp,
});
}
let mut decoder_mems: Array4<f32> = stack_along_batch(&init_hidden, beam_size)?;
let mut finished: Vec<FinishedBeam> = Vec::new();
while !active.is_empty() && active[0].tokens.len() < max_len && finished.len() < beam_size {
let b = active.len();
let last_tokens: Vec<i64> = active.iter().map(|x| *x.tokens.last().unwrap()).collect();
let input_ids: Array2<i64> = Array2::from_shape_vec((b, 1), last_tokens.clone())
.map_err(|e| AsrError::Inference(format!("input_ids reshape (step): {e}")))?;
let enc_emb_b: Array3<f32> = repeat_along_axis0(encoder_embeddings, b);
let enc_mask_b: Array2<i64> = repeat_along_axis0_i64(encoder_mask, b);
let (mut logits, hidden) =
run_decoder_step(&mut session, input_ids, enc_emb_b, enc_mask_b, decoder_mems)?;
let suffix_len = active[0].tokens.len() - prefix_len;
let n_enc_frames = valid_frame_count(encoder_mask, 0);
for (bi, beam) in active.iter().enumerate() {
if opts.repetition_penalty != 1.0 {
apply_repetition_penalty_one_batch(
&mut logits,
bi,
&beam.tokens[prefix_len..],
opts.repetition_penalty,
);
}
if opts.min_token_to_frame_ratio > 0.0 {
suppress_eos_until_min_length_one_batch(
&mut logits,
bi,
eos,
suffix_len,
n_enc_frames,
opts.min_token_to_frame_ratio,
);
}
if opts.eos_confidence_margin > 0.0 {
enforce_eos_confidence_margin_one_batch(
&mut logits,
bi,
eos,
opts.eos_confidence_margin,
);
}
}
let last_t = logits.shape()[1] - 1;
let vocab_size = logits.shape()[2];
let mut all_candidates: Vec<(usize, i64, f32)> = Vec::with_capacity(b * beam_size);
for bi in 0..b {
let row: Vec<f32> = (0..vocab_size).map(|v| logits[[bi, last_t, v]]).collect();
let lp = log_softmax_row(&row);
let topk = topk_indices(&lp, beam_size);
for (vid, score) in topk {
all_candidates.push((bi, vid as i64, active[bi].cum_logprob + score));
}
}
all_candidates.sort_by(|a, c| {
c.2.partial_cmp(&a.2)
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.0.cmp(&c.0))
});
let mut next_active: Vec<ActiveBeam> = Vec::new();
let mut parent_idx_for_active: Vec<usize> = Vec::new();
for (parent, vid, cum) in all_candidates {
if next_active.len() + finished.len() >= beam_size + finished.len() {
}
if next_active.len() >= beam_size {
break;
}
let mut tokens = active[parent].tokens.clone();
if vid as u32 == eos {
finished.push(FinishedBeam {
tokens,
cum_logprob: active[parent].cum_logprob,
length: suffix_len,
});
if finished.len() >= beam_size {
break;
}
} else {
tokens.push(vid);
next_active.push(ActiveBeam {
tokens,
cum_logprob: cum,
});
parent_idx_for_active.push(parent);
}
}
if next_active.is_empty() {
break;
}
decoder_mems = hidden.select(Axis(1), &parent_idx_for_active);
active = next_active;
}
let alpha = opts.length_penalty;
let mut best: Option<(Vec<i64>, Vec<f32>)> = None;
let mut best_score = f32::NEG_INFINITY;
for f in &finished {
let score = length_normalised_score(f.cum_logprob, f.length, alpha);
if score > best_score {
best_score = score;
let logprobs = vec![0.0_f32; f.length];
best = Some((f.tokens.clone(), logprobs));
}
}
for a in &active {
let length = a.tokens.len() - prefix_len;
let score = length_normalised_score(a.cum_logprob, length, alpha);
if score > best_score {
best_score = score;
let logprobs = vec![0.0_f32; length];
best = Some((a.tokens.clone(), logprobs));
}
}
let (best_tokens, best_logprobs) = best.ok_or_else(|| AsrError::Inference("beam search produced no candidates".into()))?;
let all_ids: Vec<u32> = best_tokens.iter().map(|&x| x as u32).collect();
let kept = strip_prefix_and_specials(&all_ids, prefix_len, vocab);
let kept_logprobs: Vec<f32> = vec![0.0_f32; kept.len()];
let _ = best_logprobs;
Ok(DecodeOutput {
tokens: kept,
logprobs: kept_logprobs,
})
}
}
fn run_decoder_step(
session: &mut Session,
input_ids: Array2<i64>,
encoder_embeddings: Array3<f32>,
encoder_mask: Array2<i64>,
decoder_mems: Array4<f32>,
) -> Result<(Array3<f32>, Array4<f32>), AsrError> {
let input_ids_v = Value::from_array(input_ids).map_err(|e| AsrError::Inference(format!("input_ids Value: {e}")))?;
let enc_emb_v = Value::from_array(encoder_embeddings).map_err(|e| AsrError::Inference(format!("encoder_embeddings Value: {e}")))?;
let enc_mask_v = Value::from_array(encoder_mask).map_err(|e| AsrError::Inference(format!("encoder_mask Value: {e}")))?;
let mems_v = Value::from_array(decoder_mems).map_err(|e| AsrError::Inference(format!("decoder_mems Value: {e}")))?;
let outputs = session
.run(vec![
(DECODER_INPUT_IDS, input_ids_v.into_dyn()),
(DECODER_INPUT_ENCODER_EMBEDDINGS, enc_emb_v.into_dyn()),
(DECODER_INPUT_ENCODER_MASK, enc_mask_v.into_dyn()),
(DECODER_INPUT_DECODER_MEMS, mems_v.into_dyn()),
])
.map_err(|e| AsrError::Inference(format!("Canary decoder run: {e}")))?;
let logits_idx = output_index(&outputs, DECODER_OUTPUT_LOGITS).ok_or_else(|| AsrError::Inference(format!("decoder missing output {DECODER_OUTPUT_LOGITS}")))?;
let mems_idx =
output_index(&outputs, DECODER_OUTPUT_HIDDEN_STATES).ok_or_else(|| AsrError::Inference(format!("decoder missing output {DECODER_OUTPUT_HIDDEN_STATES}")))?;
let logits: Array3<f32> = outputs[logits_idx]
.try_extract_array::<f32>()
.map_err(|e| AsrError::Inference(format!("extract logits: {e}")))?
.to_owned()
.into_dimensionality()
.map_err(|e| AsrError::Inference(format!("logits rank: {e}")))?;
let hidden: Array4<f32> = outputs[mems_idx]
.try_extract_array::<f32>()
.map_err(|e| AsrError::Inference(format!("extract hidden: {e}")))?
.to_owned()
.into_dimensionality()
.map_err(|e| AsrError::Inference(format!("hidden rank: {e}")))?;
Ok((logits, hidden))
}
fn stack_along_batch(hidden: &Array4<f32>, n_copies: usize) -> Result<Array4<f32>, AsrError> {
let (l, b, t, h) = hidden.dim();
if b != 1 {
return Err(AsrError::Inference(format!("stack_along_batch expected batch=1, got {b}")));
}
let mut out = Array4::<f32>::zeros((l, n_copies, t, h));
for k in 0..n_copies {
out.slice_mut(ndarray::s![.., k..k + 1, .., ..])
.assign(hidden);
}
Ok(out)
}
fn repeat_along_axis0(arr: &Array3<f32>, n_copies: usize) -> Array3<f32> {
let shape = arr.shape();
let new_shape = (n_copies, shape[1], shape[2]);
let mut out = Array3::<f32>::zeros(new_shape);
for k in 0..n_copies {
out.slice_mut(ndarray::s![k..k + 1, .., ..]).assign(arr);
}
out
}
fn repeat_along_axis0_i64(arr: &Array2<i64>, n_copies: usize) -> Array2<i64> {
let shape = arr.shape();
let new_shape = (n_copies, shape[1]);
let mut out = Array2::<i64>::zeros(new_shape);
for k in 0..n_copies {
out.slice_mut(ndarray::s![k..k + 1, ..]).assign(arr);
}
out
}
fn apply_repetition_penalty_one_batch(
logits: &mut Array3<f32>,
batch: usize,
history: &[i64],
penalty: f32,
) {
if penalty == 1.0 || history.is_empty() {
return;
}
let (_, time, vocab_size) = logits.dim();
if time == 0 {
return;
}
let last_t = time - 1;
let mut seen = std::collections::HashSet::<i64>::new();
for &tok in history {
if tok < 0 || !seen.insert(tok) {
continue;
}
let v = tok as usize;
if v >= vocab_size {
continue;
}
let score = logits[[batch, last_t, v]];
let new_score = if score < 0.0 {
score * penalty
} else {
score / penalty
};
logits[[batch, last_t, v]] = new_score;
}
}
fn suppress_eos_until_min_length_one_batch(
logits: &mut Array3<f32>,
batch: usize,
eos_token_id: u32,
suffix_len: usize,
n_encoder_frames: usize,
ratio: f32,
) {
if ratio <= 0.0 || n_encoder_frames == 0 {
return;
}
let min_len = (ratio * n_encoder_frames as f32).ceil() as usize;
if suffix_len >= min_len {
return;
}
let (_, time, vocab_size) = logits.dim();
if time == 0 {
return;
}
let v = eos_token_id as usize;
if v >= vocab_size {
return;
}
let last_t = time - 1;
logits[[batch, last_t, v]] = f32::NEG_INFINITY;
}
fn enforce_eos_confidence_margin_one_batch(
logits: &mut Array3<f32>,
batch: usize,
eos_token_id: u32,
margin: f32,
) {
if margin <= 0.0 {
return;
}
let (_, time, vocab_size) = logits.dim();
if time == 0 {
return;
}
let v = eos_token_id as usize;
if v >= vocab_size {
return;
}
let last_t = time - 1;
let eos_logit = logits[[batch, last_t, v]];
if !eos_logit.is_finite() {
return;
}
let mut max_other = f32::NEG_INFINITY;
for vi in 0..vocab_size {
if vi == v {
continue;
}
let l = logits[[batch, last_t, vi]];
if l > max_other {
max_other = l;
}
}
if !max_other.is_finite() {
return;
}
if eos_logit - max_other < margin {
logits[[batch, last_t, v]] = f32::NEG_INFINITY;
}
}
fn validate_decoder_io(session: &Session, path: &Path) -> Result<(), AsrError> {
let want_in = [
DECODER_INPUT_IDS,
DECODER_INPUT_ENCODER_EMBEDDINGS,
DECODER_INPUT_ENCODER_MASK,
DECODER_INPUT_DECODER_MEMS,
];
let got_in: Vec<String> = session
.inputs()
.iter()
.map(|i| i.name().to_string())
.collect();
for name in &want_in {
if !got_in.iter().any(|n| n == name) {
return Err(AsrError::Inference(format!(
"decoder {} missing input {} (have: {got_in:?})",
path.display(),
name
)));
}
}
let want_out = [DECODER_OUTPUT_LOGITS, DECODER_OUTPUT_HIDDEN_STATES];
let got_out: Vec<String> = session
.outputs()
.iter()
.map(|o| o.name().to_string())
.collect();
for name in &want_out {
if !got_out.iter().any(|n| n == name) {
return Err(AsrError::Inference(format!(
"decoder {} missing output {} (have: {got_out:?})",
path.display(),
name
)));
}
}
Ok(())
}
fn read_mems_shape(session: &Session, path: &Path) -> Result<(usize, usize), AsrError> {
let input = session
.inputs()
.iter()
.find(|i| i.name() == DECODER_INPUT_DECODER_MEMS)
.ok_or_else(|| AsrError::Inference(format!(
"decoder {} missing input {DECODER_INPUT_DECODER_MEMS}",
path.display()
)))?;
let shape = input.dtype().tensor_shape().ok_or_else(|| AsrError::Inference(format!(
"decoder {} input {DECODER_INPUT_DECODER_MEMS} is not a tensor",
path.display()
)))?;
if shape.len() != 4 {
return Err(AsrError::Inference(format!(
"decoder {} input {DECODER_INPUT_DECODER_MEMS} expected rank 4, got {}",
path.display(),
shape.len()
)));
}
let layers = shape[0];
let hidden = shape[3];
if layers <= 0 || hidden <= 0 {
return Err(AsrError::Inference(format!(
"decoder {} input {DECODER_INPUT_DECODER_MEMS} has non-static \
layers/hidden dims (L={layers}, H={hidden})",
path.display(),
)));
}
Ok((layers as usize, hidden as usize))
}
fn output_index(outputs: &ort::session::SessionOutputs<'_>, name: &str) -> Option<usize> {
outputs.keys().position(|k| k == name)
}
#[allow(dead_code)]
fn _unused_array1_marker() -> Array1<f32> {
Array1::zeros(0)
}
#[cfg(test)]
mod tests {
use super::*;
fn mini_vocab_text() -> String {
let entries: Vec<(&str, u32)> = vec![
("<unk>", 0),
("<|nospeech|>", 1),
("<pad>", 2),
("<|endoftext|>", 3),
("<|startoftranscript|>", 4),
("<|pnc|>", 5),
("<|nopnc|>", 6),
("<|startofcontext|>", 7),
("<|noitn|>", 8),
("<|nodiarize|>", 9),
("<|notimestamp|>", 10),
("<|emo:undefined|>", 11),
("<|en|>", 12),
("<|de|>", 13),
("<|fr|>", 14),
("<|es|>", 15),
("\u{2581}", 16),
("hello", 17),
("world", 18),
("\u{2581}", 19),
];
let mut s = String::new();
for (piece, id) in entries {
s.push_str(&format!("{piece} {id}\n"));
}
s
}
#[test]
fn prefix_layout_for_spanish_asr_onnx_asr() {
let v = Vocab::from_text(&mini_vocab_text()).unwrap();
let mut opts = DecodeOptions::for_asr("es");
opts.prefix_format = PrefixFormat::OnnxAsr;
let p = build_decoder_prefix(&v, &opts).unwrap();
assert_eq!(p.len(), PREFIX_LEN);
assert_eq!(p[0], 19);
assert_eq!(p[1], 7);
assert_eq!(p[2], 4);
assert_eq!(p[3], 11);
assert_eq!(p[4], 15);
assert_eq!(p[5], 15);
assert_eq!(p[6], 5);
assert_eq!(p[7], 8);
assert_eq!(p[8], 10);
assert_eq!(p[9], 9);
}
#[test]
fn prefix_nemo_canary2_layout_is_nine_tokens() {
let v = Vocab::from_text(&mini_vocab_text()).unwrap();
let mut opts = DecodeOptions::for_asr("es");
opts.prefix_format = PrefixFormat::NemoCanary2;
let p = build_decoder_prefix(&v, &opts).unwrap();
assert_eq!(p.len(), 9);
assert_eq!(p.len(), PrefixFormat::NemoCanary2.token_count());
assert_eq!(p[0], 7);
assert_eq!(p[1], 4); assert_eq!(p[2], 11); assert_eq!(p[3], 15); assert_eq!(p[4], 15); assert_eq!(p[5], 5); assert_eq!(p[6], 8); assert_eq!(p[7], 10); assert_eq!(p[8], 9); }
#[test]
fn prefix_uses_nopnc_when_disabled() {
let v = Vocab::from_text(&mini_vocab_text()).unwrap();
let mut opts = DecodeOptions::for_asr("en");
opts.pnc = false;
let p = build_decoder_prefix(&v, &opts).unwrap();
assert_eq!(p[5], 6); }
#[test]
fn prefix_supports_all_canary_languages() {
let v = Vocab::from_text(&mini_vocab_text()).unwrap();
for (lang, expected_id) in [("en", 12), ("de", 13), ("fr", 14), ("es", 15)] {
let opts = DecodeOptions::for_asr(lang);
let p = build_decoder_prefix(&v, &opts).unwrap();
assert_eq!(p[3], expected_id, "lang={lang}");
assert_eq!(p[4], expected_id, "lang={lang}");
}
}
#[test]
fn prefix_target_language_independent_of_source() {
let v = Vocab::from_text(&mini_vocab_text()).unwrap();
let opts = DecodeOptions {
source_language: "en".into(),
target_language: "es".into(),
pnc: true,
max_sequence_length: 1024,
repetition_penalty: 1.0,
min_token_to_frame_ratio: 0.0,
eos_confidence_margin: 0.0,
beam_size: 1,
length_penalty: 0.0,
prefix_format: PrefixFormat::OnnxAsr,
};
let p = build_decoder_prefix(&v, &opts).unwrap();
assert_eq!(p[4], 12, "source = en");
assert_eq!(p[5], 15, "target = es");
}
#[test]
fn prefix_rejects_unsupported_language() {
let v = Vocab::from_text(&mini_vocab_text()).unwrap();
let opts = DecodeOptions::for_asr("ja"); let err = match build_decoder_prefix(&v, &opts) {
Ok(_) => panic!("expected error"),
Err(e) => e,
};
assert!(err.to_string().contains("ja"), "{}", err);
}
#[test]
fn prefix_errors_when_required_token_missing() {
let mut text = String::new();
text.push_str("<unk> 0\n");
text.push_str("<|nospeech|> 1\n");
text.push_str("<pad> 2\n");
text.push_str("<|endoftext|> 3\n");
text.push_str("<|startoftranscript|> 4\n");
text.push_str("<|pnc|> 5\n");
text.push_str("<|nopnc|> 6\n");
text.push_str("<|startofcontext|> 7\n");
text.push_str("<|notimestamp|> 8\n");
text.push_str("<|nodiarize|> 9\n");
text.push_str("<|emo:undefined|> 10\n");
text.push_str("<|en|> 11\n");
text.push_str("\u{2581} 12\n");
let v = Vocab::from_text(&text).unwrap();
let opts = DecodeOptions::for_asr("en");
let err = match build_decoder_prefix(&v, &opts) {
Ok(_) => panic!("expected error"),
Err(e) => e,
};
assert!(err.to_string().contains("<|noitn|>"), "{}", err);
}
#[test]
fn argmax_last_position_picks_max_at_final_time() {
let logits = Array3::from_shape_vec(
(1, 2, 4),
vec![
10.0, 1.0, 1.0, 1.0, 0.0, 0.5, 9.9, 0.5,
],
)
.unwrap();
assert_eq!(argmax_last_position(&logits), vec![2]);
}
#[test]
fn argmax_last_position_handles_batch() {
let logits = Array3::from_shape_vec(
(2, 1, 3),
vec![
0.0, 5.0, 0.0, 1.0, 1.0, 9.0,
],
)
.unwrap();
assert_eq!(argmax_last_position(&logits), vec![1, 2]);
}
fn make_logits_1xv(row: &[f32]) -> Array3<f32> {
let v = row.len();
Array3::from_shape_vec((1, 1, v), row.to_vec()).unwrap()
}
#[test]
fn repetition_penalty_no_op_at_one() {
let mut logits = make_logits_1xv(&[1.0, 2.0, 3.0]);
let before = logits.clone();
apply_repetition_penalty(&mut logits, &[0, 1, 2], 1.0);
assert_eq!(logits, before);
}
#[test]
fn repetition_penalty_no_op_for_empty_history() {
let mut logits = make_logits_1xv(&[1.0, -2.0, 3.0]);
let before = logits.clone();
apply_repetition_penalty(&mut logits, &[], 1.5);
assert_eq!(logits, before);
}
#[test]
fn repetition_penalty_divides_positive_logits() {
let mut logits = make_logits_1xv(&[2.0, 4.0, 6.0]);
apply_repetition_penalty(&mut logits, &[1], 2.0);
assert_eq!(logits[[0, 0, 0]], 2.0); assert_eq!(logits[[0, 0, 1]], 2.0); assert_eq!(logits[[0, 0, 2]], 6.0); }
#[test]
fn repetition_penalty_multiplies_negative_logits() {
let mut logits = make_logits_1xv(&[-1.0, -2.0, -3.0]);
apply_repetition_penalty(&mut logits, &[2], 2.0);
assert_eq!(logits[[0, 0, 0]], -1.0);
assert_eq!(logits[[0, 0, 1]], -2.0);
assert_eq!(logits[[0, 0, 2]], -6.0); }
#[test]
fn repetition_penalty_dedups_history() {
let mut logits = make_logits_1xv(&[1.0, 4.0, 9.0]);
apply_repetition_penalty(&mut logits, &[1, 1, 1, 1, 1], 2.0);
assert_eq!(logits[[0, 0, 1]], 2.0);
}
#[test]
fn repetition_penalty_skips_out_of_range_tokens() {
let mut logits = make_logits_1xv(&[1.0, 2.0, 3.0]);
apply_repetition_penalty(&mut logits, &[99, -1, 2], 2.0);
assert_eq!(logits[[0, 0, 2]], 1.5);
assert_eq!(logits[[0, 0, 0]], 1.0);
assert_eq!(logits[[0, 0, 1]], 2.0);
}
#[test]
fn repetition_penalty_only_touches_last_time_position() {
let mut logits = Array3::from_shape_vec(
(1, 3, 4),
vec![
1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0,
],
)
.unwrap();
apply_repetition_penalty(&mut logits, &[1], 2.0);
assert_eq!(logits[[0, 0, 1]], 2.0);
assert_eq!(logits[[0, 1, 1]], 6.0);
assert_eq!(logits[[0, 2, 1]], 5.0); }
#[test]
fn repetition_penalty_changes_argmax_when_penalty_active() {
let mut logits = make_logits_1xv(&[5.0, 4.0, 3.0]);
assert_eq!(argmax_last_position(&logits), vec![0]);
apply_repetition_penalty(&mut logits, &[0], 2.0);
assert_eq!(argmax_last_position(&logits), vec![1]);
}
#[test]
fn decode_options_default_repetition_penalty() {
let o = DecodeOptions::for_asr("es");
assert_eq!(o.repetition_penalty, DEFAULT_REPETITION_PENALTY);
assert_eq!(o.repetition_penalty, 2.0);
}
#[test]
fn min_length_no_op_at_zero_ratio() {
let mut logits = make_logits_1xv(&[1.0, 2.0, 3.0]);
let before = logits.clone();
suppress_eos_until_min_length(&mut logits, 1, 0, 100, 0.0);
assert_eq!(logits, before);
}
#[test]
fn min_length_no_op_when_suffix_long_enough() {
let mut logits = make_logits_1xv(&[1.0, 5.0, 3.0]); let before = logits.clone();
suppress_eos_until_min_length(&mut logits, 1, 3, 10, 0.3);
assert_eq!(logits, before);
}
#[test]
fn min_length_suppresses_eos_when_below_min() {
let mut logits = make_logits_1xv(&[1.0, 5.0, 3.0]);
suppress_eos_until_min_length(&mut logits, 1, 2, 10, 0.5);
assert_eq!(logits[[0, 0, 0]], 1.0); assert!(logits[[0, 0, 1]].is_infinite() && logits[[0, 0, 1]] < 0.0);
assert_eq!(logits[[0, 0, 2]], 3.0); }
#[test]
fn min_length_changes_argmax_when_eos_was_winning() {
let mut logits = make_logits_1xv(&[1.0, 5.0, 3.0]);
assert_eq!(argmax_last_position(&logits), vec![1]);
suppress_eos_until_min_length(&mut logits, 1, 0, 10, 0.3);
assert_eq!(argmax_last_position(&logits), vec![2]);
}
#[test]
fn min_length_uses_ceil_for_fractional_min() {
let mut a = make_logits_1xv(&[1.0, 5.0, 3.0]);
suppress_eos_until_min_length(&mut a, 1, 2, 10, 0.25);
assert!(a[[0, 0, 1]].is_infinite());
let mut b = make_logits_1xv(&[1.0, 5.0, 3.0]);
suppress_eos_until_min_length(&mut b, 1, 3, 10, 0.25);
assert_eq!(b[[0, 0, 1]], 5.0);
}
#[test]
fn valid_frame_count_returns_count_of_ones() {
let m =
ndarray::Array2::<i64>::from_shape_vec((1, 8), vec![1, 1, 1, 1, 1, 0, 0, 0]).unwrap();
assert_eq!(valid_frame_count(&m, 0), 5);
}
#[test]
fn valid_frame_count_full_mask_equals_shape() {
let m = ndarray::Array2::<i64>::from_shape_vec((1, 4), vec![1, 1, 1, 1]).unwrap();
assert_eq!(valid_frame_count(&m, 0), 4);
}
#[test]
fn valid_frame_count_handles_out_of_range_batch() {
let m = ndarray::Array2::<i64>::from_shape_vec((1, 4), vec![1, 1, 0, 0]).unwrap();
assert_eq!(valid_frame_count(&m, 7), 0);
}
#[test]
fn valid_frame_count_treats_nonzero_as_valid() {
let m = ndarray::Array2::<i64>::from_shape_vec((1, 4), vec![1, 2, 0, 1]).unwrap();
assert_eq!(valid_frame_count(&m, 0), 3);
}
#[test]
fn valid_frame_count_per_batch_independent() {
let m = ndarray::Array2::<i64>::from_shape_vec(
(2, 4),
vec![
1, 1, 1, 0, 1, 1, 0, 0, ],
)
.unwrap();
assert_eq!(valid_frame_count(&m, 0), 3);
assert_eq!(valid_frame_count(&m, 1), 2);
}
#[test]
fn min_length_skips_when_zero_encoder_frames() {
let mut logits = make_logits_1xv(&[1.0, 5.0, 3.0]);
let before = logits.clone();
suppress_eos_until_min_length(&mut logits, 1, 0, 0, 0.5);
assert_eq!(logits, before);
}
#[test]
fn min_length_skips_out_of_range_eos() {
let mut logits = make_logits_1xv(&[1.0, 2.0, 3.0]);
let before = logits.clone();
suppress_eos_until_min_length(&mut logits, 99, 0, 10, 0.5);
assert_eq!(logits, before);
}
#[test]
fn min_length_only_touches_last_time_position() {
let mut logits = Array3::from_shape_vec(
(1, 2, 3),
vec![
1.0, 2.0, 3.0, 4.0, 5.0, 6.0,
],
)
.unwrap();
suppress_eos_until_min_length(&mut logits, 1, 0, 10, 0.5);
assert_eq!(logits[[0, 0, 1]], 2.0);
assert!(logits[[0, 1, 1]].is_infinite() && logits[[0, 1, 1]] < 0.0);
}
#[test]
fn decode_options_default_min_token_ratio() {
let o = DecodeOptions::for_asr("es");
assert_eq!(o.min_token_to_frame_ratio, DEFAULT_MIN_TOKEN_TO_FRAME_RATIO);
assert_eq!(o.min_token_to_frame_ratio, 0.2);
}
#[test]
fn eos_margin_no_op_at_zero() {
let mut logits = make_logits_1xv(&[5.0, 1.0, 2.0]); let before = logits.clone();
enforce_eos_confidence_margin(&mut logits, 0, 0.0);
assert_eq!(logits, before);
}
#[test]
fn eos_margin_keeps_eos_when_dominant() {
let mut logits = make_logits_1xv(&[5.0, 1.0, 2.0]);
enforce_eos_confidence_margin(&mut logits, 0, 2.0);
assert_eq!(logits[[0, 0, 0]], 5.0);
}
#[test]
fn eos_margin_demotes_eos_when_marginal() {
let mut logits = make_logits_1xv(&[5.0, 4.0, 2.0]);
enforce_eos_confidence_margin(&mut logits, 0, 2.0);
assert!(logits[[0, 0, 0]].is_infinite() && logits[[0, 0, 0]] < 0.0);
assert_eq!(logits[[0, 0, 1]], 4.0);
assert_eq!(logits[[0, 0, 2]], 2.0);
}
#[test]
fn eos_margin_skips_already_suppressed_eos() {
let mut logits = make_logits_1xv(&[f32::NEG_INFINITY, 1.0, 2.0]);
let before = logits.clone();
enforce_eos_confidence_margin(&mut logits, 0, 2.0);
assert_eq!(format!("{:?}", logits), format!("{:?}", before));
}
#[test]
fn eos_margin_skips_out_of_range_id() {
let mut logits = make_logits_1xv(&[1.0, 2.0, 3.0]);
let before = logits.clone();
enforce_eos_confidence_margin(&mut logits, 99, 2.0);
assert_eq!(logits, before);
}
#[test]
fn eos_margin_changes_argmax_when_demoting() {
let mut logits = make_logits_1xv(&[5.0, 4.0, 2.0]);
assert_eq!(argmax_last_position(&logits), vec![0]);
enforce_eos_confidence_margin(&mut logits, 0, 2.0);
assert_eq!(argmax_last_position(&logits), vec![1]);
}
#[test]
fn eos_margin_only_touches_last_time_position() {
let mut logits = Array3::from_shape_vec(
(1, 2, 3),
vec![
5.0, 4.0, 2.0, 5.0, 4.0, 3.0,
],
)
.unwrap();
enforce_eos_confidence_margin(&mut logits, 0, 2.0);
assert_eq!(logits[[0, 0, 0]], 5.0);
assert!(logits[[0, 1, 0]].is_infinite() && logits[[0, 1, 0]] < 0.0);
}
#[test]
fn eos_margin_keeps_eos_when_no_competing_token() {
let mut logits = make_logits_1xv(&[3.0]);
let before = logits.clone();
enforce_eos_confidence_margin(&mut logits, 0, 2.0);
assert_eq!(logits, before);
}
#[test]
fn decode_options_default_eos_confidence_margin() {
let o = DecodeOptions::for_asr("es");
assert_eq!(o.eos_confidence_margin, DEFAULT_EOS_CONFIDENCE_MARGIN);
assert_eq!(o.eos_confidence_margin, 2.0);
}
#[test]
fn log_softmax_row_sums_to_one_in_probability_space() {
let logits = vec![1.0, 2.0, 3.0, 4.0];
let lp = log_softmax_row(&logits);
let sum_p: f32 = lp.iter().map(|&x| x.exp()).sum();
assert!((sum_p - 1.0).abs() < 1e-5, "sum p = {sum_p}");
}
#[test]
fn log_softmax_row_handles_neg_inf_inputs() {
let logits = vec![1.0_f32, f32::NEG_INFINITY, 2.0, 3.0];
let lp = log_softmax_row(&logits);
for v in &lp {
assert!(v.is_finite() || *v == f32::NEG_INFINITY, "v={v}");
}
assert_eq!(lp[1], f32::NEG_INFINITY);
}
#[test]
fn topk_indices_picks_largest_in_descending_order() {
let xs = vec![1.0, 5.0, 2.0, 4.0, 3.0];
let top = topk_indices(&xs, 3);
assert_eq!(top.len(), 3);
assert_eq!(top[0].0, 1); assert_eq!(top[1].0, 3); assert_eq!(top[2].0, 4); }
#[test]
fn topk_indices_breaks_ties_by_lower_index() {
let xs = vec![3.0, 3.0, 3.0, 1.0];
let top = topk_indices(&xs, 2);
assert_eq!(top[0].0, 0);
assert_eq!(top[1].0, 1);
}
#[test]
fn topk_indices_caps_at_input_length() {
let xs = vec![1.0, 2.0];
let top = topk_indices(&xs, 5);
assert_eq!(top.len(), 2);
}
#[test]
fn length_normalised_score_alpha_zero_returns_raw_logprob() {
assert_eq!(length_normalised_score(-3.0, 5, 0.0), -3.0);
}
#[test]
fn length_normalised_score_alpha_one_divides_by_length() {
let s = length_normalised_score(-6.0, 3, 1.0);
assert!((s - (-2.0)).abs() < 1e-6, "s={s}");
}
#[test]
fn length_normalised_score_zero_length_returns_raw() {
assert_eq!(length_normalised_score(0.0, 0, 0.6), 0.0);
}
#[test]
fn length_normalised_score_alpha_below_one_penalises_longer_seqs() {
let short = length_normalised_score(-5.0, 5, 0.6);
let long = length_normalised_score(-10.0, 10, 0.6);
assert!(short > long, "short={short} long={long}");
}
#[test]
fn length_normalised_score_alpha_above_one_promotes_longer_seqs() {
let short = length_normalised_score(-5.0, 5, 1.5);
let long = length_normalised_score(-10.0, 10, 1.5);
assert!(long > short, "short={short} long={long}");
}
#[test]
fn decode_options_default_beam_size_is_one() {
let o = DecodeOptions::for_asr("es");
assert_eq!(o.beam_size, DEFAULT_BEAM_SIZE);
assert_eq!(o.beam_size, 1);
}
#[test]
fn stack_along_batch_replicates_correctly() {
let h: Array4<f32> =
Array4::from_shape_vec((2, 1, 3, 2), (0..12).map(|x| x as f32).collect()).unwrap();
let stacked = stack_along_batch(&h, 4).unwrap();
assert_eq!(stacked.dim(), (2, 4, 3, 2));
for k in 0..4 {
for l in 0..2 {
for t in 0..3 {
for c in 0..2 {
assert_eq!(stacked[[l, k, t, c]], h[[l, 0, t, c]]);
}
}
}
}
}
#[test]
fn stack_along_batch_rejects_non_unit_batch() {
let h: Array4<f32> = Array4::zeros((2, 3, 1, 1));
assert!(stack_along_batch(&h, 4).is_err());
}
#[test]
fn repeat_along_axis0_f32_replicates() {
let a: Array3<f32> =
Array3::from_shape_vec((1, 2, 3), (0..6).map(|x| x as f32).collect()).unwrap();
let r = repeat_along_axis0(&a, 3);
assert_eq!(r.dim(), (3, 2, 3));
for k in 0..3 {
for i in 0..2 {
for j in 0..3 {
assert_eq!(r[[k, i, j]], a[[0, i, j]]);
}
}
}
}
#[test]
fn repeat_along_axis0_i64_replicates() {
let a: Array2<i64> = Array2::from_shape_vec((1, 4), vec![10, 20, 30, 40]).unwrap();
let r = repeat_along_axis0_i64(&a, 2);
assert_eq!(r.dim(), (2, 4));
assert_eq!(r[[0, 2]], 30);
assert_eq!(r[[1, 2]], 30);
}
#[test]
fn apply_repetition_penalty_one_batch_only_touches_target_batch() {
let mut logits: Array3<f32> =
Array3::from_shape_vec((2, 1, 3), vec![1.0, 2.0, 3.0, 1.0, 2.0, 3.0]).unwrap();
apply_repetition_penalty_one_batch(&mut logits, 1, &[1], 2.0);
assert_eq!(logits[[0, 0, 1]], 2.0); assert_eq!(logits[[1, 0, 1]], 1.0); }
#[test]
fn strip_prefix_and_specials_drops_specials() {
let v = Vocab::from_text(&mini_vocab_text()).unwrap();
let mut all = vec![0u32; PREFIX_LEN];
all.extend_from_slice(&[17, 3, 18]);
let kept = strip_prefix_and_specials(&all, PREFIX_LEN, &v);
assert_eq!(kept, vec![17, 18]);
}
#[test]
fn strip_prefix_and_specials_handles_no_emissions() {
let v = Vocab::from_text(&mini_vocab_text()).unwrap();
let only_prefix = vec![0u32; PREFIX_LEN];
let kept = strip_prefix_and_specials(&only_prefix, PREFIX_LEN, &v);
assert!(kept.is_empty());
}
#[test]
fn strip_prefix_and_specials_handles_oversized_prefix() {
let v = Vocab::from_text(&mini_vocab_text()).unwrap();
let kept = strip_prefix_and_specials(&[1, 2, 3], 10, &v);
assert!(kept.is_empty());
}
#[test]
fn load_nonexistent_decoder_returns_error() {
match CanaryDecoder::load("/nonexistent/path/to/decoder.onnx") {
Ok(_) => panic!("expected error"),
Err(e) => assert!(e.to_string().contains("load Canary decoder"), "{}", e.to_string()),
}
}
#[test]
fn io_name_constants_match_onnx_asr_conventions() {
assert_eq!(DECODER_INPUT_IDS, "input_ids");
assert_eq!(DECODER_INPUT_ENCODER_EMBEDDINGS, "encoder_embeddings");
assert_eq!(DECODER_INPUT_ENCODER_MASK, "encoder_mask");
assert_eq!(DECODER_INPUT_DECODER_MEMS, "decoder_mems");
assert_eq!(DECODER_OUTPUT_LOGITS, "logits");
assert_eq!(DECODER_OUTPUT_HIDDEN_STATES, "decoder_hidden_states");
}
#[test]
fn prefix_len_constant_matches_layout() {
assert_eq!(PREFIX_LEN, 10);
}
}