brain2qwerty 0.0.1

Brain2Qwerty V1/V2 MEG neural decoding inference in Rust (parity-tested vs Python)
Documentation
//! Channel merger (neuraltrain ChannelMergerModel).

use crate::model::fourier::FourierEmb;
use crate::tensor::Tensor;

pub const INVALID_POS: f32 = -0.1;

pub struct ChannelMerger {
    pub embedding: FourierEmb,
    pub heads: Tensor,
    pub per_subject: bool,
    pub n_virtual: usize,
    pub invalid_value: f32,
}

impl ChannelMerger {
    pub fn forward(&self, meg: &Tensor, subject_ids: &[usize], positions: &Tensor) -> Tensor {
        assert_eq!(meg.ndim(), 3);
        let (b, c_in, t) = (meg.shape[0], meg.shape[1], meg.shape[2]);
        let embedding = self.embedding.forward(positions);
        let pos_dim = self.embedding.total_dim;
        let mut out = Tensor::zeros(&[b, self.n_virtual, t]);

        for bi in 0..b {
            let sid = subject_ids.get(bi).copied().unwrap_or(0);
            let heads = if self.per_subject {
                self.gather_heads(sid)
            } else {
                self.heads.clone()
            };

            for vo in 0..self.n_virtual {
                let mut scores = vec![f32::NEG_INFINITY; c_in];
                for ci in 0..c_in {
                    let pos_base = (bi * c_in + ci) * 2;
                    if positions.data[pos_base] == self.invalid_value
                        && positions.data[pos_base + 1] == self.invalid_value
                    {
                        continue;
                    }
                    let emb_base = (bi * c_in + ci) * pos_dim;
                    let head_base = vo * pos_dim;
                    let mut dot = 0.0f32;
                    for j in 0..pos_dim {
                        dot += embedding.data[emb_base + j] * heads.data[head_base + j];
                    }
                    scores[ci] = dot;
                }
                let max_s = scores.iter().copied().fold(f32::NEG_INFINITY, f32::max);
                let mut sum = 0.0f32;
                let mut weights = vec![0.0f32; c_in];
                for ci in 0..c_in {
                    if scores[ci].is_finite() {
                        weights[ci] = (scores[ci] - max_s).exp();
                        sum += weights[ci];
                    }
                }
                if sum > 0.0 {
                    for w in &mut weights {
                        *w /= sum;
                    }
                }
                for ti in 0..t {
                    let mut v = 0.0f32;
                    for ci in 0..c_in {
                        v += weights[ci] * meg.data[(bi * c_in + ci) * t + ti];
                    }
                    out.data[(bi * self.n_virtual + vo) * t + ti] =
                        if v.is_nan() { 0.0 } else { v };
                }
            }
        }
        out
    }

    fn gather_heads(&self, subject_id: usize) -> Tensor {
        let pos_dim = self.embedding.total_dim;
        let n_subjects = self.heads.shape[0];
        let sid = subject_id.min(n_subjects.saturating_sub(1));
        let n_virt = self.heads.shape[1];
        let mut data = vec![0.0f32; n_virt * pos_dim];
        let src_base = sid * n_virt * pos_dim;
        data.copy_from_slice(&self.heads.data[src_base..src_base + n_virt * pos_dim]);
        Tensor::from_vec(data, vec![n_virt, pos_dim])
    }
}