1use crate::quantized_nn::RmsNorm;
2use crate::utils::repeat_kv;
3use candle::quantized::gguf_file;
4use candle::quantized::QMatMul;
5use candle::{bail, DType, Device, IndexOp, Result, Tensor};
6use candle_nn::{Conv1d, Conv1dConfig, Embedding, Module};
7use std::collections::HashMap;
8
9fn get_qtensor<R: std::io::Seek + std::io::Read>(
10 ct: &gguf_file::Content,
11 reader: &mut R,
12 device: &Device,
13 names: &[String],
14) -> Result<candle::quantized::QTensor> {
15 for name in names {
16 if let Ok(t) = ct.tensor(reader, name, device) {
17 return Ok(t);
18 }
19 }
20 bail!("cannot find tensor info for {}", names.join(" | "))
21}
22
23fn get_dequantized<R: std::io::Seek + std::io::Read>(
24 ct: &gguf_file::Content,
25 reader: &mut R,
26 device: &Device,
27 names: &[String],
28) -> Result<Tensor> {
29 get_qtensor(ct, reader, device, names)?.dequantize(device)
30}
31
32#[derive(Debug, Clone)]
33struct Mlp {
34 w1: QMatMul,
35 w2: QMatMul,
36 w3: QMatMul,
37}
38
39impl Module for Mlp {
40 fn forward(&self, xs: &Tensor) -> Result<Tensor> {
41 let w1 = self.w1.forward(xs)?;
42 let w3 = self.w3.forward(xs)?;
43 self.w2.forward(&(candle_nn::ops::silu(&w1)? * w3)?)
44 }
45}
46
47#[derive(Debug, Clone)]
48struct AttentionLayer {
49 wq: QMatMul,
50 wk: QMatMul,
51 wv: QMatMul,
52 wo: QMatMul,
53 q_norm: RmsNorm,
54 k_norm: RmsNorm,
55 n_head: usize,
56 n_kv_head: usize,
57 head_dim: usize,
58 cos: Tensor,
59 sin: Tensor,
60 neg_inf: Tensor,
61 kv_cache: Option<(Tensor, Tensor)>,
62 span_attn: tracing::Span,
63 span_rot: tracing::Span,
64}
65
66#[derive(Debug, Clone)]
67struct ShortConvLayer {
68 in_proj: QMatMul,
69 out_proj: QMatMul,
70 conv: Tensor,
71 l_cache: usize,
72 cache: Option<Tensor>,
73}
74
75#[allow(clippy::large_enum_variant)]
76#[derive(Debug, Clone)]
77enum LayerKind {
78 Attention(AttentionLayer),
79 ShortConv(ShortConvLayer),
80}
81
82#[derive(Debug, Clone)]
83struct LayerWeights {
84 operator_norm: RmsNorm,
85 ffn_norm: RmsNorm,
86 mlp: Mlp,
87 kind: LayerKind,
88 span_mlp: tracing::Span,
89}
90
91fn masked_fill(on_false: &Tensor, mask: &Tensor, on_true: &Tensor) -> Result<Tensor> {
92 let shape = mask.shape();
93 let m = mask.where_cond(&on_true.broadcast_as(shape.dims())?, on_false)?;
94 Ok(m)
95}
96
97fn precomput_freqs_cis(
98 head_dim: usize,
99 freq_base: f32,
100 context_length: usize,
101 device: &Device,
102) -> Result<(Tensor, Tensor)> {
103 let theta: Vec<_> = (0..head_dim)
104 .step_by(2)
105 .map(|i| 1f32 / freq_base.powf(i as f32 / head_dim as f32))
106 .collect();
107 let theta = Tensor::new(theta.as_slice(), device)?;
108 let idx_theta = Tensor::arange(0, context_length as u32, device)?
109 .to_dtype(DType::F32)?
110 .reshape((context_length, 1))?
111 .matmul(&theta.reshape((1, theta.elem_count()))?)?;
112 let cos = idx_theta.cos()?;
113 let sin = idx_theta.sin()?;
114 Ok((cos, sin))
115}
116
117impl AttentionLayer {
118 fn apply_rotary_emb(&self, x: &Tensor, index_pos: usize) -> Result<Tensor> {
119 let _enter = self.span_rot.enter();
120 let (_b, _n, seq_len, _d) = x.dims4()?;
121 let cos = self.cos.narrow(0, index_pos, seq_len)?;
122 let sin = self.sin.narrow(0, index_pos, seq_len)?;
123 candle_nn::rotary_emb::rope(&x.contiguous()?, &cos, &sin)
124 }
125
126 fn forward(&mut self, xs: &Tensor, mask: Option<&Tensor>, index_pos: usize) -> Result<Tensor> {
127 let _enter = self.span_attn.enter();
128 let (b_sz, seq_len, n_embd) = xs.dims3()?;
129
130 let q = self.wq.forward(xs)?;
131 let k = self.wk.forward(xs)?;
132 let v = self.wv.forward(xs)?;
133
134 let q = q
135 .reshape((b_sz, seq_len, self.n_head, self.head_dim))?
136 .transpose(1, 2)?;
137 let k = k
138 .reshape((b_sz, seq_len, self.n_kv_head, self.head_dim))?
139 .transpose(1, 2)?;
140 let v = v
141 .reshape((b_sz, seq_len, self.n_kv_head, self.head_dim))?
142 .transpose(1, 2)?
143 .contiguous()?;
144
145 let q = self.q_norm.forward(&q.contiguous()?)?;
146 let k = self.k_norm.forward(&k.contiguous()?)?;
147
148 let q = self.apply_rotary_emb(&q, index_pos)?;
149 let k = self.apply_rotary_emb(&k, index_pos)?;
150
151 let (k, v) = match &self.kv_cache {
152 None => (k, v),
153 Some((k_cache, v_cache)) => {
154 if index_pos == 0 {
155 (k, v)
156 } else {
157 let k = Tensor::cat(&[k_cache, &k], 2)?;
158 let v = Tensor::cat(&[v_cache, &v], 2)?;
159 (k, v)
160 }
161 }
162 };
163 self.kv_cache = Some((k.clone(), v.clone()));
164
165 let k = repeat_kv(k, self.n_head / self.n_kv_head)?;
166 let v = repeat_kv(v, self.n_head / self.n_kv_head)?;
167
168 let att = (q.matmul(&k.t()?)? / (self.head_dim as f64).sqrt())?;
169 let att = match mask {
170 None => att,
171 Some(mask) => {
172 let mask = mask.broadcast_as(att.shape())?;
173 masked_fill(&att, &mask, &self.neg_inf)?
174 }
175 };
176 let att = candle_nn::ops::softmax_last_dim(&att)?;
177 let y = att.matmul(&v.contiguous()?)?;
178
179 let y = y.transpose(1, 2)?.reshape(&[b_sz, seq_len, n_embd])?;
180 self.wo.forward(&y)
181 }
182}
183
184impl ShortConvLayer {
185 fn forward(&mut self, xs: &Tensor, _index_pos: usize) -> Result<Tensor> {
186 let (b_sz, seq_len, hidden) = xs.dims3()?;
187 let bcx = self.in_proj.forward(xs)?.transpose(1, 2)?;
188 let b = bcx.narrow(1, 0, hidden)?;
189 let c = bcx.narrow(1, hidden, hidden)?;
190 let x = bcx.narrow(1, 2 * hidden, hidden)?;
191 let bx = (b * &x)?.contiguous()?;
192
193 let mut conv_weight = self.conv.clone();
195 if conv_weight.dims().len() == 3 {
196 conv_weight = conv_weight.squeeze(1)?;
197 } else if conv_weight.dims().len() == 2 && conv_weight.dims2()? == (self.l_cache, hidden) {
198 conv_weight = conv_weight.t()?.contiguous()?;
199 }
200 let conv_weight = conv_weight.contiguous()?;
201
202 let mut conv_out = if seq_len == 1 {
203 let mut state = if let Some(cache) = &self.cache {
204 cache.clone()
205 } else {
206 Tensor::zeros((b_sz, hidden, self.l_cache), bx.dtype(), bx.device())?
207 };
208
209 if self.l_cache > 1 {
210 let tail = state.narrow(2, 1, self.l_cache - 1)?;
211 state = Tensor::cat(&[tail, bx.clone()], 2)?;
212 } else {
213 state = bx.clone();
214 }
215 self.cache = Some(state.clone());
216
217 (state * &conv_weight.unsqueeze(0)?)?
218 .sum_keepdim(2)?
219 .contiguous()?
220 } else {
221 let conv = Conv1d::new(
222 conv_weight
223 .reshape((hidden, 1, self.l_cache))?
224 .contiguous()?,
225 None,
226 Conv1dConfig {
227 padding: self.l_cache.saturating_sub(1),
228 groups: hidden,
229 ..Default::default()
230 },
231 );
232 let mut out = conv.forward(&bx.contiguous()?)?;
233 out = out.narrow(2, 0, seq_len)?;
234
235 if self.l_cache > 0 {
236 let (_, _, cur_len) = bx.dims3()?;
237 let start = cur_len.saturating_sub(self.l_cache);
238 let mut cache_src = bx.narrow(2, start, cur_len - start)?;
239 if cache_src.dims3()?.2 < self.l_cache {
240 let pad = self.l_cache - cache_src.dims3()?.2;
241 let zeros =
242 Tensor::zeros((b_sz, hidden, pad), cache_src.dtype(), cache_src.device())?;
243 cache_src = Tensor::cat(&[zeros, cache_src], 2)?;
244 }
245 self.cache = Some(cache_src);
246 }
247
248 out
249 };
250
251 conv_out = (c * &conv_out)?;
252 let conv_out = conv_out.transpose(1, 2)?.contiguous()?;
253 self.out_proj.forward(&conv_out)
254 }
255}
256
257pub struct ModelWeights {
258 tok_embeddings: Embedding,
259 layers: Vec<LayerWeights>,
260 norm: RmsNorm,
261 output: QMatMul,
262 masks: HashMap<(usize, usize), Tensor>,
263 span: tracing::Span,
264 span_output: tracing::Span,
265}
266
267fn value_to_usize(v: &gguf_file::Value) -> Result<usize> {
268 use gguf_file::Value::*;
269 match v {
270 U8(x) => Ok(*x as usize),
271 I8(x) => Ok(*x as usize),
272 U16(x) => Ok(*x as usize),
273 I16(x) => Ok(*x as usize),
274 U32(x) => Ok(*x as usize),
275 I32(x) => Ok(*x as usize),
276 U64(x) => Ok(*x as usize),
277 I64(x) => Ok(*x as usize),
278 F32(x) => Ok(*x as usize),
279 F64(x) => Ok(*x as usize),
280 Bool(x) => Ok(usize::from(*x)),
281 String(_) => bail!("unexpected string metadata"),
282 Array(_) => bail!("array should be handled separately"),
283 }
284}
285
286fn read_usize_list(v: &gguf_file::Value, len: usize) -> Result<Vec<usize>> {
287 use gguf_file::Value::Array;
288 match v {
289 Array(arr) => {
290 let mut out = Vec::with_capacity(arr.len());
291 for item in arr {
292 out.push(value_to_usize(item)?);
293 }
294 if out.len() == len {
295 Ok(out)
296 } else if out.len() == 1 {
297 Ok(vec![out[0]; len])
298 } else {
299 bail!(
300 "unexpected array length in metadata, expected {len} got {}",
301 out.len()
302 )
303 }
304 }
305 _ => Ok(vec![value_to_usize(v)?; len]),
306 }
307}
308
309impl ModelWeights {
310 pub fn from_gguf<R: std::io::Seek + std::io::Read>(
311 ct: gguf_file::Content,
312 reader: &mut R,
313 device: &Device,
314 ) -> Result<Self> {
315 let md_get = |s: &str| match ct.metadata.get(s) {
316 None => bail!("cannot find {s} in metadata"),
317 Some(v) => Ok(v),
318 };
319
320 let head_count = md_get("lfm2.attention.head_count")?.to_u32()? as usize;
321 let head_count_kv_meta = md_get("lfm2.attention.head_count_kv")?;
322 let embedding_length = md_get("lfm2.embedding_length")?.to_u32()? as usize;
323 let context_length = md_get("lfm2.context_length")?.to_u32()? as usize;
324 let block_count = md_get("lfm2.block_count")?.to_u32()? as usize;
325 let rms_norm_eps = md_get("lfm2.attention.layer_norm_rms_epsilon")?.to_f32()? as f64;
326 let rope_freq_base = md_get("lfm2.rope.freq_base")
327 .and_then(|m| m.to_f32())
328 .unwrap_or(1_000_000f32);
329 let l_cache = md_get("lfm2.shortconv.l_cache")?.to_u32()? as usize;
330
331 let head_count_kv = read_usize_list(head_count_kv_meta, block_count)?;
332 let head_dim = embedding_length / head_count;
333 let (cos, sin) = precomput_freqs_cis(head_dim, rope_freq_base, context_length, device)?;
334 let neg_inf = Tensor::new(f32::NEG_INFINITY, device)?;
335
336 let tok_embeddings_q = get_qtensor(
337 &ct,
338 reader,
339 device,
340 &[
341 "token_embd.weight",
342 "tok_embeddings.weight",
343 "model.embed_tokens.weight",
344 ]
345 .iter()
346 .map(|s| s.to_string())
347 .collect::<Vec<_>>(),
348 )?;
349 let tok_embeddings = tok_embeddings_q.dequantize(device)?;
350 tracing::debug!(
351 tok_embd_shape = ?tok_embeddings.shape().dims(),
352 "loaded lfm2 token embeddings"
353 );
354
355 let norm = RmsNorm::from_qtensor(
356 get_qtensor(
357 &ct,
358 reader,
359 device,
360 &[
361 "output_norm.weight",
362 "embedding_norm.weight",
363 "model.embedding_norm.weight",
364 "model.embedding_norm",
365 "token_embd_norm.weight",
366 ]
367 .iter()
368 .map(|s| s.to_string())
369 .collect::<Vec<_>>(),
370 )?,
371 rms_norm_eps,
372 )?;
373 let output_q = get_qtensor(
374 &ct,
375 reader,
376 device,
377 &[
378 "output.weight",
379 "lm_head.weight",
380 "model.output.weight",
381 "model.lm_head.weight",
382 ]
383 .iter()
384 .map(|s| s.to_string())
385 .collect::<Vec<_>>(),
386 )
387 .unwrap_or(tok_embeddings_q);
388 tracing::debug!(
389 output_shape = ?output_q.shape().dims(),
390 "loaded lfm2 output weight (using tok_embd if missing)"
391 );
392
393 let mut layers = Vec::with_capacity(block_count);
394 for layer_idx in 0..block_count {
395 let prefix = format!("blk.{layer_idx}");
396 let is_attention = head_count_kv.get(layer_idx).copied().unwrap_or(head_count) > 0;
397
398 let operator_norm = get_qtensor(
399 &ct,
400 reader,
401 device,
402 &[
403 format!("{prefix}.attn_norm.weight"),
404 format!("{prefix}.operator_norm.weight"),
405 format!("{prefix}.attention_norm.weight"),
406 ],
407 )?;
408 let ffn_norm = get_qtensor(
409 &ct,
410 reader,
411 device,
412 &[
413 format!("{prefix}.ffn_norm.weight"),
414 format!("{prefix}.ffn_norm"),
415 ],
416 )?;
417 let mlp = {
418 let w1 = get_qtensor(
419 &ct,
420 reader,
421 device,
422 &[
423 format!("{prefix}.ffn_gate.weight"),
424 format!("{prefix}.feed_forward.w1.weight"),
425 format!("{prefix}.mlp.gate_proj.weight"),
426 ],
427 )?;
428 let w2 = get_qtensor(
429 &ct,
430 reader,
431 device,
432 &[
433 format!("{prefix}.ffn_down.weight"),
434 format!("{prefix}.feed_forward.w2.weight"),
435 format!("{prefix}.mlp.down_proj.weight"),
436 ],
437 )?;
438 let w3 = get_qtensor(
439 &ct,
440 reader,
441 device,
442 &[
443 format!("{prefix}.ffn_up.weight"),
444 format!("{prefix}.feed_forward.w3.weight"),
445 format!("{prefix}.mlp.up_proj.weight"),
446 ],
447 )?;
448 Mlp {
449 w1: QMatMul::from_qtensor(w1)?,
450 w2: QMatMul::from_qtensor(w2)?,
451 w3: QMatMul::from_qtensor(w3)?,
452 }
453 };
454
455 let kind = if is_attention {
456 let n_kv_head = head_count_kv[layer_idx];
457 let wq = get_qtensor(
458 &ct,
459 reader,
460 device,
461 &[
462 format!("{prefix}.attn_q.weight"),
463 format!("{prefix}.self_attn.q_proj.weight"),
464 ],
465 )?;
466 let wk = get_qtensor(
467 &ct,
468 reader,
469 device,
470 &[
471 format!("{prefix}.attn_k.weight"),
472 format!("{prefix}.self_attn.k_proj.weight"),
473 ],
474 )?;
475 let wv = get_qtensor(
476 &ct,
477 reader,
478 device,
479 &[
480 format!("{prefix}.attn_v.weight"),
481 format!("{prefix}.self_attn.v_proj.weight"),
482 ],
483 )?;
484 let wo = get_qtensor(
485 &ct,
486 reader,
487 device,
488 &[
489 format!("{prefix}.attn_output.weight"),
490 format!("{prefix}.self_attn.out_proj.weight"),
491 ],
492 )?;
493 let q_norm = get_qtensor(
494 &ct,
495 reader,
496 device,
497 &[
498 format!("{prefix}.attn_q_norm.weight"),
499 format!("{prefix}.self_attn.q_layernorm.weight"),
500 format!("{prefix}.attention.q_norm.weight"),
501 ],
502 )?;
503 let k_norm = get_qtensor(
504 &ct,
505 reader,
506 device,
507 &[
508 format!("{prefix}.attn_k_norm.weight"),
509 format!("{prefix}.self_attn.k_layernorm.weight"),
510 format!("{prefix}.attention.k_norm.weight"),
511 ],
512 )?;
513
514 LayerKind::Attention(AttentionLayer {
515 wq: QMatMul::from_qtensor(wq)?,
516 wk: QMatMul::from_qtensor(wk)?,
517 wv: QMatMul::from_qtensor(wv)?,
518 wo: QMatMul::from_qtensor(wo)?,
519 q_norm: RmsNorm::from_qtensor(q_norm, rms_norm_eps)?,
520 k_norm: RmsNorm::from_qtensor(k_norm, rms_norm_eps)?,
521 n_head: head_count,
522 n_kv_head,
523 head_dim,
524 cos: cos.clone(),
525 sin: sin.clone(),
526 neg_inf: neg_inf.clone(),
527 kv_cache: None,
528 span_attn: tracing::span!(tracing::Level::TRACE, "attn"),
529 span_rot: tracing::span!(tracing::Level::TRACE, "attn-rot"),
530 })
531 } else {
532 let in_proj = get_qtensor(
533 &ct,
534 reader,
535 device,
536 &[
537 format!("{prefix}.shortconv.in_proj.weight"),
538 format!("{prefix}.conv.in_proj.weight"),
539 ],
540 )?;
541 let out_proj = get_qtensor(
542 &ct,
543 reader,
544 device,
545 &[
546 format!("{prefix}.shortconv.out_proj.weight"),
547 format!("{prefix}.conv.out_proj.weight"),
548 ],
549 )?;
550 let conv = get_dequantized(
551 &ct,
552 reader,
553 device,
554 &[
555 format!("{prefix}.shortconv.conv.weight"),
556 format!("{prefix}.conv.conv.weight"),
557 format!("{prefix}.shortconv.conv"),
558 ],
559 )?;
560 LayerKind::ShortConv(ShortConvLayer {
561 in_proj: QMatMul::from_qtensor(in_proj)?,
562 out_proj: QMatMul::from_qtensor(out_proj)?,
563 conv,
564 l_cache,
565 cache: None,
566 })
567 };
568
569 layers.push(LayerWeights {
570 operator_norm: RmsNorm::from_qtensor(operator_norm, rms_norm_eps)?,
571 ffn_norm: RmsNorm::from_qtensor(ffn_norm, rms_norm_eps)?,
572 mlp,
573 kind,
574 span_mlp: tracing::span!(tracing::Level::TRACE, "ffn"),
575 });
576 }
577
578 Ok(Self {
579 tok_embeddings: Embedding::new(tok_embeddings, embedding_length),
580 layers,
581 norm,
582 output: QMatMul::from_qtensor(output_q)?,
583 masks: HashMap::new(),
584 span: tracing::span!(tracing::Level::TRACE, "model"),
585 span_output: tracing::span!(tracing::Level::TRACE, "output"),
586 })
587 }
588
589 fn mask(&mut self, seq_len: usize, index_pos: usize, device: &Device) -> Result<Tensor> {
590 let kv_len = index_pos + seq_len;
591 if let Some(mask) = self.masks.get(&(seq_len, kv_len)) {
592 Ok(mask.clone())
593 } else {
594 let mask = crate::utils::build_causal_mask(seq_len, index_pos, device)?;
595 self.masks.insert((seq_len, kv_len), mask.clone());
596 Ok(mask)
597 }
598 }
599
600 pub fn forward(&mut self, x: &Tensor, index_pos: usize) -> Result<Tensor> {
601 let (_b_sz, seq_len) = x.dims2()?;
602 let mask = if seq_len == 1 {
603 None
604 } else {
605 Some(self.mask(seq_len, index_pos, x.device())?)
606 };
607
608 let _enter = self.span.enter();
609 let mut hidden = self.tok_embeddings.forward(x)?;
610 for layer in self.layers.iter_mut() {
611 let residual = hidden.clone();
612 let normed = layer.operator_norm.forward(&hidden)?;
613 hidden = match &mut layer.kind {
614 LayerKind::Attention(attn) => attn.forward(&normed, mask.as_ref(), index_pos)?,
615 LayerKind::ShortConv(conv) => conv.forward(&normed, index_pos)?,
616 };
617 hidden = (hidden + residual)?;
618
619 let residual = hidden.clone();
620 let ff = layer.ffn_norm.forward(&hidden)?;
621 let _enter = layer.span_mlp.enter();
622 let ff = layer.mlp.forward(&ff)?;
623 hidden = (ff + residual)?;
624 }
625 let hidden = self.norm.forward(&hidden)?;
626 let hidden = hidden.i((.., seq_len - 1, ..))?;
627 let _enter = self.span_output.enter();
628 self.output.forward(&hidden)
629 }
630}