brain2qwerty 0.0.1

Brain2Qwerty V1/V2 MEG neural decoding inference in Rust (parity-tested vs Python)
Documentation
//! Bahdanau attention temporal aggregation (neuraltrain `BahdanauAttention`).

use crate::tensor::Tensor;

pub struct BahdanauAttention {
    pub wa_w: Tensor,
    pub wa_b: Tensor,
    pub va_w: Tensor,
    pub va_b: Tensor,
}

impl BahdanauAttention {
    /// `keys`: (B, F, T) channel-first conv features.
    pub fn forward(&self, keys: &Tensor) -> Tensor {
        assert_eq!(keys.ndim(), 3);
        let (b, f, t) = (keys.shape[0], keys.shape[1], keys.shape[2]);
        // (B, T, F)
        let keys_t = keys.transpose(&[0, 2, 1]);
        let mut sum = keys_t.linear(&self.wa_w, Some(&self.wa_b));
        let sum_data: Vec<f32> = sum.data.iter().map(|x| x.tanh()).collect();
        sum = Tensor::from_vec(sum_data, sum.shape.clone());
        let scores = sum.linear(&self.va_w, Some(&self.va_b));
        // scores: (B, T, 1) -> softmax over T
        assert_eq!(scores.shape, vec![b, t, 1]);
        let mut weights = vec![0.0f32; b * t];
        for bi in 0..b {
            let mut max_s = f32::NEG_INFINITY;
            for ti in 0..t {
                max_s = max_s.max(scores.data[(bi * t + ti) * 1]);
            }
            let mut denom = 0.0f32;
            for ti in 0..t {
                let e = (scores.data[(bi * t + ti) * 1] - max_s).exp();
                weights[bi * t + ti] = e;
                denom += e;
            }
            for ti in 0..t {
                weights[bi * t + ti] /= denom;
            }
        }
        // context = weights @ keys_t -> (B, 1, F)
        let mut context = vec![0.0f32; b * f];
        for bi in 0..b {
            for fi in 0..f {
                let mut v = 0.0f32;
                for ti in 0..t {
                    v += weights[bi * t + ti] * keys_t.data[(bi * t + ti) * f + fi];
                }
                context[bi * f + fi] = v;
            }
        }
        // (B, F, 1)
        Tensor::from_vec(context, vec![b, f, 1])
    }
}