brain2qwerty 0.0.1

Brain2Qwerty V1/V2 MEG neural decoding inference in Rust (parity-tested vs Python)
Documentation
//! CTC greedy decode (brain2qwerty_v2/utils.py).

use crate::tensor::Tensor;

pub const NUM_CLASSES: usize = 29;
pub const SPACE_IDX: usize = 27;

pub fn letters_withblank() -> Vec<char> {
    let mut v = vec!['-'];
    v.extend('a'..='z');
    v.push('&');
    v
}

pub fn label_to_text(ids: &[usize]) -> String {
    let vocab = letters_withblank();
    ids.iter()
        .filter_map(|&i| {
            if i > 0 && i < vocab.len() {
                Some(if vocab[i] == '&' { ' ' } else { vocab[i] })
            } else {
                None
            }
        })
        .collect()
}

pub fn ctc_greedy_decode(ctc_logits: &Tensor) -> Vec<String> {
    assert_eq!(ctc_logits.ndim(), 3);
    let (b, t, c) = (
        ctc_logits.shape[0],
        ctc_logits.shape[1],
        ctc_logits.shape[2],
    );
    let vocab = letters_withblank();
    let mut texts = Vec::with_capacity(b);
    for bi in 0..b {
        let mut chars = Vec::new();
        let mut prev = 0usize;
        for ti in 0..t {
            let base = (bi * t + ti) * c;
            let mut best = 0usize;
            let mut best_v = ctc_logits.data[base];
            for j in 1..c {
                if ctc_logits.data[base + j] > best_v {
                    best_v = ctc_logits.data[base + j];
                    best = j;
                }
            }
            if best != prev && best != 0 && best < vocab.len() {
                let ch = vocab[best];
                chars.push(if ch == '&' { ' ' } else { ch });
            }
            prev = best;
        }
        texts.push(chars.into_iter().collect());
    }
    texts
}