brain2qwerty 0.0.1

Brain2Qwerty V1/V2 MEG neural decoding inference in Rust (parity-tested vs Python)
Documentation
//! Per-subject linear projection (neuraltrain SubjectLayersModel).

use crate::tensor::Tensor;

pub struct SubjectLayers {
    pub weights: Tensor,
    pub bias: Option<Tensor>,
    pub average_subjects: bool,
    pub n_subjects: usize,
}

impl SubjectLayers {
    pub fn forward(&self, x: &Tensor, subject_ids: &[usize]) -> Tensor {
        assert_eq!(x.ndim(), 3);
        let (b, c_in, t) = (x.shape[0], x.shape[1], x.shape[2]);
        let c_out = self.weights.shape[2];
        let mut out = Tensor::zeros(&[b, c_out, t]);

        for bi in 0..b {
            let (w, bias) = if self.average_subjects {
                let sid = self.n_subjects;
                (
                    self.slice_weight(sid.min(self.weights.shape[0] - 1)),
                    self.bias
                        .as_ref()
                        .map(|bias| self.slice_bias(sid.min(bias.shape[0] - 1))),
                )
            } else {
                let sid = subject_ids.get(bi).copied().unwrap_or(0);
                let bias = self.bias.as_ref().map(|_| self.slice_bias(sid));
                (self.slice_weight(sid), bias)
            };
            for ti in 0..t {
                for co in 0..c_out {
                    let mut sum = 0.0f32;
                    for ci in 0..c_in {
                        sum += x.data[(bi * c_in + ci) * t + ti] * w.data[ci * c_out + co];
                    }
                    if let Some(ref bias) = bias {
                        sum += bias.data[co];
                    }
                    out.data[(bi * c_out + co) * t + ti] = sum;
                }
            }
        }
        out
    }

    fn slice_weight(&self, sid: usize) -> Tensor {
        let c_in = self.weights.shape[1];
        let c_out = self.weights.shape[2];
        let base = sid * c_in * c_out;
        Tensor::from_vec(
            self.weights.data[base..base + c_in * c_out].to_vec(),
            vec![c_in, c_out],
        )
    }

    fn slice_bias(&self, sid: usize) -> Tensor {
        let c_out = self.weights.shape[2];
        let base = sid * c_out;
        Tensor::from_vec(
            self.bias.as_ref().unwrap().data[base..base + c_out].to_vec(),
            vec![c_out],
        )
    }
}