use anyhow::{Context, Result};
use crate::runtime::{
session::RuntimeSession,
tensor::{Shape, Tensor, TensorData, TensorView},
};
use super::bias::Biaser;
use super::{DecoderState, PRED_HIDDEN};
const MAX_TOKENS_PER_STEP: usize = 10;
const MAX_BIAS_OVERRIDES_PER_STEP: usize = 1;
const ENC_DIM: usize = 768;
pub(crate) const ENDPOINT_BLANK_THRESHOLD: usize = 15;
#[derive(Debug, Clone)]
pub(crate) struct TokenInfo {
pub token_id: usize,
pub frame_index: usize,
pub confidence: f32,
}
#[derive(Debug)]
pub(crate) struct DecodeResult {
pub tokens: Vec<TokenInfo>,
pub endpoint_detected: bool,
}
pub(crate) fn extract_encoder_frame(
encoded: &[f32],
encoded_len: usize,
t: usize,
enc_frame: &mut [f32],
) {
for ch in 0..enc_frame.len() {
enc_frame[ch] = encoded[ch * encoded_len + t];
}
}
pub(crate) fn argmax(logits: &[f32], blank_id: usize) -> usize {
logits
.iter()
.enumerate()
.max_by(|(_i, a), (_j, b)| a.total_cmp(b))
.map(|(idx, _)| idx)
.unwrap_or(blank_id)
}
pub(crate) fn token_confidence(logits: &[f32], token: usize) -> f32 {
let Some(&logit) = logits.get(token) else {
return 0.0;
};
let max_logit = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let sum_exp: f32 = logits.iter().map(|&l| (l - max_logit).exp()).sum();
(logit - max_logit).exp() / sum_exp
}
pub(crate) fn argmax_with_confidence(logits: &[f32], blank_id: usize) -> (usize, f32) {
if logits.is_empty() {
return (blank_id, 0.0);
}
let token = argmax(logits, blank_id);
(token, token_confidence(logits, token))
}
#[derive(Default)]
pub(crate) struct DecoderOutput {
dec_data: Vec<f32>,
new_h: Vec<f32>,
new_c: Vec<f32>,
}
impl DecoderOutput {
fn fill(dst: &mut Vec<f32>, src: &[f32]) {
if dst.len() != src.len() {
dst.resize(src.len(), 0.0);
}
dst.copy_from_slice(src);
}
}
#[derive(Debug)]
pub(crate) struct DecodeBuffers {
decoder_inputs: Vec<Tensor>,
joiner_inputs: Vec<Tensor>,
}
impl DecodeBuffers {
fn new() -> Self {
Self {
decoder_inputs: vec![
Tensor::new_checked(Shape::new(vec![1, 1]), TensorData::I64(vec![0])),
Tensor::new_checked(
Shape::new(vec![1, 1, PRED_HIDDEN]),
TensorData::F32(vec![0.0; PRED_HIDDEN]),
),
Tensor::new_checked(
Shape::new(vec![1, 1, PRED_HIDDEN]),
TensorData::F32(vec![0.0; PRED_HIDDEN]),
),
],
joiner_inputs: vec![
Tensor::new_checked(
Shape::new(vec![1, ENC_DIM, 1]),
TensorData::F32(vec![0.0; ENC_DIM]),
),
Tensor::new_checked(
Shape::new(vec![1, PRED_HIDDEN, 1]),
TensorData::F32(vec![0.0; PRED_HIDDEN]),
),
],
}
}
}
fn run_decoder(
decoder: &dyn RuntimeSession,
state: &DecoderState,
out: &mut DecoderOutput,
bufs: &mut DecodeBuffers,
) -> Result<()> {
bufs.decoder_inputs[0]
.as_i64_mut()
.context("decoder prev_token tensor is not i64")?[0] = state.prev_token;
bufs.decoder_inputs[1]
.as_f32_mut()
.context("decoder h tensor is not f32")?
.copy_from_slice(&state.h);
bufs.decoder_inputs[2]
.as_f32_mut()
.context("decoder c tensor is not f32")?
.copy_from_slice(&state.c);
let decoder_outputs = decoder
.run(&bufs.decoder_inputs)
.context("Decoder inference failed")?;
let dec_data = decoder_outputs[0]
.view()
.data()
.as_f32()
.context("Failed to extract decoder output")?;
let new_h_data = decoder_outputs[1]
.view()
.data()
.as_f32()
.context("Failed to extract decoder h state")?;
let new_c_data = decoder_outputs[2]
.view()
.data()
.as_f32()
.context("Failed to extract decoder c state")?;
DecoderOutput::fill(&mut out.dec_data, dec_data);
DecoderOutput::fill(&mut out.new_h, new_h_data);
DecoderOutput::fill(&mut out.new_c, new_c_data);
Ok(())
}
fn run_joiner_single(
joiner: &dyn RuntimeSession,
enc_frame: &[f32],
dec_data: &[f32],
logits_buf: &mut Vec<f32>,
bufs: &mut DecodeBuffers,
) -> Result<()> {
bufs.joiner_inputs[0]
.as_f32_mut()
.context("joiner enc_frame tensor is not f32")?
.copy_from_slice(enc_frame);
bufs.joiner_inputs[1]
.as_f32_mut()
.context("joiner dec_data tensor is not f32")?
.copy_from_slice(dec_data);
let joiner_outputs = joiner
.run(&bufs.joiner_inputs)
.context("Joiner inference failed")?;
let logits = joiner_outputs[0]
.view()
.data()
.as_f32()
.context("Failed to extract joiner output")?;
DecoderOutput::fill(logits_buf, logits);
Ok(())
}
pub(crate) trait DecodeBackend {
fn decode_step(
&mut self,
state: &DecoderState,
out: &mut DecoderOutput,
bufs: &mut DecodeBuffers,
) -> Result<()>;
fn joiner_step(
&mut self,
enc_frame: &[f32],
dec_data: &[f32],
logits_buf: &mut Vec<f32>,
bufs: &mut DecodeBuffers,
) -> Result<()>;
}
struct OrtBackend<'a> {
decoder: &'a dyn RuntimeSession,
joiner: &'a dyn RuntimeSession,
}
impl DecodeBackend for OrtBackend<'_> {
fn decode_step(
&mut self,
state: &DecoderState,
out: &mut DecoderOutput,
bufs: &mut DecodeBuffers,
) -> Result<()> {
run_decoder(self.decoder, state, out, bufs)
}
fn joiner_step(
&mut self,
enc_frame: &[f32],
dec_data: &[f32],
logits_buf: &mut Vec<f32>,
bufs: &mut DecodeBuffers,
) -> Result<()> {
run_joiner_single(self.joiner, enc_frame, dec_data, logits_buf, bufs)
}
}
pub fn greedy_decode(
decoder: &dyn RuntimeSession,
joiner: &dyn RuntimeSession,
encoded: &TensorView<'_>, encoded_len: usize,
blank_id: usize,
state: &mut DecoderState,
biaser: Option<&Biaser>,
) -> Result<DecodeResult> {
let mut backend = OrtBackend { decoder, joiner };
greedy_decode_impl(&mut backend, encoded, encoded_len, blank_id, state, biaser)
}
fn select_token(
logits: &[f32],
blank_id: usize,
biaser: Option<&Biaser>,
bias_state: Option<&super::bias::BiasState>,
bias_overrides: usize,
biased_buf: &mut Vec<f32>,
) -> (usize, f32, bool) {
match (biaser, bias_state) {
(Some(b), Some(bs)) if bias_overrides < MAX_BIAS_OVERRIDES_PER_STEP => {
biased_buf.clear();
biased_buf.extend_from_slice(logits);
b.boost_logits(bs, biased_buf);
let boosted = argmax(biased_buf, blank_id);
let spent = boosted != argmax(logits, blank_id);
(boosted, token_confidence(logits, boosted), spent)
}
_ => {
let (token, confidence) = argmax_with_confidence(logits, blank_id);
(token, confidence, false)
}
}
}
fn commit_non_blank(
state: &mut DecoderState,
decoder_out: &DecoderOutput,
token: usize,
biaser: Option<&Biaser>,
bias_state: Option<&mut super::bias::BiasState>,
) -> Result<()> {
state.consecutive_blanks = 0;
state.prev_token = token as i64;
if decoder_out.new_h.len() != PRED_HIDDEN || decoder_out.new_c.len() != PRED_HIDDEN {
anyhow::bail!(
"Unexpected decoder state shape: h={}, c={}, expected {}",
decoder_out.new_h.len(),
decoder_out.new_c.len(),
PRED_HIDDEN
);
}
state.h.copy_from_slice(&decoder_out.new_h);
state.c.copy_from_slice(&decoder_out.new_c);
if let (Some(b), Some(bs)) = (biaser, bias_state) {
b.advance(bs, token);
}
Ok(())
}
fn greedy_decode_impl<B: DecodeBackend>(
backend: &mut B,
encoded: &TensorView<'_>, encoded_len: usize,
blank_id: usize,
state: &mut DecoderState,
biaser: Option<&Biaser>,
) -> Result<DecodeResult> {
let encoded = encoded
.data()
.as_f32()
.context("encoder output must be f32")?;
let mut tokens = Vec::new();
let mut endpoint_detected = false;
let mut bufs = DecodeBuffers::new();
let mut enc_frame = vec![0.0_f32; ENC_DIM];
let mut logits_buf = Vec::new();
let mut decoder_calls: u32 = 0;
let mut joiner_calls: u32 = 0;
let mut skipped_decoder_calls: u32 = 0;
let mut decoder_out = DecoderOutput::default();
let mut cache_valid = false;
let mut in_blank_run = false;
let mut bias_state = biaser.map(|b| b.new_state());
let mut biased_buf: Vec<f32> = Vec::new();
anyhow::ensure!(
encoded.len() >= ENC_DIM * encoded_len,
"Encoder output size mismatch: got {}, expected >= {}",
encoded.len(),
ENC_DIM * encoded_len
);
for t in 0..encoded_len {
let mut tokens_this_step = 0;
let mut bias_overrides = 0usize;
extract_encoder_frame(encoded, encoded_len, t, &mut enc_frame);
loop {
if in_blank_run {
skipped_decoder_calls += 1;
if !cache_valid {
anyhow::bail!("blank run invariant violated: decoder output cache is stale");
}
} else {
decoder_calls += 1;
backend.decode_step(state, &mut decoder_out, &mut bufs)?;
cache_valid = true;
}
joiner_calls += 1;
backend.joiner_step(
&enc_frame,
&decoder_out.dec_data,
&mut logits_buf,
&mut bufs,
)?;
let (token, confidence, spent) = select_token(
&logits_buf,
blank_id,
biaser,
bias_state.as_ref(),
bias_overrides,
&mut biased_buf,
);
if spent {
bias_overrides += 1;
}
if token == blank_id {
in_blank_run = true;
state.consecutive_blanks += 1;
if state.consecutive_blanks >= ENDPOINT_BLANK_THRESHOLD && !tokens.is_empty() {
endpoint_detected = true;
}
break;
}
if tokens_this_step >= MAX_TOKENS_PER_STEP {
in_blank_run = false;
cache_valid = false;
state.consecutive_blanks = 0;
break;
}
in_blank_run = false;
commit_non_blank(state, &decoder_out, token, biaser, bias_state.as_mut())?;
tokens.push(TokenInfo {
token_id: token,
frame_index: t,
confidence,
});
tokens_this_step += 1;
}
}
tracing::debug!(
decoder_calls,
joiner_calls,
skipped_decoder_calls,
encoded_len,
"decode_loop_stats"
);
Ok(DecodeResult {
tokens,
endpoint_detected,
})
}
#[cfg(test)]
mod tests;