1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
//! `impl Engine` methods — split out of the former god-file.
use super::*;
impl Engine {
/// Run encoder + decode. `low_latency` selects the streaming encoder path
/// ([`RuntimeSession::run_low_latency`]) so ANE can pad underfilled short
/// windows; file mode keeps the calibrated 0.5 fill floor. `biaser` is the
/// effective per-call hotword biaser (boot, request override, or off).
#[allow(clippy::too_many_arguments)] // encoder/decode call site; bundle later if it grows again
pub(crate) fn run_inference(
&self,
triplet: &mut SessionTriplet,
features: &[f32],
num_frames: usize,
decoder_state: &mut DecoderState,
frame_offset: usize,
low_latency: bool,
biaser: Option<&bias::Biaser>,
) -> anyhow::Result<(Vec<WordInfo>, bool)> {
// Reuse the encoder input tensors: resize the signal tensor to the
// current frame count and overwrite both buffers in place.
triplet.encoder_inputs[0].resize_to(Shape::new(vec![1, N_MELS, num_frames]));
triplet.encoder_inputs[0]
.as_f32_mut()
.context("encoder signal tensor is not f32")?
.copy_from_slice(features);
triplet.encoder_inputs[1]
.as_i64_mut()
.context("encoder length tensor is not i64")?[0] = num_frames as i64;
let enc_start = std::time::Instant::now();
let encoder_outputs = if low_latency {
triplet
.encoder
.run_low_latency(&triplet.encoder_inputs)
.context("Encoder inference failed")?
} else {
triplet
.encoder
.run(&triplet.encoder_inputs)
.context("Encoder inference failed")?
};
tracing::info!(
elapsed_ms = enc_start.elapsed().as_millis() as u64,
"encoder_inference"
);
let enc_len = match encoder_outputs[1].view().data() {
TensorDataView::I32(v) => usize::try_from(v[0]).context("Negative encoder length")?,
TensorDataView::I64(v) => usize::try_from(v[0]).context("Negative encoder length")?,
_ => anyhow::bail!("Unexpected encoder length tensor type"),
};
tracing::debug!("Encoder output: {} frames", enc_len);
// CTC head: the single encoder emits per-frame class log-probs
// (`[1, T', 71]`, row-major). Decode them directly — there is no
// prediction network / joiner, so we return before the RNN-T block
// borrows `encoder_outputs` for the decode loop.
//
// A glossary switches the decode to a prefix beam, which is the only
// form that can act on one: a per-frame argmax has no continuation
// state to steer. Without hotwords the greedy path runs untouched, so
// output for everyone else is byte-for-byte what it was.
if self.variant.is_ctc() {
let log_probs = encoder_outputs[0]
.view()
.data()
.as_f32()
.context("CTC log_probs tensor is not f32")?;
let tokens = match biaser {
Some(b) => ctc::ctc_prefix_beam_decode(
log_probs,
enc_len,
self.tokenizer.vocab_size(),
self.tokenizer.blank_id(),
b,
),
None => ctc::ctc_greedy_decode(
log_probs,
enc_len,
self.tokenizer.vocab_size(),
self.tokenizer.blank_id(),
),
};
let words = ctc::ctc_tokens_to_words(&self.tokenizer, &tokens, frame_offset);
return Ok((words, false)); // CTC has no endpoint signal
}
// RNN-T greedy decode — the encoder output is borrowed for the decode loop.
let dec_start = std::time::Instant::now();
let decoder = triplet
.decoder
.as_deref()
.ok_or_else(|| anyhow::anyhow!("RNN-T decoder session missing for a non-CTC head"))?;
let joiner = triplet
.joiner
.as_deref()
.ok_or_else(|| anyhow::anyhow!("RNN-T joiner session missing for a non-CTC head"))?;
let result = decode::greedy_decode(
decoder,
joiner,
&encoder_outputs[0].view(),
enc_len,
self.tokenizer.blank_id(),
decoder_state,
biaser,
)?;
tracing::info!(
elapsed_ms = dec_start.elapsed().as_millis() as u64,
"greedy_decode"
);
// Convert token infos to words with timestamps
let words = self.tokens_to_words(&result.tokens, frame_offset);
tracing::info!(
tokens = result.tokens.len(),
words = words.len(),
duration_ms = dec_start.elapsed().as_millis() as u64,
"Decoded tokens"
);
Ok((words, result.endpoint_detected))
}
/// Convert decoded tokens into words with timestamps and confidence.
pub(crate) fn tokens_to_words(
&self,
tokens: &[decode::TokenInfo],
frame_offset: usize,
) -> Vec<WordInfo> {
TokenFormatter::tokens_to_words(&self.tokenizer, tokens, frame_offset)
}
}