brain2qwerty 0.0.1

Brain2Qwerty V1/V2 MEG neural decoding inference in Rust (parity-tested vs Python)
Documentation
//! x-transformers-style encoder for V1 sentence transformer (ALiBi + ScaleNorm).

use crate::tensor::Tensor;
use crate::v1::config::V1TransformerConfig;
use crate::weights::{get, param_to_tensor, WeightStore};

fn alibi_slopes(heads: usize) -> Vec<f32> {
    fn slopes_power_of_2(n: usize) -> Vec<f32> {
        let start = 2f32.powf(-2f32.powf(-((n as f32).log2() - 3.0)));
        let ratio = start;
        (0..n).map(|i| start * ratio.powi(i as i32)).collect()
    }
    if heads.is_power_of_two() {
        return slopes_power_of_2(heads);
    }
    let closest = 2usize.pow((heads as f32).log2().floor() as u32);
    let mut out = slopes_power_of_2(closest);
    let extra: Vec<f32> = slopes_power_of_2(2 * closest)
        .into_iter()
        .step_by(2)
        .take(heads - closest)
        .collect();
    out.extend(extra);
    out
}

pub struct ScaleNorm {
    pub g: f32,
    pub dim: usize,
}

impl ScaleNorm {
    pub fn forward(&self, x: &Tensor) -> Tensor {
        let d = self.dim;
        let scale = (d as f32).sqrt();
        // x-transformers uses `unit_offset=True` by default (gamma = g + 1).
        let gamma = self.g + 1.0;
        let (b, n, _) = (x.shape[0], x.shape[1], x.shape[2]);
        let mut out = vec![0.0f32; x.data.len()];
        for bi in 0..b {
            for ni in 0..n {
                let base = (bi * n + ni) * d;
                let mut norm = 0.0f32;
                for i in 0..d {
                    norm += x.data[base + i] * x.data[base + i];
                }
                norm = norm.sqrt().max(1e-12);
                for i in 0..d {
                    out[base + i] = (x.data[base + i] / norm) * scale * gamma;
                }
            }
        }
        Tensor::from_vec(out, x.shape.clone())
    }
}

struct AttnLayer {
    pre_norm: ScaleNorm,
    q_w: Tensor,
    k_w: Tensor,
    v_w: Tensor,
    residual_scale: Option<Tensor>,
}

struct FfLayer {
    pre_norm: ScaleNorm,
    w1: Tensor,
    b1: Tensor,
    w2: Tensor,
    b2: Tensor,
    residual_scale: Option<Tensor>,
}

pub struct SentenceTransformer {
    pub dim: usize,
    pub heads: usize,
    pub head_dim: usize,
    pub alibi_heads: usize,
    pub alibi_slopes: Vec<f32>,
    attn_layers: Vec<AttnLayer>,
    ff_layers: Vec<FfLayer>,
    final_norm: ScaleNorm,
}

impl SentenceTransformer {
    pub fn from_config_and_weights(
        cfg: &V1TransformerConfig,
        dim: usize,
        store: &mut WeightStore,
        prefix: &str,
    ) -> anyhow::Result<Self> {
        anyhow::ensure!(
            !cfg.rotary_pos_emb,
            "rotary_pos_emb is not supported in the Rust V1 transformer yet"
        );
        anyhow::ensure!(dim % cfg.heads == 0, "dim must divide heads");
        let head_dim = dim / cfg.heads;
        let alibi_heads = cfg.heads;
        let mut attn_layers = Vec::new();
        let mut ff_layers = Vec::new();
        for layer in 0..cfg.depth {
            let attn_i = layer * 2;
            let ff_i = layer * 2 + 1;
            let ap = format!("{prefix}layers.{attn_i}.");
            let fp = format!("{prefix}layers.{ff_i}.");
            attn_layers.push(AttnLayer {
                pre_norm: ScaleNorm {
                    g: get(store, &format!("{ap}0.0.g"))
                        .map(|p| p.data[0])
                        .unwrap_or(1.0),
                    dim,
                },
                q_w: param_to_tensor(get(store, &format!("{ap}1.to_q.weight"))?),
                k_w: param_to_tensor(get(store, &format!("{ap}1.to_k.weight"))?),
                v_w: param_to_tensor(get(store, &format!("{ap}1.to_v.weight"))?),
                residual_scale: get(store, &format!("{ap}2.residual_scale"))
                    .ok()
                    .map(param_to_tensor),
            });
            ff_layers.push(FfLayer {
                pre_norm: ScaleNorm {
                    g: get(store, &format!("{fp}0.0.g"))
                        .map(|p| p.data[0])
                        .unwrap_or(1.0),
                    dim,
                },
                w1: param_to_tensor(get(store, &format!("{fp}1.ff.0.0.weight"))?),
                b1: param_to_tensor(get(store, &format!("{fp}1.ff.0.0.bias"))?),
                w2: param_to_tensor(get(store, &format!("{fp}1.ff.2.weight"))?),
                b2: param_to_tensor(get(store, &format!("{fp}1.ff.2.bias"))?),
                residual_scale: get(store, &format!("{fp}2.residual_scale"))
                    .ok()
                    .map(param_to_tensor),
            });
        }
        let final_g = get(store, &format!("{prefix}final_norm.g"))
            .map(|p| p.data[0])
            .unwrap_or(1.0);
        Ok(Self {
            dim,
            heads: cfg.heads,
            head_dim,
            alibi_heads,
            alibi_slopes: alibi_slopes(alibi_heads),
            attn_layers,
            ff_layers,
            final_norm: ScaleNorm { g: final_g, dim },
        })
    }

    fn alibi_bias(&self, n: usize) -> Vec<f32> {
        let h = self.heads;
        let mut bias = vec![0.0f32; h * n * n];
        for hi in 0..h {
            let slope = self.alibi_slopes[hi.min(self.alibi_slopes.len() - 1)];
            for i in 0..n {
                for j in 0..n {
                    bias[hi * n * n + i * n + j] = -((j as f32 - i as f32).abs()) * slope;
                }
            }
        }
        bias
    }

    fn self_attention(&self, layer: &AttnLayer, x: &Tensor, mask: &[bool]) -> Tensor {
        let (b, n, d) = (x.shape[0], x.shape[1], x.shape[2]);
        let h = self.heads;
        let hd = self.head_dim;
        let xn = layer.pre_norm.forward(x);
        let q = xn.linear(&layer.q_w, None);
        let k = xn.linear(&layer.k_w, None);
        let v = xn.linear(&layer.v_w, None);
        let alibi = self.alibi_bias(n);
        let scale = (hd as f32).sqrt();
        let mut out = vec![0.0f32; b * n * d];
        for bi in 0..b {
            for ni in 0..n {
                for hi in 0..h {
                    let mut scores = vec![0.0f32; n];
                    let q_base = (bi * n + ni) * d + hi * hd;
                    for kj in 0..n {
                        if !mask[kj] {
                            scores[kj] = f32::NEG_INFINITY;
                            continue;
                        }
                        let mut dot = 0.0f32;
                        let k_base = (bi * n + kj) * d + hi * hd;
                        for di in 0..hd {
                            dot += q.data[q_base + di] * k.data[k_base + di];
                        }
                        scores[kj] = dot / scale + alibi[hi * n * n + ni * n + kj];
                    }
                    let max_s = scores.iter().copied().fold(f32::NEG_INFINITY, f32::max);
                    let mut denom = 0.0f32;
                    for kj in 0..n {
                        if scores[kj].is_finite() {
                            scores[kj] = (scores[kj] - max_s).exp();
                            denom += scores[kj];
                        } else {
                            scores[kj] = 0.0;
                        }
                    }
                    if denom > 0.0 {
                        for s in &mut scores {
                            *s /= denom;
                        }
                    }
                    let out_base = (bi * n + ni) * d + hi * hd;
                    for di in 0..hd {
                        let mut sum = 0.0f32;
                        for kj in 0..n {
                            let v_base = (bi * n + kj) * d + hi * hd;
                            sum += scores[kj] * v.data[v_base + di];
                        }
                        out[out_base + di] = sum;
                    }
                }
            }
        }
        Tensor::from_vec(out, vec![b, n, d])
    }

    fn feed_forward(&self, layer: &FfLayer, x: &Tensor) -> Tensor {
        let xn = layer.pre_norm.forward(x);
        let mut h = xn.linear(&layer.w1, Some(&layer.b1)).gelu();
        h = h.linear(&layer.w2, Some(&layer.b2));
        h
    }

    fn apply_residual(&self, branch: &Tensor, residual: &Tensor, scale: Option<&Tensor>) -> Tensor {
        let mut out = branch.data.clone();
        if let Some(s) = scale {
            for i in 0..out.len() {
                out[i] += residual.data[i] * s.data[i % s.data.len()];
            }
        } else {
            for (o, &r) in out.iter_mut().zip(residual.data.iter()) {
                *o += r;
            }
        }
        Tensor::from_vec(out, branch.shape.clone())
    }

    /// `x`: (B, N, D); `mask`: length N, true = valid token.
    pub fn forward(&self, x: &Tensor, mask: &[bool]) -> Tensor {
        assert_eq!(x.ndim(), 3);
        assert_eq!(x.shape[2], self.dim);
        assert_eq!(mask.len(), x.shape[1]);
        let mut h = x.clone();
        for (attn, ff) in self.attn_layers.iter().zip(self.ff_layers.iter()) {
            let residual = h.clone();
            let attn_out = self.self_attention(attn, &h, mask);
            h = self.apply_residual(&attn_out, &residual, attn.residual_scale.as_ref());
            let residual = h.clone();
            let ff_out = self.feed_forward(ff, &h);
            h = self.apply_residual(&ff_out, &residual, ff.residual_scale.as_ref());
        }
        self.final_norm.forward(&h)
    }
}