#[cfg(any(test, feature = "bench"))]
use ndarray::{Array2, ArrayViewMut1};
use ndarray::{ArrayView1, ArrayView2};
use super::charset::Charset;
use super::recognizer::RecognizedText;
const BLANK_CLASS: usize = 0;
const CUSTOM_MEAN_EXPONENT_NUMERATOR: f32 = 2.0;
#[cfg(any(test, feature = "bench"))]
pub(crate) fn decode_greedy(logits: ArrayView2<f32>, charset: &Charset, ignore: &[usize]) -> RecognizedText {
let ignore_mask = IgnoreMask::new(logits.ncols(), ignore);
decode_greedy_with_mask(logits, charset, &ignore_mask)
}
pub(super) fn decode_greedy_with_mask(
logits: ArrayView2<f32>,
charset: &Charset,
ignore_mask: &IgnoreMask,
) -> RecognizedText {
let mut weights: Vec<f32> = Vec::with_capacity(logits.ncols());
let per_timestep: Vec<(usize, f32)> = logits
.rows()
.into_iter()
.map(|row| decode_row(row, ignore_mask, &mut weights))
.collect();
let confidence = custom_mean(&collect_max_probs(&per_timestep));
let text = collapse(&per_timestep, charset);
RecognizedText { text, confidence }
}
pub(super) struct IgnoreMask {
classes: Vec<bool>,
}
impl IgnoreMask {
pub(super) fn new(num_classes: usize, ignore: &[usize]) -> Self {
let mut classes = vec![false; num_classes];
for &class in ignore {
if let Some(slot) = classes.get_mut(class) {
*slot = true;
}
}
Self { classes }
}
fn contains(&self, class: usize) -> bool {
self.classes.get(class).copied().unwrap_or(false)
}
}
fn decode_row(input: ArrayView1<f32>, ignore_mask: &IgnoreMask, weights: &mut Vec<f32>) -> (usize, f32) {
let max = input.iter().copied().fold(f32::NEG_INFINITY, f32::max);
weights.clear();
let mut sum = 0.0f32;
let mut best_class = BLANK_CLASS;
let mut best_weight = f32::NEG_INFINITY;
for (class, &value) in input.iter().enumerate() {
let weight = (value - max).exp();
weights.push(weight);
sum += weight;
if !ignore_mask.contains(class) && weight > best_weight {
best_weight = weight;
best_class = class;
}
}
let mut renorm = 0.0f32;
for (class, &weight) in weights.iter().enumerate() {
if !ignore_mask.contains(class) {
renorm += weight / sum;
}
}
if renorm > 0.0 {
(best_class, (best_weight / sum) / renorm)
} else {
(BLANK_CLASS, 0.0)
}
}
#[cfg(any(test, feature = "bench"))]
fn probability_rows(logits: ArrayView2<f32>, ignore: &[usize]) -> Array2<f32> {
let mut probs = Array2::<f32>::zeros(logits.raw_dim());
for (out_row, in_row) in probs.rows_mut().into_iter().zip(logits.rows()) {
fill_probability_row(out_row, in_row, ignore);
}
probs
}
#[cfg(any(test, feature = "bench"))]
fn fill_probability_row(mut out: ArrayViewMut1<f32>, input: ArrayView1<f32>, ignore: &[usize]) {
let max = input.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0f32;
for (cell, &value) in out.iter_mut().zip(input.iter()) {
let weight = (value - max).exp();
*cell = weight;
sum += weight;
}
if sum > 0.0 {
out.iter_mut().for_each(|cell| *cell /= sum);
}
for &class in ignore {
if let Some(cell) = out.get_mut(class) {
*cell = 0.0;
}
}
let renorm: f32 = out.iter().sum();
if renorm > 0.0 {
out.iter_mut().for_each(|cell| *cell /= renorm);
}
}
#[cfg(any(test, feature = "bench"))]
fn argmax_row(row: ArrayView1<f32>) -> (usize, f32) {
let mut best_class = BLANK_CLASS;
let mut best_prob = f32::NEG_INFINITY;
for (class, &prob) in row.iter().enumerate() {
if prob > best_prob {
best_prob = prob;
best_class = class;
}
}
(best_class, best_prob)
}
#[cfg(any(test, feature = "bench"))]
pub(crate) fn decode_greedy_reference(logits: ArrayView2<f32>, charset: &Charset, ignore: &[usize]) -> RecognizedText {
let probs = probability_rows(logits, ignore);
let per_timestep: Vec<(usize, f32)> = probs.rows().into_iter().map(argmax_row).collect();
let confidence = custom_mean(&collect_max_probs(&per_timestep));
let text = collapse(&per_timestep, charset);
RecognizedText { text, confidence }
}
fn collect_max_probs(per_timestep: &[(usize, f32)]) -> Vec<f32> {
let max_probs: Vec<f32> = per_timestep
.iter()
.filter(|(class, _)| *class != BLANK_CLASS)
.map(|(_, prob)| *prob)
.collect();
if max_probs.is_empty() { vec![0.0] } else { max_probs }
}
fn collapse(per_timestep: &[(usize, f32)], charset: &Charset) -> String {
let mut text = String::new();
let mut previous: Option<usize> = None;
for &(class, _) in per_timestep {
let is_new = previous != Some(class);
previous = Some(class);
if is_new && class != BLANK_CLASS {
if let Some(character) = charset.char_at_class(class) {
text.push(character);
}
}
}
text
}
fn custom_mean(values: &[f32]) -> f32 {
let product: f32 = values.iter().product();
let exponent = CUSTOM_MEAN_EXPONENT_NUMERATOR / (values.len() as f32).sqrt();
product.powf(exponent)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::Language;
use ndarray::arr2;
fn english() -> Charset {
Charset::for_language(Language::English)
}
#[test]
fn should_decode_two_distinct_timesteps_into_two_chars() {
let logits = arr2(&[[0.0f32, 5.0, 0.0], [0.0, 0.0, 5.0]]);
let decoded = decode_greedy(logits.view(), &english(), &[]);
assert_eq!(decoded.text, "01");
assert!(decoded.confidence > 0.0, "two confident timesteps score above zero");
}
#[test]
fn should_collapse_repeated_argmax_into_single_char() {
let logits = arr2(&[[0.0f32, 5.0, 0.0], [0.0, 5.0, 0.0]]);
let decoded = decode_greedy(logits.view(), &english(), &[]);
assert_eq!(decoded.text, "0");
}
#[test]
fn should_return_empty_text_and_zero_confidence_when_all_blank() {
let logits = arr2(&[[5.0f32, 0.0, 0.0], [5.0, 0.0, 0.0]]);
let decoded = decode_greedy(logits.view(), &english(), &[]);
assert_eq!(decoded.text, "");
assert_eq!(decoded.confidence, 0.0);
}
#[test]
fn should_compute_exact_custom_mean_confidence() {
let logits = arr2(&[[(0.25f32).ln(), (0.75f32).ln()]]);
let decoded = decode_greedy(logits.view(), &english(), &[]);
assert_eq!(decoded.text, "0");
assert!(
(decoded.confidence - 0.5625).abs() < 1e-5,
"confidence {} should equal 0.5625",
decoded.confidence
);
}
#[test]
fn should_let_ignore_change_the_decoded_character() {
let logits = arr2(&[[(0.1f32).ln(), (0.6f32).ln(), (0.3f32).ln()]]);
let without_ignore = decode_greedy(logits.view(), &english(), &[]);
assert_eq!(without_ignore.text, "0");
let with_ignore = decode_greedy(logits.view(), &english(), &[1]);
assert_eq!(with_ignore.text, "1");
}
#[test]
fn prebuilt_ignore_mask_matches_index_decoder_bitwise() {
let charset = english();
let logits = arr2(&[
[(0.1f32).ln(), (0.6f32).ln(), (0.3f32).ln()],
[(0.1f32).ln(), (0.2f32).ln(), (0.7f32).ln()],
]);
let ignored = [1usize];
let expected = decode_greedy(logits.view(), &charset, &ignored);
let mask = IgnoreMask::new(charset.num_classes(), &ignored);
let actual = decode_greedy_with_mask(logits.view(), &charset, &mask);
assert_eq!(actual.text, expected.text);
assert_eq!(actual.confidence.to_bits(), expected.confidence.to_bits());
}
fn pseudo_random_logits(timesteps: usize, classes: usize, seed: u32) -> Array2<f32> {
let mut state = seed;
let mut next = || {
state = state.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
(state >> 8) as f32 / (1u32 << 24) as f32 * 16.0 - 8.0
};
Array2::from_shape_fn((timesteps, classes), |_| next())
}
#[test]
fn fused_decode_matches_reference_bitwise() {
let charset = english();
let classes = charset.num_classes();
let shapes = [(1usize, classes), (5, classes), (16, classes), (37, 10), (128, classes)];
let ignores: [Vec<usize>; 4] = [vec![], vec![0], vec![1, 2, 3], vec![0, 5, 40, 96]];
for (index, (timesteps, class_count)) in shapes.into_iter().enumerate() {
let logits = pseudo_random_logits(timesteps, class_count, 0x9E37_79B9 ^ index as u32);
for ignore in &ignores {
let fused = decode_greedy(logits.view(), &charset, ignore);
let reference = decode_greedy_reference(logits.view(), &charset, ignore);
assert_eq!(
fused.text, reference.text,
"shape {timesteps}x{class_count}, ignore {ignore:?}"
);
assert_eq!(
fused.confidence.to_bits(),
reference.confidence.to_bits(),
"confidence differs bitwise for shape {timesteps}x{class_count}, ignore {ignore:?}"
);
}
}
}
#[test]
fn fused_decode_matches_reference_on_small_fixtures() {
let charset = english();
let fixtures = [
arr2(&[[0.0f32, 5.0, 0.0], [0.0, 0.0, 5.0]]),
arr2(&[[0.0f32, 5.0, 0.0], [0.0, 5.0, 0.0]]),
arr2(&[[5.0f32, 0.0, 0.0], [5.0, 0.0, 0.0]]),
arr2(&[[(0.25f32).ln(), (0.75f32).ln()]]),
arr2(&[[(0.1f32).ln(), (0.6f32).ln(), (0.3f32).ln()]]),
];
for logits in &fixtures {
for ignore in [Vec::new(), vec![1usize]] {
let fused = decode_greedy(logits.view(), &charset, &ignore);
let reference = decode_greedy_reference(logits.view(), &charset, &ignore);
assert_eq!(fused.text, reference.text, "text differs for ignore {ignore:?}");
assert_eq!(
fused.confidence.to_bits(),
reference.confidence.to_bits(),
"confidence differs bitwise for ignore {ignore:?}"
);
}
}
}
}