use anyhow::{Context, Result, bail, ensure};
use crate::backend::cpu;
use crate::backend::cpu::RopeType;
use crate::gguf::GgufFile;
use crate::kv_cache::InferenceState;
use crate::model::transformer::{self, AttnDims, AttnExtras, AttnWeights, FfnWeights, WeightRef};
use crate::model::{BlockType, Model, ModelConfig, ScalarMultipliers};
struct LayerWeightRefs {
attn_q: WeightRef,
attn_k: WeightRef,
attn_v: WeightRef,
attn_output: WeightRef,
ffn_gate: WeightRef,
ffn_up: WeightRef,
ffn_down: WeightRef,
}
pub struct LlamaModel {
gguf: GgufFile,
config: ModelConfig,
head_dim: usize,
rope_type: RopeType,
rope_freqs: Option<Vec<f32>>,
output_norm_weight: Vec<f32>,
attn_norm_weights: Vec<Vec<f32>>,
ffn_norm_weights: Vec<Vec<f32>>,
attn_q_norm_weights: Vec<Option<Vec<f32>>>,
attn_k_norm_weights: Vec<Option<Vec<f32>>>,
attn_q_bias: Vec<Option<Vec<f32>>>,
attn_k_bias: Vec<Option<Vec<f32>>>,
attn_v_bias: Vec<Option<Vec<f32>>>,
embd_ref: WeightRef,
output_ref: Option<WeightRef>,
layer_refs: Vec<LayerWeightRefs>,
#[allow(dead_code)]
model_id: String,
}
impl LlamaModel {
#[allow(dead_code)]
pub fn from_gguf(gguf: GgufFile, context_size: usize) -> Result<Self> {
Self::from_gguf_with_id(gguf, context_size, String::new())
}
pub fn from_gguf_with_id(
gguf: GgufFile,
context_size: usize,
model_id: String,
) -> Result<Self> {
ensure!(context_size > 0, "context_size must be > 0");
let arch = gguf
.get_str("general.architecture")
.context("missing general.architecture")?
.to_string();
let prefix = arch.as_str();
let rope_type = match prefix {
"qwen2" | "qwen3" => RopeType::Neox,
"llama" | "granite" => RopeType::Norm,
other => bail!(
"LlamaModel: no RoPE layout mapping for arch {other:?}; \
add it to the rope_type match in llama.rs"
),
};
let scalars = ScalarMultipliers::from_gguf(&gguf, prefix)?;
let n_layers =
gguf.get_u32(&format!("{prefix}.block_count"))
.with_context(|| format!("missing {prefix}.block_count"))? as usize;
let hidden_size = gguf
.get_u32(&format!("{prefix}.embedding_length"))
.with_context(|| format!("missing {prefix}.embedding_length"))?
as usize;
let intermediate_size = gguf
.get_u32(&format!("{prefix}.feed_forward_length"))
.with_context(|| format!("missing {prefix}.feed_forward_length"))?
as usize;
let n_heads = gguf
.get_u32(&format!("{prefix}.attention.head_count"))
.with_context(|| format!("missing {prefix}.attention.head_count"))?
as usize;
let n_kv_heads = gguf
.get_u32(&format!("{prefix}.attention.head_count_kv"))
.with_context(|| format!("missing {prefix}.attention.head_count_kv"))?
as usize;
ensure!(
n_heads > 0 && n_kv_heads > 0 && n_heads % n_kv_heads == 0,
"n_heads ({n_heads}) must be a positive multiple of n_kv_heads ({n_kv_heads})"
);
let vocab_size = match gguf.get_u32(&format!("{prefix}.vocab_size")) {
Some(v) => v as usize,
None => {
let info = gguf
.tensors
.get("token_embd.weight")
.context("missing token_embd.weight (cannot derive vocab_size)")?;
ensure!(
info.shape.len() >= 2,
"token_embd.weight has unexpected shape {:?}",
info.shape
);
info.shape[1]
}
};
let gguf_max_seq_len = gguf
.get_u32(&format!("{prefix}.context_length"))
.unwrap_or(128000) as usize;
let max_seq_len = context_size.min(gguf_max_seq_len);
let rope_theta = gguf
.get_f32(&format!("{prefix}.rope.freq_base"))
.unwrap_or(1_000_000.0);
let rms_norm_eps = gguf
.get_f32(&format!("{prefix}.attention.layer_norm_rms_epsilon"))
.unwrap_or(1e-6);
let head_dim = gguf
.get_u32(&format!("{prefix}.attention.key_length"))
.map(|v| v as usize)
.unwrap_or(hidden_size / n_heads);
ensure!(head_dim > 0, "head_dim must be > 0");
let block_types = vec![BlockType::Attention; n_layers];
let kv_heads_per_layer = vec![n_kv_heads; n_layers];
let config = ModelConfig {
architecture: arch.clone(),
n_layers,
hidden_size,
intermediate_size,
n_heads,
n_kv_heads,
head_dim,
vocab_size,
max_seq_len,
rope_theta,
rms_norm_eps,
block_types,
conv_kernel_size: None,
kv_heads_per_layer,
scalars,
};
let output_norm_weight = gguf.get_tensor("output_norm.weight")?.to_f32_vec();
let mut attn_norm_weights = Vec::with_capacity(n_layers);
let mut ffn_norm_weights = Vec::with_capacity(n_layers);
let mut attn_q_norm_weights = Vec::with_capacity(n_layers);
let mut attn_k_norm_weights = Vec::with_capacity(n_layers);
let mut attn_q_bias = Vec::with_capacity(n_layers);
let mut attn_k_bias = Vec::with_capacity(n_layers);
let mut attn_v_bias = Vec::with_capacity(n_layers);
let mut layer_refs = Vec::with_capacity(n_layers);
for i in 0..n_layers {
attn_norm_weights.push(
gguf.get_tensor(&format!("blk.{i}.attn_norm.weight"))?
.to_f32_vec(),
);
ffn_norm_weights.push(
gguf.get_tensor(&format!("blk.{i}.ffn_norm.weight"))?
.to_f32_vec(),
);
let q_norm_name = format!("blk.{i}.attn_q_norm.weight");
let k_norm_name = format!("blk.{i}.attn_k_norm.weight");
if gguf.tensors.contains_key(&q_norm_name) {
attn_q_norm_weights.push(Some(gguf.get_tensor(&q_norm_name)?.to_f32_vec()));
attn_k_norm_weights.push(Some(gguf.get_tensor(&k_norm_name)?.to_f32_vec()));
} else {
attn_q_norm_weights.push(None);
attn_k_norm_weights.push(None);
}
let q_bias_name = format!("blk.{i}.attn_q.bias");
let k_bias_name = format!("blk.{i}.attn_k.bias");
let v_bias_name = format!("blk.{i}.attn_v.bias");
if gguf.tensors.contains_key(&q_bias_name) {
attn_q_bias.push(Some(gguf.get_tensor(&q_bias_name)?.to_f32_vec()));
attn_k_bias.push(Some(gguf.get_tensor(&k_bias_name)?.to_f32_vec()));
attn_v_bias.push(Some(gguf.get_tensor(&v_bias_name)?.to_f32_vec()));
} else {
attn_q_bias.push(None);
attn_k_bias.push(None);
attn_v_bias.push(None);
}
layer_refs.push(LayerWeightRefs {
attn_q: transformer::resolve_weight(&gguf, &format!("blk.{i}.attn_q.weight"))?,
attn_k: transformer::resolve_weight(&gguf, &format!("blk.{i}.attn_k.weight"))?,
attn_v: transformer::resolve_weight(&gguf, &format!("blk.{i}.attn_v.weight"))?,
attn_output: transformer::resolve_weight(
&gguf,
&format!("blk.{i}.attn_output.weight"),
)?,
ffn_gate: transformer::resolve_weight(&gguf, &format!("blk.{i}.ffn_gate.weight"))?,
ffn_up: transformer::resolve_weight(&gguf, &format!("blk.{i}.ffn_up.weight"))?,
ffn_down: transformer::resolve_weight(&gguf, &format!("blk.{i}.ffn_down.weight"))?,
});
}
let embd_ref = transformer::resolve_weight(&gguf, "token_embd.weight")?;
let output_ref = if gguf.tensors.contains_key("output.weight") {
Some(transformer::resolve_weight(&gguf, "output.weight")?)
} else {
None
};
let rope_freqs = gguf
.get_tensor("rope_freqs.weight")
.ok()
.map(|t| t.to_f32_vec());
if let Some(rf) = &rope_freqs {
ensure!(
rf.len() == head_dim / 2,
"rope_freqs.weight has {} entries, expected head_dim/2 = {}",
rf.len(),
head_dim / 2
);
}
Ok(Self {
gguf,
config,
head_dim,
rope_type,
rope_freqs,
output_norm_weight,
attn_norm_weights,
ffn_norm_weights,
attn_q_norm_weights,
attn_k_norm_weights,
attn_q_bias,
attn_k_bias,
attn_v_bias,
embd_ref,
output_ref,
layer_refs,
model_id,
})
}
fn attn_dims(&self) -> AttnDims<'_> {
AttnDims {
hidden_size: self.config.hidden_size,
n_heads: self.config.n_heads,
n_kv_heads: self.config.n_kv_heads,
head_dim: self.head_dim,
rope_theta: self.config.rope_theta,
rms_norm_eps: self.config.rms_norm_eps,
rope_type: self.rope_type,
attn_scale: self.config.scalars.attn,
rope_freqs: self.rope_freqs.as_deref(),
}
}
fn run_layers(&self, hidden: &mut [f32], pos: usize, state: &mut InferenceState) {
let cfg = &self.config;
let hs = cfg.hidden_size;
let dims = self.attn_dims();
let mut normed = std::mem::take(&mut state.scratch.normed);
let mut ffn_input = std::mem::take(&mut state.scratch.ffn_input);
normed.resize(hs, 0.0);
ffn_input.resize(hs, 0.0);
for i in 0..cfg.n_layers {
normed.copy_from_slice(hidden);
cpu::rmsnorm(&mut normed, &self.attn_norm_weights[i], cfg.rms_norm_eps);
#[cfg(target_arch = "aarch64")]
transformer::quantize_to_scratch(&normed, state);
let refs = &self.layer_refs[i];
let weights = AttnWeights {
attn_q: &refs.attn_q,
attn_k: &refs.attn_k,
attn_v: &refs.attn_v,
attn_output: &refs.attn_output,
};
let extras = AttnExtras {
qkv_bias: match (
self.attn_q_bias[i].as_deref(),
self.attn_k_bias[i].as_deref(),
self.attn_v_bias[i].as_deref(),
) {
(Some(q), Some(k), Some(v)) => Some((q, k, v)),
_ => None,
},
qk_norm: match (
self.attn_q_norm_weights[i].as_deref(),
self.attn_k_norm_weights[i].as_deref(),
) {
(Some(q), Some(k)) => Some((q, k)),
_ => None,
},
};
transformer::forward_attn_block(
&self.gguf, i, &weights, &extras, dims, &normed, pos, state,
);
if self.config.scalars.residual != 1.0 {
cpu::scale_inplace(&mut state.scratch.out[..hs], self.config.scalars.residual);
}
cpu::add_inplace(hidden, &state.scratch.out[..hs]);
ffn_input.copy_from_slice(hidden);
cpu::rmsnorm(&mut ffn_input, &self.ffn_norm_weights[i], cfg.rms_norm_eps);
#[cfg(target_arch = "aarch64")]
transformer::quantize_to_scratch(&ffn_input, state);
let refs = &self.layer_refs[i];
let ffn_weights = FfnWeights {
ffn_gate: &refs.ffn_gate,
ffn_up: &refs.ffn_up,
ffn_down: &refs.ffn_down,
};
transformer::forward_ffn_block(
&self.gguf,
&ffn_weights,
hs,
cfg.intermediate_size,
&ffn_input,
state,
);
if self.config.scalars.residual != 1.0 {
cpu::scale_inplace(&mut state.scratch.out[..hs], self.config.scalars.residual);
}
cpu::add_inplace(hidden, &state.scratch.out[..hs]);
if transformer::oracle_dump::is_active() {
transformer::oracle_dump::record(&format!("l_out-{i}"), hidden);
}
}
cpu::rmsnorm(hidden, &self.output_norm_weight, cfg.rms_norm_eps);
transformer::oracle_dump::record("result_norm", hidden);
state.seq_len += 1;
state.scratch.normed = normed;
state.scratch.ffn_input = ffn_input;
}
fn project_logits(&self, hidden: &[f32], state: &mut InferenceState) -> Vec<f32> {
let cfg = &self.config;
let out_ref = self.output_ref.as_ref().unwrap_or(&self.embd_ref);
let mut logits = vec![0.0f32; cfg.vocab_size];
#[cfg(target_arch = "aarch64")]
{
transformer::quantize_to_scratch(hidden, state);
transformer::gemv_preq(
&self.gguf,
out_ref,
hidden,
&state.scratch.q8_scales,
&state.scratch.q8_quants,
&mut logits,
);
}
#[cfg(not(target_arch = "aarch64"))]
{
let _ = state;
transformer::gemv(&self.gguf, out_ref, hidden, &mut logits);
}
if self.config.scalars.logit != 1.0 {
cpu::scale_inplace(&mut logits, 1.0 / self.config.scalars.logit);
}
transformer::oracle_dump::record("result_output", &logits);
logits
}
}
impl Model for LlamaModel {
fn forward(&self, tokens: &[u32], pos: usize, state: &mut InferenceState) -> Vec<f32> {
assert_eq!(tokens.len(), 1, "LlamaModel forward expects single token");
let token_id = tokens[0] as usize;
let cfg = &self.config;
assert!(
token_id < cfg.vocab_size,
"token_id {token_id} out of range (vocab_size={})",
cfg.vocab_size
);
let mut hidden = transformer::dequantize_row(&self.gguf, &self.embd_ref, token_id);
if self.config.scalars.embedding != 1.0 {
cpu::scale_inplace(&mut hidden, self.config.scalars.embedding);
}
transformer::oracle_dump::record("embd", &hidden);
self.run_layers(&mut hidden, pos, state);
self.project_logits(&hidden, state)
}
fn forward_prefill(
&self,
tokens: &[u32],
start_pos: usize,
state: &mut InferenceState,
) -> Vec<f32> {
assert!(
!tokens.is_empty(),
"forward_prefill requires at least one token"
);
assert_eq!(
start_pos, state.seq_len,
"forward_prefill: start_pos ({start_pos}) must equal state.seq_len ({})",
state.seq_len
);
let mut logits = Vec::new();
for (i, &token) in tokens.iter().enumerate() {
logits = self.forward(&[token], start_pos + i, state);
}
logits
}
fn config(&self) -> &ModelConfig {
&self.config
}
fn supports_kv_shift(&self) -> bool {
true
}
fn shift_kv(&self, state: &mut InferenceState, n_keep: usize, shift: usize) {
state.shift_kv_with_rope(
n_keep,
shift,
self.config.rope_theta,
self.head_dim,
&self.config.kv_heads_per_layer,
self.rope_type,
self.rope_freqs.as_deref(),
);
}
}
#[cfg(feature = "gpu")]
impl crate::model::gpu_weight_source::GpuWeightSource for LlamaModel {
fn config(&self) -> &ModelConfig {
&self.config
}
fn gguf(&self) -> &GgufFile {
&self.gguf
}
fn output_norm_weight(&self) -> &[f32] {
&self.output_norm_weight
}
fn attn_norm_weight(&self, layer: usize) -> &[f32] {
&self.attn_norm_weights[layer]
}
fn ffn_norm_weight(&self, layer: usize) -> &[f32] {
&self.ffn_norm_weights[layer]
}
fn attn_q_norm_weight(&self, layer: usize) -> Option<&[f32]> {
self.attn_q_norm_weights[layer].as_deref()
}
fn attn_k_norm_weight(&self, layer: usize) -> Option<&[f32]> {
self.attn_k_norm_weights[layer].as_deref()
}
fn conv_weight(&self, _layer: usize) -> Option<&[f32]> {
None
}
fn attn_q_bias(&self, layer: usize) -> Option<&[f32]> {
self.attn_q_bias[layer].as_deref()
}
fn attn_k_bias(&self, layer: usize) -> Option<&[f32]> {
self.attn_k_bias[layer].as_deref()
}
fn attn_v_bias(&self, layer: usize) -> Option<&[f32]> {
self.attn_v_bias[layer].as_deref()
}
fn rope_freqs(&self) -> Option<&[f32]> {
self.rope_freqs.as_deref()
}
fn weight_bytes(&self, wref: &WeightRef) -> &[u8] {
transformer::weight_data(&self.gguf, wref)
}
fn dequantize_weight(&self, wref: &WeightRef) -> Vec<f32> {
transformer::dequantize_weight(&self.gguf, wref)
}
fn output_ref(&self) -> Option<&WeightRef> {
self.output_ref.as_ref()
}
fn ffn_gate_ref(&self, layer: usize) -> &WeightRef {
&self.layer_refs[layer].ffn_gate
}
fn ffn_up_ref(&self, layer: usize) -> &WeightRef {
&self.layer_refs[layer].ffn_up
}
fn ffn_down_ref(&self, layer: usize) -> &WeightRef {
&self.layer_refs[layer].ffn_down
}
fn conv_in_proj_ref(&self, _layer: usize) -> Option<&WeightRef> {
None
}
fn conv_out_proj_ref(&self, _layer: usize) -> Option<&WeightRef> {
None
}
fn attn_q_ref(&self, layer: usize) -> Option<&WeightRef> {
Some(&self.layer_refs[layer].attn_q)
}
fn attn_k_ref(&self, layer: usize) -> Option<&WeightRef> {
Some(&self.layer_refs[layer].attn_k)
}
fn attn_v_ref(&self, layer: usize) -> Option<&WeightRef> {
Some(&self.layer_refs[layer].attn_v)
}
fn attn_output_ref(&self, layer: usize) -> Option<&WeightRef> {
Some(&self.layer_refs[layer].attn_output)
}
fn rope_type(&self) -> RopeType {
self.rope_type
}
fn supports_batched_prefill(&self) -> bool {
false
}
}