use super::*;
#[test]
fn test_extract_encoder_frame_first() {
let encoded = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let mut frame = vec![0.0; 2];
extract_encoder_frame(&encoded, 3, 0, &mut frame);
assert_eq!(frame, vec![1.0, 4.0]);
}
#[test]
fn test_extract_encoder_frame_last() {
let encoded = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let mut frame = vec![0.0; 2];
extract_encoder_frame(&encoded, 3, 2, &mut frame);
assert_eq!(frame, vec![3.0, 6.0]);
}
#[test]
fn test_extract_encoder_frame_middle() {
let encoded = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let mut frame = vec![0.0; 2];
extract_encoder_frame(&encoded, 3, 1, &mut frame);
assert_eq!(frame, vec![2.0, 5.0]);
}
#[test]
fn test_argmax_clear_winner() {
let logits = vec![0.1, 0.5, 0.9, 0.2];
assert_eq!(argmax(&logits, 999), 2);
}
#[test]
fn test_argmax_tie_returns_last() {
let logits = vec![1.0, 1.0, 0.5];
assert_eq!(argmax(&logits, 999), 1);
}
#[test]
fn test_argmax_single_element() {
let logits = vec![42.0];
assert_eq!(argmax(&logits, 999), 0);
}
#[test]
fn test_argmax_negative_values() {
let logits = vec![-3.0, -1.0, -2.0];
assert_eq!(argmax(&logits, 999), 1);
}
#[test]
fn test_argmax_empty_returns_blank() {
let logits: Vec<f32> = vec![];
assert_eq!(argmax(&logits, 1024), 1024);
}
#[test]
fn test_argmax_blank_id_selected() {
let logits = vec![0.1, 0.2, 0.9]; assert_eq!(argmax(&logits, 2), 2); }
struct FakeBackend {
script: std::collections::VecDeque<usize>,
vocab: usize,
blank_id: usize,
decoder_calls: u32,
joiner_calls: u32,
}
impl FakeBackend {
fn new(script: Vec<usize>, vocab: usize, blank_id: usize) -> Self {
Self {
script: script.into(),
vocab,
blank_id,
decoder_calls: 0,
joiner_calls: 0,
}
}
}
impl DecodeBackend for FakeBackend {
fn decode_step(
&mut self,
_state: &DecoderState,
out: &mut DecoderOutput,
_bufs: &mut DecodeBuffers,
) -> Result<()> {
self.decoder_calls += 1;
DecoderOutput::fill(&mut out.dec_data, &[0.0; PRED_HIDDEN]);
DecoderOutput::fill(&mut out.new_h, &[0.0; PRED_HIDDEN]);
DecoderOutput::fill(&mut out.new_c, &[0.0; PRED_HIDDEN]);
Ok(())
}
fn joiner_step(
&mut self,
_enc_frame: &[f32],
_dec_data: &[f32],
logits_buf: &mut Vec<f32>,
_bufs: &mut DecodeBuffers,
) -> Result<()> {
self.joiner_calls += 1;
let tok = self.script.pop_front().unwrap_or(self.blank_id);
logits_buf.clear();
logits_buf.resize(self.vocab, 0.0);
logits_buf[tok] = 10.0; Ok(())
}
}
struct FlatLogitsBackend {
logits: Vec<f32>,
}
impl DecodeBackend for FlatLogitsBackend {
fn decode_step(
&mut self,
_state: &DecoderState,
out: &mut DecoderOutput,
_bufs: &mut DecodeBuffers,
) -> Result<()> {
DecoderOutput::fill(&mut out.dec_data, &[0.0; PRED_HIDDEN]);
DecoderOutput::fill(&mut out.new_h, &[0.0; PRED_HIDDEN]);
DecoderOutput::fill(&mut out.new_c, &[0.0; PRED_HIDDEN]);
Ok(())
}
fn joiner_step(
&mut self,
_enc_frame: &[f32],
_dec_data: &[f32],
logits_buf: &mut Vec<f32>,
_bufs: &mut DecodeBuffers,
) -> Result<()> {
logits_buf.clear();
logits_buf.extend_from_slice(&self.logits);
Ok(())
}
}
fn blank_leading_logits() -> Vec<f32> {
let mut logits = vec![0.0; 5];
logits[4] = 3.0; logits[1] = 1.0; logits[2] = 1.5; logits
}
#[test]
fn test_biasing_emits_at_most_one_token_per_frame() {
let biaser = Biaser::from_sequences(vec![vec![1, 2]], 5.0).expect("biaser compiles");
let mut backend = FlatLogitsBackend {
logits: blank_leading_logits(),
};
let mut state = DecoderState::new(4);
let enc = fake_enc_tensor(3);
let result =
greedy_decode_impl(&mut backend, &enc.view(), 3, 4, &mut state, Some(&biaser)).unwrap();
let frames: Vec<usize> = result.tokens.iter().map(|t| t.frame_index).collect();
assert_eq!(
frames,
vec![0, 1, 2],
"each frame may contribute one biased token, not a burst"
);
let ids: Vec<usize> = result.tokens.iter().map(|t| t.token_id).collect();
assert_eq!(&ids[..2], &[1, 2], "hotword advances across frames");
}
#[test]
fn test_biasing_reports_the_models_own_confidence() {
let logits = blank_leading_logits();
let biaser = Biaser::from_sequences(vec![vec![1, 2]], 5.0).expect("biaser compiles");
let mut backend = FlatLogitsBackend {
logits: logits.clone(),
};
let mut state = DecoderState::new(4);
let enc = fake_enc_tensor(1);
let result =
greedy_decode_impl(&mut backend, &enc.view(), 1, 4, &mut state, Some(&biaser)).unwrap();
assert_eq!(result.tokens.len(), 1);
let expected = token_confidence(&logits, 1);
assert!(
(result.tokens[0].confidence - expected).abs() < 1e-6,
"confidence {} should be the un-boosted {expected}",
result.tokens[0].confidence
);
}
#[test]
fn test_biasing_leaves_a_confident_model_pick_alone() {
let mut logits = vec![0.0; 5];
logits[3] = 10.0; logits[1] = 1.0;
let biaser = Biaser::from_sequences(vec![vec![1, 2]], 5.0).expect("biaser compiles");
let mut backend = FlatLogitsBackend { logits };
let mut state = DecoderState::new(4);
let enc = fake_enc_tensor(1);
let result =
greedy_decode_impl(&mut backend, &enc.view(), 1, 4, &mut state, Some(&biaser)).unwrap();
assert_eq!(
result.tokens.len(),
MAX_TOKENS_PER_STEP,
"an unbiased pick is not rationed by the biasing budget"
);
assert!(result.tokens.iter().all(|t| t.token_id == 3));
}
fn fake_enc(frames: usize) -> Vec<f32> {
vec![0.0_f32; ENC_DIM * frames]
}
fn fake_enc_tensor(frames: usize) -> Tensor {
Tensor::new(
Shape::new(vec![1, ENC_DIM, frames]),
TensorData::F32(fake_enc(frames)),
)
.unwrap()
}
#[test]
fn test_greedy_decode_happy_path() {
let mut backend = FakeBackend::new(vec![1, 4, 2, 4], 5, 4);
let mut state = DecoderState::new(4);
let enc = fake_enc_tensor(2);
let result = greedy_decode_impl(&mut backend, &enc.view(), 2, 4, &mut state, None).unwrap();
assert_eq!(result.tokens.len(), 2);
assert_eq!(result.tokens[0].token_id, 1);
assert_eq!(result.tokens[0].frame_index, 0);
assert_eq!(result.tokens[1].token_id, 2);
assert_eq!(result.tokens[1].frame_index, 1);
assert_eq!(state.prev_token, 2);
assert_eq!(state.h.len(), PRED_HIDDEN);
assert!(!result.endpoint_detected);
}
#[test]
fn test_greedy_decode_blank_run_skips_decoder() {
let mut backend = FakeBackend::new(vec![1, 4, 4, 4, 4], 5, 4);
let mut state = DecoderState::new(4);
let enc = fake_enc_tensor(4);
let result = greedy_decode_impl(&mut backend, &enc.view(), 4, 4, &mut state, None).unwrap();
assert_eq!(result.tokens.len(), 1);
assert_eq!(
backend.decoder_calls, 2,
"decoder must not run during the blank run"
);
assert!(backend.joiner_calls >= 5);
}
#[test]
fn test_greedy_decode_endpoint_after_threshold_blanks() {
let mut script = vec![1usize];
script.extend(std::iter::repeat_n(4usize, ENDPOINT_BLANK_THRESHOLD + 1));
let frames = ENDPOINT_BLANK_THRESHOLD + 2;
let mut backend = FakeBackend::new(script, 5, 4);
let mut state = DecoderState::new(4);
let enc = fake_enc_tensor(frames);
let result =
greedy_decode_impl(&mut backend, &enc.view(), frames, 4, &mut state, None).unwrap();
assert!(
result.endpoint_detected,
"{ENDPOINT_BLANK_THRESHOLD}+ blanks after a token must endpoint"
);
}
#[test]
fn test_greedy_decode_no_endpoint_before_first_token() {
let frames = ENDPOINT_BLANK_THRESHOLD + 5;
let mut backend = FakeBackend::new(vec![4usize; frames], 5, 4);
let mut state = DecoderState::new(4);
let enc = fake_enc_tensor(frames);
let result =
greedy_decode_impl(&mut backend, &enc.view(), frames, 4, &mut state, None).unwrap();
assert!(result.tokens.is_empty());
assert!(
!result.endpoint_detected,
"blanks before any token must not endpoint"
);
}
#[test]
fn test_greedy_decode_token_cap_does_not_inflate_blanks() {
let mut backend = FakeBackend::new(vec![1usize; MAX_TOKENS_PER_STEP + 1], 5, 4);
let mut state = DecoderState::new(4);
let enc = fake_enc_tensor(1);
let result = greedy_decode_impl(&mut backend, &enc.view(), 1, 4, &mut state, None).unwrap();
assert_eq!(result.tokens.len(), MAX_TOKENS_PER_STEP);
assert_eq!(
state.consecutive_blanks, 0,
"token cap must not inflate the blank counter"
);
assert!(!result.endpoint_detected);
}
#[test]
fn test_argmax_with_confidence_clear_winner() {
let (tok, conf) = argmax_with_confidence(&[0.1, 5.0, 0.2], 99);
assert_eq!(tok, 1);
assert!(
conf > 0.5 && conf <= 1.0,
"confidence should be a softmax prob in (0.5, 1], got {conf}"
);
}
#[test]
fn test_argmax_with_confidence_empty_returns_blank_zero() {
let (tok, conf) = argmax_with_confidence(&[], 1024);
assert_eq!(tok, 1024);
assert_eq!(conf, 0.0);
}
struct LogitBackend {
script: std::collections::VecDeque<Vec<f32>>,
vocab: usize,
blank_id: usize,
}
impl LogitBackend {
fn new(script: Vec<Vec<f32>>, vocab: usize, blank_id: usize) -> Self {
Self {
script: script.into(),
vocab,
blank_id,
}
}
}
impl DecodeBackend for LogitBackend {
fn decode_step(
&mut self,
_state: &DecoderState,
out: &mut DecoderOutput,
_bufs: &mut DecodeBuffers,
) -> Result<()> {
DecoderOutput::fill(&mut out.dec_data, &[0.0; PRED_HIDDEN]);
DecoderOutput::fill(&mut out.new_h, &[0.0; PRED_HIDDEN]);
DecoderOutput::fill(&mut out.new_c, &[0.0; PRED_HIDDEN]);
Ok(())
}
fn joiner_step(
&mut self,
_enc_frame: &[f32],
_dec_data: &[f32],
logits_buf: &mut Vec<f32>,
_bufs: &mut DecodeBuffers,
) -> Result<()> {
logits_buf.clear();
match self.script.pop_front() {
Some(v) => logits_buf.extend_from_slice(&v),
None => {
logits_buf.resize(self.vocab, 0.0);
logits_buf[self.blank_id] = 10.0;
}
}
Ok(())
}
}
fn ab_script() -> Vec<Vec<f32>> {
vec![
vec![0.0, 2.0, 1.0, 0.0],
vec![0.0, 0.0, 0.0, 100.0],
]
}
#[test]
fn test_bias_steers_argmax_to_boosted_token() {
let mut backend = LogitBackend::new(ab_script(), 4, 3);
let mut state = DecoderState::new(3);
let enc = fake_enc_tensor(2);
let unbiased = greedy_decode_impl(&mut backend, &enc.view(), 2, 3, &mut state, None).unwrap();
assert_eq!(unbiased.tokens.len(), 1);
assert_eq!(unbiased.tokens[0].token_id, 1, "no bias → model picks A");
let biaser = Biaser::from_sequences(vec![vec![2]], 5.0).unwrap();
let mut backend = LogitBackend::new(ab_script(), 4, 3);
let mut state = DecoderState::new(3);
let enc = fake_enc_tensor(2);
let biased =
greedy_decode_impl(&mut backend, &enc.view(), 2, 3, &mut state, Some(&biaser)).unwrap();
assert_eq!(biased.tokens.len(), 1);
assert_eq!(
biased.tokens[0].token_id, 2,
"boost must steer the argmax from A to the hotword token B"
);
}
#[test]
fn test_bias_prefix_advances_then_boosts_continuation() {
let script = vec![
vec![0.0, 0.0, 0.0, 0.0, 0.0, 3.0],
vec![0.0, 0.0, 0.0, 100.0, 0.0, 0.0],
vec![0.0, 2.0, 1.0, 0.0, 0.0, 0.0],
vec![0.0, 0.0, 0.0, 100.0, 0.0, 0.0],
];
let biaser = Biaser::from_sequences(vec![vec![5, 2]], 5.0).unwrap();
let mut backend = LogitBackend::new(script, 6, 3);
let mut state = DecoderState::new(3);
let enc = fake_enc_tensor(2);
let result =
greedy_decode_impl(&mut backend, &enc.view(), 2, 3, &mut state, Some(&biaser)).unwrap();
assert_eq!(
result.tokens.iter().map(|t| t.token_id).collect::<Vec<_>>(),
vec![5, 2],
"prefix [5] must advance so the boost on the continuation 2 steers frame 1"
);
}
#[test]
fn test_bias_none_is_byte_for_byte_unchanged() {
let base_script = || {
vec![
vec![0.0, 2.0, 1.0, 0.0],
vec![0.0, 0.0, 0.0, 100.0],
vec![0.0, 1.5, 2.5, 0.0],
vec![0.0, 0.0, 0.0, 100.0],
]
};
let mut b_none = LogitBackend::new(base_script(), 4, 3);
let mut s_none = DecoderState::new(3);
let enc_none = fake_enc_tensor(2);
let none = greedy_decode_impl(&mut b_none, &enc_none.view(), 2, 3, &mut s_none, None).unwrap();
let biaser = Biaser::from_sequences(vec![vec![0]], 0.5).unwrap();
let mut b_some = LogitBackend::new(base_script(), 4, 3);
let mut s_some = DecoderState::new(3);
let enc_some = fake_enc_tensor(2);
let some = greedy_decode_impl(
&mut b_some,
&enc_some.view(),
2,
3,
&mut s_some,
Some(&biaser),
)
.unwrap();
assert_eq!(
none.tokens.iter().map(|t| t.token_id).collect::<Vec<_>>(),
some.tokens.iter().map(|t| t.token_id).collect::<Vec<_>>(),
"a non-winning hotword must not change the decoded tokens"
);
assert_eq!(none.tokens.len(), 2);
}