brain2qwerty 0.0.1

Brain2Qwerty V1/V2 MEG neural decoding inference in Rust (parity-tested vs Python)
Documentation
//! TinyLlama / Llama-shaped causal LM (CPU reference for parity).

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

#[derive(Clone, Debug)]
pub struct LlamaConfig {
    pub hidden_size: usize,
    pub num_hidden_layers: usize,
    pub num_attention_heads: usize,
    pub num_key_value_heads: usize,
    pub intermediate_size: usize,
    pub vocab_size: usize,
    pub rms_norm_eps: f64,
    pub rope_theta: f32,
}

impl LlamaConfig {
    pub fn tinyllama_1_1b() -> Self {
        Self {
            hidden_size: 2048,
            num_hidden_layers: 22,
            num_attention_heads: 32,
            num_key_value_heads: 4,
            intermediate_size: 5632,
            vocab_size: 32000,
            rms_norm_eps: 1e-5,
            rope_theta: 10000.0,
        }
    }

    pub fn from_json(path: &str) -> anyhow::Result<Self> {
        let v: serde_json::Value = serde_json::from_str(&std::fs::read_to_string(path)?)?;
        let rope_theta = v["rope_parameters"]["rope_theta"]
            .as_f64()
            .or_else(|| v["rope_theta"].as_f64())
            .unwrap_or(10000.0) as f32;
        Ok(Self {
            hidden_size: v["hidden_size"].as_u64().unwrap_or(2048) as usize,
            num_hidden_layers: v["num_hidden_layers"].as_u64().unwrap_or(22) as usize,
            num_attention_heads: v["num_attention_heads"].as_u64().unwrap_or(32) as usize,
            num_key_value_heads: v["num_key_value_heads"].as_u64().unwrap_or(4) as usize,
            intermediate_size: v["intermediate_size"].as_u64().unwrap_or(5632) as usize,
            vocab_size: v["vocab_size"].as_u64().unwrap_or(32000) as usize,
            rms_norm_eps: v["rms_norm_eps"].as_f64().unwrap_or(1e-5),
            rope_theta,
        })
    }

    pub fn head_dim(&self) -> usize {
        self.hidden_size / self.num_attention_heads
    }
}

#[derive(Clone)]
pub struct LlamaLayer {
    pub attn_norm_w: Tensor,
    pub q_w: Tensor,
    pub k_w: Tensor,
    pub v_w: Tensor,
    pub o_w: Tensor,
    pub ffn_norm_w: Tensor,
    pub gate_w: Tensor,
    pub up_w: Tensor,
    pub down_w: Tensor,
}

#[derive(Clone)]
pub struct TinyLlama {
    pub config: LlamaConfig,
    pub embed_tokens: Tensor,
    pub layers: Vec<LlamaLayer>,
    pub norm_w: Tensor,
    pub lm_head: Option<Tensor>,
}

impl TinyLlama {
    pub fn load(store: &WeightStore, config: LlamaConfig, prefix: &str) -> anyhow::Result<Self> {
        let embed = param_to_tensor(get(store, &format!("{prefix}embed_tokens.weight"))?);
        let mut layers = Vec::new();
        for i in 0..config.num_hidden_layers {
            let p = format!("{prefix}layers.{i}.");
            layers.push(LlamaLayer {
                attn_norm_w: param_to_tensor(get(store, &format!("{p}input_layernorm.weight"))?),
                q_w: param_to_tensor(get(store, &format!("{p}self_attn.q_proj.weight"))?),
                k_w: param_to_tensor(get(store, &format!("{p}self_attn.k_proj.weight"))?),
                v_w: param_to_tensor(get(store, &format!("{p}self_attn.v_proj.weight"))?),
                o_w: param_to_tensor(get(store, &format!("{p}self_attn.o_proj.weight"))?),
                ffn_norm_w: param_to_tensor(get(
                    store,
                    &format!("{p}post_attention_layernorm.weight"),
                )?),
                gate_w: param_to_tensor(get(store, &format!("{p}mlp.gate_proj.weight"))?),
                up_w: param_to_tensor(get(store, &format!("{p}mlp.up_proj.weight"))?),
                down_w: param_to_tensor(get(store, &format!("{p}mlp.down_proj.weight"))?),
            });
        }
        let norm_w = param_to_tensor(get(store, &format!("{prefix}norm.weight"))?);
        let lm_head = get(store, &format!("{prefix}lm_head.weight"))
            .ok()
            .map(param_to_tensor);
        Ok(Self {
            config,
            embed_tokens: embed,
            layers,
            norm_w,
            lm_head,
        })
    }

    pub fn embed(&self, token_ids: &[usize]) -> Tensor {
        let d = self.config.hidden_size;
        let mut data = vec![0.0f32; token_ids.len() * d];
        for (i, &tid) in token_ids.iter().enumerate() {
            let src = tid.min(self.embed_tokens.shape[0] - 1) * d;
            data[i * d..(i + 1) * d].copy_from_slice(&self.embed_tokens.data[src..src + d]);
        }
        Tensor::from_vec(data, vec![1, token_ids.len(), d])
    }

    pub fn forward_embeds(&self, input_embeds: &Tensor, attention_mask: Option<&[f32]>) -> Tensor {
        let mut h = input_embeds.clone();
        for layer in &self.layers {
            h = self.layer_forward(layer, &h, attention_mask);
        }
        self.rms_norm(&h, &self.norm_w)
    }

    pub fn logits(&self, hidden: &Tensor) -> Tensor {
        let head = self.lm_head.as_ref().unwrap_or(&self.embed_tokens);
        hidden.linear(head, None)
    }

    fn layer_forward(
        &self,
        layer: &LlamaLayer,
        x: &Tensor,
        attention_mask: Option<&[f32]>,
    ) -> Tensor {
        let residual = x.clone();
        let xn = self.rms_norm(x, &layer.attn_norm_w);
        let attn = self.self_attn(layer, &xn, attention_mask);
        let x = residual.add(&attn);
        let residual = x.clone();
        let xn = self.rms_norm(&x, &layer.ffn_norm_w);
        let gate = xn.linear(&layer.gate_w, None).silu();
        let up = xn.linear(&layer.up_w, None);
        let ff = gate.mul(&up).linear(&layer.down_w, None);
        residual.add(&ff)
    }

    fn rms_norm(&self, x: &Tensor, weight: &Tensor) -> Tensor {
        let (b, t, d) = (x.shape[0], x.shape[1], x.shape[2]);
        let eps = self.config.rms_norm_eps;
        let mut out = vec![0.0f32; x.data.len()];
        for bi in 0..b {
            for ti in 0..t {
                let base = (bi * t + ti) * d;
                let var: f32 = x.data[base..base + d].iter().map(|v| v * v).sum::<f32>() / d as f32;
                let inv = (var + eps as f32).sqrt().recip();
                for j in 0..d {
                    out[base + j] = x.data[base + j] * inv * weight.data[j];
                }
            }
        }
        Tensor::from_vec(out, x.shape.clone())
    }

    fn self_attn(&self, layer: &LlamaLayer, x: &Tensor, attention_mask: Option<&[f32]>) -> Tensor {
        let cfg = &self.config;
        let (b, t, d) = (x.shape[0], x.shape[1], x.shape[2]);
        let nh = cfg.num_attention_heads;
        let nkv = cfg.num_key_value_heads;
        let dh = cfg.head_dim();
        let mut q = x.linear(&layer.q_w, None);
        let mut k = x.linear(&layer.k_w, None);
        let v = x.linear(&layer.v_w, None);
        self.apply_rope(&mut q, dh);
        self.apply_rope(&mut k, dh);
        let scale = (dh as f32).sqrt().recip();
        let mut out = Tensor::zeros(&[b, t, d]);
        for bi in 0..b {
            for hi in 0..nh {
                let k_hi = hi * nkv / nh;
                for ti in 0..t {
                    let mut scores = vec![0.0f32; ti + 1];
                    for tj in 0..=ti {
                        if attention_mask.is_some_and(|m| m.get(tj).copied().unwrap_or(1.0) <= 0.0)
                        {
                            scores[tj] = f32::NEG_INFINITY;
                            continue;
                        }
                        let mut dot = 0.0f32;
                        for j in 0..dh {
                            let q_idx = (bi * t + ti) * d + hi * dh + j;
                            let k_idx = (bi * t + tj) * d + k_hi * dh + j;
                            dot += q.data[q_idx] * k.data[k_idx];
                        }
                        scores[tj] = dot * scale;
                    }
                    let max_s = scores.iter().copied().fold(f32::NEG_INFINITY, f32::max);
                    let mut sum = 0.0f32;
                    for s in &mut scores {
                        if s.is_finite() {
                            *s = (*s - max_s).exp();
                            sum += *s;
                        } else {
                            *s = 0.0;
                        }
                    }
                    if sum > 0.0 {
                        for s in &mut scores {
                            *s /= sum;
                        }
                    }
                    for j in 0..dh {
                        let mut acc = 0.0f32;
                        for (tj, &sc) in scores.iter().enumerate() {
                            let v_idx = (bi * t + tj) * d + k_hi * dh + j;
                            acc += sc * v.data[v_idx];
                        }
                        out.data[(bi * t + ti) * d + hi * dh + j] = acc;
                    }
                }
            }
        }
        out.linear(&layer.o_w, None)
    }

    fn apply_rope(&self, x: &mut Tensor, dh: usize) {
        let (b, t, d) = (x.shape[0], x.shape[1], x.shape[2]);
        let nh = d / dh;
        let half = dh / 2;
        for bi in 0..b {
            for ti in 0..t {
                for hi in 0..nh {
                    let base = (bi * t + ti) * d + hi * dh;
                    for i in 0..half {
                        let theta = self.config.rope_theta.powf(-2.0 * i as f32 / dh as f32);
                        let angle = ti as f32 * theta;
                        let cos = angle.cos();
                        let sin = angle.sin();
                        let x0 = x.data[base + i];
                        let x1 = x.data[base + half + i];
                        x.data[base + i] = x0 * cos - x1 * sin;
                        x.data[base + half + i] = x0 * sin + x1 * cos;
                    }
                }
            }
        }
    }
}