use std::ops::Range;
use burn::tensor::{Device, Int, Tensor, backend::Backend};
use combs_formats::{ModelMetadata, ModelSource};
use crate::archspec::{ArchSpec, LayerKind, NormFlavor};
use crate::kv::{CacheConfig, CacheKind, ContiguousKVCache, KVCache, PagedKVCache};
use crate::matmul::safe_matmul;
use crate::norm::{gemma_rms_norm, rms_norm};
use crate::precision::{to_f32, to_float};
use crate::qlinear::{Linear, try_quant_linear};
use crate::rope::RotaryEmbedding;
use crate::traits::GenerativeModel;
use crate::{ModelError, Result};
struct LlamaLayer<B: Backend> {
q: Linear<B>,
k: Linear<B>,
v: Linear<B>,
o: Linear<B>,
q_bias: Option<Tensor<B, 1>>,
k_bias: Option<Tensor<B, 1>>,
v_bias: Option<Tensor<B, 1>>,
o_bias: Option<Tensor<B, 1>>,
q_norm: Option<Tensor<B, 1>>,
k_norm: Option<Tensor<B, 1>>,
gate: Linear<B>,
up: Linear<B>,
down: Linear<B>,
input_norm: Tensor<B, 1>,
pre_mlp_norm: Tensor<B, 1>,
attn_out_norm: Option<Tensor<B, 1>>,
mlp_out_norm: Option<Tensor<B, 1>>,
}
pub struct LlamaModel<B: Backend> {
metadata: ModelMetadata,
spec: ArchSpec,
embed: Tensor<B, 2>,
lm_head: Option<Linear<B>>, final_norm: Tensor<B, 1>,
layers: Vec<LlamaLayer<B>>,
rotary: RotaryEmbedding<B>,
rotary_local: Option<RotaryEmbedding<B>>,
scale: f64,
}
pub(crate) fn linear<B: Backend>(
x: Tensor<B, 3>,
w: &Tensor<B, 2>,
bias: Option<&Tensor<B, 1>>,
) -> Tensor<B, 3> {
let out = safe_matmul(x, w.clone().transpose().unsqueeze_dim::<3>(0));
match bias {
Some(b) => {
let [batch, seq, dim] = out.dims();
out + b.clone().reshape([1, 1, dim]).expand([batch, seq, dim])
}
None => out,
}
}
fn load_weight<B: Backend, const D: usize>(
source: &dyn ModelSource,
device: &Device<B>,
name: &str,
) -> Result<Tensor<B, D>> {
source
.open_tensor(name)
.map_err(|e| match e {
combs_formats::FormatError::TensorNotFound(_) => {
ModelError::MissingTensor(name.to_string())
}
other => ModelError::Format(other),
})?
.load_to_tensor::<B, D>(device)
.map_err(ModelError::Format)
}
pub(crate) fn load_tensor<B: Backend, const D: usize>(
source: &dyn ModelSource,
device: &Device<B>,
name: &str,
) -> Result<Tensor<B, D>> {
load_weight(source, device, name)
}
pub(crate) fn load_linear<B: Backend>(
source: &dyn ModelSource,
device: &Device<B>,
name: &str,
) -> Result<Linear<B>> {
if let Some(op) = try_quant_linear::<B>(source, name, device)? {
return Ok(Linear::Quant(op));
}
Ok(Linear::Dense(load_weight(source, device, name)?))
}
fn load_optional_bias<B: Backend>(
source: &dyn ModelSource,
device: &Device<B>,
name: &str,
) -> Result<Option<Tensor<B, 1>>> {
match source.open_tensor(name) {
Ok(reader) => Ok(Some(
reader.load_to_tensor::<B, 1>(device).map_err(ModelError::Format)?,
)),
Err(combs_formats::FormatError::TensorNotFound(_)) => Ok(None),
Err(e) => Err(ModelError::Format(e)),
}
}
fn split_fused_rows<B: Backend, const N: usize>(
source: &dyn ModelSource,
device: &Device<B>,
name: &str,
rows: [usize; N],
) -> Result<[Linear<B>; N]> {
let w: Tensor<B, 2> = load_weight(source, device, name)?;
let [total, cols] = w.dims();
let expect: usize = rows.iter().sum();
if total != expect {
return Err(ModelError::BadShape {
tensor: name.to_string(),
expected: vec![expect, cols],
got: vec![total, cols],
});
}
let mut at = 0;
Ok(rows.map(|r| {
let part = w.clone().narrow(0, at, r);
at += r;
Linear::Dense(part)
}))
}
fn split_fused_bias<B: Backend, const N: usize>(
source: &dyn ModelSource,
device: &Device<B>,
name: &str,
rows: [usize; N],
) -> Result<Option<[Tensor<B, 1>; N]>> {
let Some(b) = load_optional_bias(source, device, name)? else {
return Ok(None);
};
let [total] = b.dims();
let expect: usize = rows.iter().sum();
if total != expect {
return Err(ModelError::BadShape {
tensor: name.to_string(),
expected: vec![expect],
got: vec![total],
});
}
let mut at = 0;
Ok(Some(rows.map(|r| {
let part = b.clone().narrow(0, at, r);
at += r;
part
})))
}
impl<B: Backend> LlamaModel<B> {
pub(crate) fn expect_shape(name: &str, got: &[usize], expected: &[usize]) -> Result<()> {
if got == expected {
Ok(())
} else {
Err(ModelError::BadShape {
tensor: name.to_string(),
expected: expected.to_vec(),
got: got.to_vec(),
})
}
}
fn norm<const D: usize>(&self, x: Tensor<B, D>, w: &Tensor<B, 1>) -> Tensor<B, D> {
match self.spec.norm_flavor {
NormFlavor::RmsNorm => rms_norm(x, w.clone(), self.metadata.rms_norm_eps),
NormFlavor::GemmaRmsNorm => {
gemma_rms_norm(x, w.clone(), self.metadata.rms_norm_eps)
}
}
}
fn rotary_for(&self, layer_idx: usize) -> &RotaryEmbedding<B> {
match (self.spec.layers.get(layer_idx), &self.rotary_local) {
(Some(LayerKind::Sliding(_)), Some(local)) => local,
_ => &self.rotary,
}
}
pub(crate) fn forward_hidden(
&self,
mut x: Tensor<B, 3>,
cache: &mut dyn KVCache<B>,
pos: usize,
) -> Tensor<B, 3> {
let m = &self.metadata;
let [_, seq, _] = x.dims();
for (layer_idx, layer) in self.layers.iter().enumerate() {
let window = match self.spec.layers.get(layer_idx) {
Some(LayerKind::Sliding(w)) => Some(*w),
_ => None,
};
let h = self.norm(x.clone(), &layer.input_norm);
let q = layer.q.forward(h.clone(), layer.q_bias.as_ref());
let k = layer.k.forward(h.clone(), layer.k_bias.as_ref());
let v = layer.v.forward(h, layer.v_bias.as_ref());
let mut q = q
.reshape([1, seq, m.num_attention_heads, m.head_dim])
.swap_dims(1, 2);
let mut k = k
.reshape([1, seq, m.num_key_value_heads, m.head_dim])
.swap_dims(1, 2);
let v = v
.reshape([1, seq, m.num_key_value_heads, m.head_dim])
.swap_dims(1, 2);
if let Some(qn) = &layer.q_norm {
q = self.norm(q, qn);
}
if let Some(kn) = &layer.k_norm {
k = self.norm(k, kn);
}
let rotary = self.rotary_for(layer_idx);
let q = rotary.apply(q, pos);
let k = rotary.apply(k, pos);
let ctx = cache.attention_opts(layer_idx, q, k, v, pos, self.scale, window);
let ctx = ctx
.swap_dims(1, 2)
.reshape([1, seq, m.num_attention_heads * m.head_dim]);
let mut attn_out = layer.o.forward(ctx, layer.o_bias.as_ref());
if let Some(n) = &layer.attn_out_norm {
attn_out = self.norm(attn_out, n);
}
x = x + attn_out;
let h = self.norm(x.clone(), &layer.pre_mlp_norm);
let gated = crate::act::apply(self.spec.activation, layer.gate.forward(h.clone(), None))
* layer.up.forward(h.clone(), None);
let mut mlp_out = layer.down.forward(gated, None);
if let Some(n) = &layer.mlp_out_norm {
mlp_out = self.norm(mlp_out, n);
}
x = x + mlp_out;
}
self.norm(x, &self.final_norm)
}
pub(crate) fn all_logits(&self, hidden: Tensor<B, 3>) -> Tensor<B, 3> {
let [_, seq, hidden_size] = hidden.dims();
let logits: Tensor<B, 3> = match &self.lm_head {
Some(head) => head.forward(hidden, None),
None => {
let flat = hidden.reshape([seq, hidden_size]);
let out = safe_matmul(flat, self.embed.clone().transpose());
let [_, vocab] = out.dims();
out.reshape([1, seq, vocab])
}
};
match self.spec.final_logit_softcap {
Some(cap) => logits.div_scalar(cap as f32).tanh().mul_scalar(cap as f32),
None => logits,
}
}
pub(crate) fn last_logits(&self, hidden: Tensor<B, 3>) -> Tensor<B, 2> {
let [_, seq, hidden_size] = hidden.dims();
let last = hidden.narrow(1, seq - 1, 1); let logits: Tensor<B, 2> = match &self.lm_head {
Some(head) => {
let logits = head.forward(last, None);
let [_, _, vocab] = logits.dims();
logits.reshape([1, vocab])
}
None => {
let last = last.reshape([1, hidden_size]);
safe_matmul(last, self.embed.clone().transpose())
}
};
match self.spec.final_logit_softcap {
Some(cap) => logits.div_scalar(cap as f32).tanh().mul_scalar(cap as f32),
None => logits,
}
}
}
impl<B: Backend> GenerativeModel<B> for LlamaModel<B> {
fn metadata(&self) -> &ModelMetadata {
&self.metadata
}
fn load(source: &dyn ModelSource, device: &Device<B>) -> Result<Self> {
let prefix = if source
.tensor_names()
.iter()
.any(|n| n == "model.embed_tokens.weight" || n == "model.embed_tokens")
{
"model"
} else {
""
};
Self::load_with_prefix(source, device, prefix)
}
fn create_kv_cache(&self, config: &CacheConfig) -> Box<dyn KVCache<B>> {
match config.kind {
CacheKind::Contiguous => {
Box::new(ContiguousKVCache::<B>::new(self.metadata.num_hidden_layers))
}
CacheKind::Paged => Box::new(PagedKVCache::<B>::new_with_windows(
self.metadata.num_hidden_layers,
*config,
self.spec.windows(),
)),
}
}
fn embed(&self, tokens: Tensor<B, 2, Int>) -> Tensor<B, 3> {
let [batch, seq] = tokens.dims();
let flat = tokens.reshape([batch * seq]);
let embedded = self
.embed
.clone()
.select(0, flat)
.reshape([batch, seq, self.metadata.hidden_size]);
if self.spec.embed_scale_sqrt_hidden {
let out_dtype = embedded.dtype();
let scale = (self.metadata.hidden_size as f64).sqrt();
to_float(to_f32(embedded).mul_scalar(scale), out_dtype)
} else {
embedded
}
}
fn prefill(
&mut self,
input: Tensor<B, 3>,
cache: &mut dyn KVCache<B>,
pos: Range<u32>,
) -> Tensor<B, 2> {
let [_, seq, _] = input.dims();
assert_eq!(
seq,
(pos.end - pos.start) as usize,
"prefill pos range must match the input sequence length"
);
let hidden = self.forward_hidden(input, cache, pos.start as usize);
self.last_logits(hidden)
}
fn prefill_hidden(
&mut self,
input: Tensor<B, 3>,
cache: &mut dyn KVCache<B>,
pos: Range<u32>,
) -> Result<Tensor<B, 3>> {
let [_, seq, _] = input.dims();
assert_eq!(
seq,
(pos.end - pos.start) as usize,
"prefill pos range must match the input sequence length"
);
Ok(self.forward_hidden(input, cache, pos.start as usize))
}
fn supports_hidden_states(&self) -> bool {
true
}
fn prefill_all_logits(
&mut self,
input: Tensor<B, 3>,
cache: &mut dyn KVCache<B>,
pos: Range<u32>,
) -> Result<Tensor<B, 3>> {
let hidden = self.prefill_hidden(input, cache, pos)?;
Ok(self.all_logits(hidden))
}
fn decode(&mut self, input: Tensor<B, 3>, cache: &mut dyn KVCache<B>) -> Tensor<B, 2> {
let pos = cache.seq_len();
let hidden = self.forward_hidden(input, cache, pos);
self.last_logits(hidden)
}
fn decode_all_logits(
&mut self,
input: Tensor<B, 3>,
cache: &mut dyn KVCache<B>,
) -> Result<Tensor<B, 3>> {
let pos = cache.seq_len();
let hidden = self.forward_hidden(input, cache, pos);
Ok(self.all_logits(hidden))
}
fn supports_decode_all_logits(&self) -> bool {
true
}
}
impl<B: Backend> LlamaModel<B> {
pub(crate) fn load_with_prefix(
source: &dyn ModelSource,
device: &Device<B>,
prefix: &str,
) -> Result<Self> {
let m = source.metadata().clone();
let spec = ArchSpec::resolve(&m);
let prefix = if prefix.is_empty() {
String::new()
} else {
format!("{prefix}.")
};
let prefix = prefix.as_str();
let embed: Tensor<B, 2> =
load_weight(source, device, &format!("{prefix}embed_tokens.weight"))?;
Self::expect_shape(
"embed_tokens.weight",
&embed.dims(),
&[m.vocab_size, m.hidden_size],
)?;
let lm_head = if m.tie_word_embeddings {
None
} else {
match load_linear(source, device, "lm_head.weight") {
Ok(w) => {
Self::expect_shape(
"lm_head.weight",
&w.dims(),
&[m.vocab_size, m.hidden_size],
)?;
Some(w)
}
Err(ModelError::MissingTensor(_)) => {
eprintln!(
"[load] lm_head.weight absent; falling back to tied embeddings"
);
None
}
Err(e) => return Err(e),
}
};
let final_norm: Tensor<B, 1> =
load_weight(source, device, &format!("{prefix}norm.weight"))?;
let pre_mlp_name = if spec.sandwich_norms {
"pre_feedforward_layernorm"
} else {
"post_attention_layernorm"
};
let q_rows = m.num_attention_heads * m.head_dim;
let kv_rows = m.num_key_value_heads * m.head_dim;
let mut layers = Vec::with_capacity(m.num_hidden_layers);
for i in 0..m.num_hidden_layers {
let p = format!("{prefix}layers.{i}");
let (q, k, v, fused_qkv_bias) =
match load_linear(source, device, &format!("{p}.self_attn.q_proj.weight")) {
Ok(q) => (
q,
load_linear(source, device, &format!("{p}.self_attn.k_proj.weight"))?,
load_linear(source, device, &format!("{p}.self_attn.v_proj.weight"))?,
None,
),
Err(ModelError::MissingTensor(_)) => {
let name = format!("{p}.self_attn.qkv_proj");
let [q, k, v] = split_fused_rows(
source,
device,
&format!("{name}.weight"),
[q_rows, kv_rows, kv_rows],
)?;
let b = split_fused_bias(
source,
device,
&format!("{name}.bias"),
[q_rows, kv_rows, kv_rows],
)?;
(q, k, v, b)
}
Err(e) => return Err(e),
};
let o = load_linear(source, device, &format!("{p}.self_attn.o_proj.weight"))?;
Self::expect_shape(
&format!("{p}.self_attn.q_proj.weight"),
&q.dims(),
&[q_rows, m.hidden_size],
)?;
Self::expect_shape(
&format!("{p}.self_attn.k_proj.weight"),
&k.dims(),
&[kv_rows, m.hidden_size],
)?;
let (gate, up) =
match load_linear(source, device, &format!("{p}.mlp.gate_proj.weight")) {
Ok(gate) => (
gate,
load_linear(source, device, &format!("{p}.mlp.up_proj.weight"))?,
),
Err(ModelError::MissingTensor(_)) => {
let [gate, up] = split_fused_rows(
source,
device,
&format!("{p}.mlp.gate_up_proj.weight"),
[m.intermediate_size, m.intermediate_size],
)?;
(gate, up)
}
Err(e) => return Err(e),
};
let bias = |proj: &str| -> Result<Option<Tensor<B, 1>>> {
load_optional_bias(source, device, &format!("{p}.{proj}.bias"))
};
let (q_bias, k_bias, v_bias) = match fused_qkv_bias {
Some([qb, kb, vb]) => (Some(qb), Some(kb), Some(vb)),
None => (
bias("self_attn.q_proj")?,
bias("self_attn.k_proj")?,
bias("self_attn.v_proj")?,
),
};
let optional_norm = |name: &str| -> Result<Option<Tensor<B, 1>>> {
match source.open_tensor(&format!("{p}.{name}.weight")) {
Ok(reader) => Ok(Some(
reader.load_to_tensor::<B, 1>(device).map_err(ModelError::Format)?,
)),
Err(combs_formats::FormatError::TensorNotFound(_)) => Ok(None),
Err(e) => Err(ModelError::Format(e)),
}
};
layers.push(LlamaLayer {
q,
k,
v,
o,
q_bias,
k_bias,
v_bias,
o_bias: bias("self_attn.o_proj")?,
q_norm: if spec.qk_norm { optional_norm("self_attn.q_norm")? } else { None },
k_norm: if spec.qk_norm { optional_norm("self_attn.k_norm")? } else { None },
gate,
up,
down: load_linear(source, device, &format!("{p}.mlp.down_proj.weight"))?,
input_norm: load_weight(source, device, &format!("{p}.input_layernorm.weight"))?,
pre_mlp_norm: load_weight(
source,
device,
&format!("{p}.{pre_mlp_name}.weight"),
)?,
attn_out_norm: if spec.sandwich_norms {
optional_norm("post_attention_layernorm")?
} else {
None
},
mlp_out_norm: if spec.sandwich_norms {
optional_norm("post_feedforward_layernorm")?
} else {
None
},
});
}
let rotary = RotaryEmbedding::new_scaled(
m.head_dim,
spec.rope_theta,
m.max_position_embeddings,
&spec.rope_scaling,
device,
);
let rotary_local = spec.rope_local_theta.map(|theta| {
RotaryEmbedding::new(m.head_dim, theta, m.max_position_embeddings, device)
});
Ok(LlamaModel {
scale: 1.0
/ spec
.query_pre_attn_scalar
.unwrap_or(m.head_dim as f64)
.sqrt(),
metadata: m,
spec,
embed,
lm_head,
final_norm,
layers,
rotary,
rotary_local,
})
}
}