Skip to main content

candle_transformers/models/
quantized_lfm2.rs

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        // conv_weight shape -> [hidden, l_cache]
194        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}