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;
}
}
}
}
}
}