brain2qwerty 0.0.1

Brain2Qwerty V1/V2 MEG neural decoding inference in Rust (parity-tested vs Python)
Documentation
//! CTC space segmenter + intra-word MLP pooler.

use crate::decode::ctc::SPACE_IDX;
use crate::tensor::Tensor;

pub struct IntraWordPooler {
    pub layers: Vec<(Tensor, Tensor, Tensor, Tensor)>,
}

pub fn build_intra_word_pooler(
    dim: usize,
    n_layers: usize,
    store: Option<&crate::weights::WeightStore>,
) -> IntraWordPooler {
    let mut layers = Vec::new();
    if let Some(store) = store {
        for i in 0..n_layers {
            let linear_idx = i * 3;
            let ln_idx = linear_idx + 1;
            let p = "word_segmenter.intra_word_pooler.mlp.";
            if let (Ok(w), Ok(b), Ok(lnw), Ok(lnb)) = (
                crate::weights::get(store, &format!("{p}{linear_idx}.weight")),
                crate::weights::get(store, &format!("{p}{linear_idx}.bias")),
                crate::weights::get(store, &format!("{p}{ln_idx}.weight")),
                crate::weights::get(store, &format!("{p}{ln_idx}.bias")),
            ) {
                layers.push((
                    crate::weights::param_to_tensor(w),
                    crate::weights::param_to_tensor(b),
                    crate::weights::param_to_tensor(lnw),
                    crate::weights::param_to_tensor(lnb),
                ));
            }
        }
    }
    if layers.is_empty() {
        for li in 0..n_layers {
            let in_d = if li == 0 { dim } else { dim };
            let out_d = dim;
            layers.push((
                identity_linear(in_d, out_d),
                Tensor::zeros(&[out_d]),
                Tensor::from_vec(vec![1.0; out_d], vec![out_d]),
                Tensor::zeros(&[out_d]),
            ));
        }
    }
    IntraWordPooler { layers }
}

fn identity_linear(in_d: usize, out_d: usize) -> Tensor {
    let mut data = vec![0.0f32; in_d * out_d];
    for i in 0..in_d.min(out_d) {
        data[i * out_d + i] = 1.0;
    }
    Tensor::from_vec(data, vec![out_d, in_d])
}

impl IntraWordPooler {
    pub fn forward_frames(&self, frames: &Tensor) -> Tensor {
        let (n, d) = (frames.shape[0], frames.shape[1]);
        let mut x = frames.reshape(&[1, n, d]);
        for (w, b, lnw, lnb) in &self.layers {
            x = x.linear(w, Some(b)).layer_norm(lnw, lnb, 1e-5).gelu();
        }
        let x = x.reshape(&[n, d]);
        let mut mean = vec![0.0f32; d];
        for fi in 0..n {
            for j in 0..d {
                mean[j] += x.data[fi * d + j];
            }
        }
        for v in &mut mean {
            *v /= n as f32;
        }
        Tensor::from_vec(mean, vec![d])
    }
}

pub struct CTCSpaceSegmenter {
    pub include_blanks: bool,
    pub min_word_frames: usize,
    pub pooler: IntraWordPooler,
}

impl CTCSpaceSegmenter {
    pub fn forward(&self, z_final: &Tensor, ctc_logits: &Tensor) -> Vec<Tensor> {
        assert_eq!(z_final.ndim(), 3);
        let (b, t, d) = (z_final.shape[0], z_final.shape[1], z_final.shape[2]);
        let c = ctc_logits.shape[2];
        let mut results = Vec::with_capacity(b);
        for bi in 0..b {
            let mut preds = vec![0usize; t];
            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;
                    }
                }
                preds[ti] = best;
            }
            let mut segments: Vec<Vec<usize>> = Vec::new();
            let mut current: Vec<usize> = Vec::new();
            for ti in 0..t {
                if preds[ti] == SPACE_IDX {
                    if !current.is_empty() {
                        segments.push(current);
                        current = Vec::new();
                    }
                } else if self.include_blanks || preds[ti] != 0 {
                    current.push(ti);
                }
            }
            if !current.is_empty() {
                segments.push(current);
            }
            segments.retain(|s| s.len() >= self.min_word_frames);
            if segments.is_empty() {
                let mut all_frames = Vec::new();
                for ti in 0..t {
                    all_frames
                        .push(z_final.data[(bi * t + ti) * d..(bi * t + ti + 1) * d].to_vec());
                }
                let flat: Vec<f32> = all_frames.into_iter().flatten().collect();
                let frames = Tensor::from_vec(flat, vec![t, d]);
                results.push(self.pooler.forward_frames(&frames).reshape(&[1, d]));
            } else {
                let mut embeds = Vec::new();
                for seg in segments {
                    let n = seg.len();
                    let mut flat = vec![0.0f32; n * d];
                    for (si, &ti) in seg.iter().enumerate() {
                        let src = (bi * t + ti) * d;
                        flat[si * d..(si + 1) * d].copy_from_slice(&z_final.data[src..src + d]);
                    }
                    let frames = Tensor::from_vec(flat, vec![n, d]);
                    embeds.push(self.pooler.forward_frames(&frames));
                }
                let n_words = embeds.len();
                let flat: Vec<f32> = embeds.into_iter().flat_map(|t| t.data).collect();
                results.push(Tensor::from_vec(flat, vec![n_words, d]));
            }
        }
        results
    }
}