1use crate::models::with_tracing::{linear_no_bias as linear, Embedding, Linear, RmsNorm};
9use crate::utils::repeat_kv;
10use candle::{DType, Device, IndexOp, Module, Result, Tensor};
11use candle_nn::{Conv1d, Conv1dConfig, VarBuilder};
12use std::collections::HashMap;
13
14#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Deserialize)]
15#[serde(rename_all = "snake_case")]
16pub enum LayerType {
17 FullAttention,
18 Conv,
19}
20
21#[derive(Debug, Clone, serde::Deserialize)]
22pub struct Lfm2Config {
23 pub vocab_size: usize,
24 pub hidden_size: usize,
25 pub num_hidden_layers: usize,
26 pub num_attention_heads: usize,
27 #[serde(default = "default_num_key_value_heads")]
28 pub num_key_value_heads: usize,
29 #[serde(default = "default_norm_eps")]
30 pub norm_eps: f64,
31 #[serde(default = "default_rope_theta")]
32 pub rope_theta: f32,
33 #[serde(default = "default_max_position_embeddings")]
34 pub max_position_embeddings: usize,
35 #[serde(default = "default_conv_l_cache", alias = "conv_L_cache")]
36 pub conv_l_cache: usize,
37 #[serde(default)]
38 pub conv_bias: bool,
39 pub layer_types: Vec<LayerType>,
40 #[serde(default)]
41 pub tie_embedding: bool,
42 pub bos_token_id: Option<u32>,
43 pub eos_token_id: Option<u32>,
44 #[serde(default = "default_ffn_dim_multiplier")]
46 pub block_ffn_dim_multiplier: f32,
47 #[serde(default = "default_block_multiple_of")]
48 pub block_multiple_of: usize,
49}
50
51fn default_num_key_value_heads() -> usize {
52 8
53}
54
55fn default_norm_eps() -> f64 {
56 1e-5
57}
58
59fn default_rope_theta() -> f32 {
60 1_000_000.0
61}
62
63fn default_max_position_embeddings() -> usize {
64 128000
65}
66
67fn default_conv_l_cache() -> usize {
68 3
69}
70
71fn default_ffn_dim_multiplier() -> f32 {
72 1.0
73}
74
75fn default_block_multiple_of() -> usize {
76 256
77}
78
79impl Lfm2Config {
80 pub fn head_dim(&self) -> usize {
81 self.hidden_size / self.num_attention_heads
82 }
83
84 fn compute_intermediate_size(&self) -> usize {
87 let base_size = (self.hidden_size as f32 * 4.0 * self.block_ffn_dim_multiplier) as usize;
88 let multiple = self.block_multiple_of;
89 base_size.div_ceil(multiple) * multiple
90 }
91
92 pub fn into_config(self, use_flash_attn: bool) -> Config {
93 let intermediate_size = self.compute_intermediate_size();
95 Config {
96 vocab_size: self.vocab_size,
97 hidden_size: self.hidden_size,
98 intermediate_size,
99 num_hidden_layers: self.num_hidden_layers,
100 num_attention_heads: self.num_attention_heads,
101 num_key_value_heads: self.num_key_value_heads,
102 norm_eps: self.norm_eps,
103 rope_theta: self.rope_theta,
104 max_position_embeddings: self.max_position_embeddings,
105 conv_l_cache: self.conv_l_cache,
106 conv_bias: self.conv_bias,
107 layer_types: self.layer_types,
108 tie_embedding: self.tie_embedding,
109 bos_token_id: self.bos_token_id,
110 eos_token_id: self.eos_token_id,
111 use_flash_attn,
112 }
113 }
114}
115
116#[derive(Debug, Clone)]
117pub struct Config {
118 pub vocab_size: usize,
119 pub hidden_size: usize,
120 pub intermediate_size: usize,
121 pub num_hidden_layers: usize,
122 pub num_attention_heads: usize,
123 pub num_key_value_heads: usize,
124 pub norm_eps: f64,
125 pub rope_theta: f32,
126 pub max_position_embeddings: usize,
127 pub conv_l_cache: usize,
128 pub conv_bias: bool,
129 pub layer_types: Vec<LayerType>,
130 pub tie_embedding: bool,
131 pub bos_token_id: Option<u32>,
132 pub eos_token_id: Option<u32>,
133 pub use_flash_attn: bool,
134}
135
136impl Config {
137 pub fn head_dim(&self) -> usize {
138 self.hidden_size / self.num_attention_heads
139 }
140}
141
142#[derive(Debug, Clone)]
144pub struct Cache {
145 masks: HashMap<(usize, usize), Tensor>,
146 pub use_kv_cache: bool,
147 kvs: Vec<Option<(Tensor, Tensor)>>,
149 conv_states: Vec<Option<Tensor>>,
151 cos: Tensor,
152 sin: Tensor,
153 device: Device,
154}
155
156fn calculate_default_inv_freq(cfg: &Config) -> Vec<f32> {
157 let head_dim = cfg.head_dim();
158 (0..head_dim)
159 .step_by(2)
160 .map(|i| 1f32 / cfg.rope_theta.powf(i as f32 / head_dim as f32))
161 .collect()
162}
163
164impl Cache {
165 pub fn new(use_kv_cache: bool, dtype: DType, config: &Config, device: &Device) -> Result<Self> {
166 let theta = calculate_default_inv_freq(config);
167 let theta = Tensor::new(theta, device)?;
168
169 let idx_theta = Tensor::arange(0, config.max_position_embeddings as u32, device)?
170 .to_dtype(DType::F32)?
171 .reshape((config.max_position_embeddings, 1))?
172 .matmul(&theta.reshape((1, theta.elem_count()))?)?;
173 let cos = idx_theta.cos()?.to_dtype(dtype)?;
174 let sin = idx_theta.sin()?.to_dtype(dtype)?;
175
176 let num_layers = config.num_hidden_layers;
177 Ok(Self {
178 masks: HashMap::new(),
179 use_kv_cache,
180 kvs: vec![None; num_layers],
181 conv_states: vec![None; num_layers],
182 device: device.clone(),
183 cos,
184 sin,
185 })
186 }
187
188 fn mask(&mut self, seq_len: usize, index_pos: usize) -> Result<Tensor> {
189 let kv_len = index_pos + seq_len;
190 if let Some(mask) = self.masks.get(&(seq_len, kv_len)) {
191 Ok(mask.clone())
192 } else {
193 let mask = crate::utils::build_causal_mask(seq_len, index_pos, &self.device)?;
194 self.masks.insert((seq_len, kv_len), mask.clone());
195 Ok(mask)
196 }
197 }
198
199 pub fn clear(&mut self) {
200 self.kvs.iter_mut().for_each(|v| *v = None);
201 self.conv_states.iter_mut().for_each(|v| *v = None);
202 }
203}
204
205fn masked_fill(on_false: &Tensor, mask: &Tensor, on_true: f32) -> Result<Tensor> {
206 let shape = mask.shape();
207 let on_true = Tensor::new(on_true, on_false.device())?.broadcast_as(shape.dims())?;
208 let m = mask.where_cond(&on_true, on_false)?;
209 Ok(m)
210}
211
212#[cfg(feature = "flash-attn")]
213fn flash_attn(
214 q: &Tensor,
215 k: &Tensor,
216 v: &Tensor,
217 softmax_scale: f32,
218 causal: bool,
219) -> Result<Tensor> {
220 candle_flash_attn::flash_attn(q, k, v, softmax_scale, causal)
221}
222
223#[cfg(not(feature = "flash-attn"))]
224fn flash_attn(_: &Tensor, _: &Tensor, _: &Tensor, _: f32, _: bool) -> Result<Tensor> {
225 unimplemented!("compile with '--features flash-attn'")
226}
227
228#[derive(Debug, Clone)]
230struct Mlp {
231 gate_proj: Linear,
232 up_proj: Linear,
233 down_proj: Linear,
234 span: tracing::Span,
235}
236
237impl Mlp {
238 fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
239 let hidden_size = cfg.hidden_size;
240 let intermediate_size = cfg.intermediate_size;
241 let gate_proj = linear(hidden_size, intermediate_size, vb.pp("w1"))?;
243 let up_proj = linear(hidden_size, intermediate_size, vb.pp("w3"))?;
244 let down_proj = linear(intermediate_size, hidden_size, vb.pp("w2"))?;
245 Ok(Self {
246 gate_proj,
247 up_proj,
248 down_proj,
249 span: tracing::span!(tracing::Level::TRACE, "mlp"),
250 })
251 }
252
253 fn forward(&self, x: &Tensor) -> Result<Tensor> {
254 let _enter = self.span.enter();
255 let gate = candle_nn::ops::silu(&self.gate_proj.forward(x)?)?;
256 let up = self.up_proj.forward(x)?;
257 self.down_proj.forward(&(gate * up)?)
258 }
259}
260
261#[derive(Debug, Clone)]
263struct Attention {
264 q_proj: Linear,
265 k_proj: Linear,
266 v_proj: Linear,
267 o_proj: Linear,
268 q_norm: RmsNorm,
269 k_norm: RmsNorm,
270 num_attention_heads: usize,
271 num_key_value_heads: usize,
272 head_dim: usize,
273 use_flash_attn: bool,
274 span: tracing::Span,
275 span_rot: tracing::Span,
276}
277
278impl Attention {
279 fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
280 let hidden_size = cfg.hidden_size;
281 let num_attention_heads = cfg.num_attention_heads;
282 let num_key_value_heads = cfg.num_key_value_heads;
283 let head_dim = cfg.head_dim();
284
285 let q_proj = linear(hidden_size, num_attention_heads * head_dim, vb.pp("q_proj"))?;
286 let k_proj = linear(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?;
287 let v_proj = linear(hidden_size, num_key_value_heads * head_dim, vb.pp("v_proj"))?;
288 let o_proj = linear(
289 num_attention_heads * head_dim,
290 hidden_size,
291 vb.pp("out_proj"),
292 )?;
293
294 let q_norm = RmsNorm::new(head_dim, cfg.norm_eps, vb.pp("q_layernorm"))?;
295 let k_norm = RmsNorm::new(head_dim, cfg.norm_eps, vb.pp("k_layernorm"))?;
296
297 Ok(Self {
298 q_proj,
299 k_proj,
300 v_proj,
301 o_proj,
302 q_norm,
303 k_norm,
304 num_attention_heads,
305 num_key_value_heads,
306 head_dim,
307 use_flash_attn: cfg.use_flash_attn,
308 span: tracing::span!(tracing::Level::TRACE, "attn"),
309 span_rot: tracing::span!(tracing::Level::TRACE, "attn-rot"),
310 })
311 }
312
313 fn apply_rotary_emb(&self, x: &Tensor, index_pos: usize, cache: &Cache) -> Result<Tensor> {
314 let _enter = self.span_rot.enter();
315 let (_, _, seq_len, _) = x.dims4()?;
316 let cos = cache.cos.narrow(0, index_pos, seq_len)?;
317 let sin = cache.sin.narrow(0, index_pos, seq_len)?;
318 candle_nn::rotary_emb::rope(&x.contiguous()?, &cos, &sin)
319 }
320
321 fn forward(
322 &self,
323 x: &Tensor,
324 index_pos: usize,
325 block_idx: usize,
326 cache: &mut Cache,
327 ) -> Result<Tensor> {
328 let _enter = self.span.enter();
329 let (b_sz, seq_len, _) = x.dims3()?;
330
331 let q = self.q_proj.forward(x)?;
332 let k = self.k_proj.forward(x)?;
333 let v = self.v_proj.forward(x)?;
334
335 let q = q
337 .reshape((b_sz, seq_len, self.num_attention_heads, self.head_dim))?
338 .transpose(1, 2)?;
339 let k = k
340 .reshape((b_sz, seq_len, self.num_key_value_heads, self.head_dim))?
341 .transpose(1, 2)?;
342 let v = v
343 .reshape((b_sz, seq_len, self.num_key_value_heads, self.head_dim))?
344 .transpose(1, 2)?
345 .contiguous()?;
346
347 let q = self.q_norm.forward(&q.contiguous()?)?;
349 let k = self.k_norm.forward(&k.contiguous()?)?;
350
351 let q = self.apply_rotary_emb(&q, index_pos, cache)?;
353 let k = self.apply_rotary_emb(&k, index_pos, cache)?;
354
355 let (k, v) = if cache.use_kv_cache {
357 match &cache.kvs[block_idx] {
358 Some((k_cache, v_cache)) if index_pos > 0 => {
359 let k = Tensor::cat(&[k_cache, &k], 2)?.contiguous()?;
360 let v = Tensor::cat(&[v_cache, &v], 2)?.contiguous()?;
361 (k, v)
362 }
363 _ => (k, v),
364 }
365 } else {
366 (k, v)
367 };
368
369 if cache.use_kv_cache {
370 cache.kvs[block_idx] = Some((k.clone(), v.clone()));
371 }
372
373 let k = repeat_kv(k, self.num_attention_heads / self.num_key_value_heads)?;
375 let v = repeat_kv(v, self.num_attention_heads / self.num_key_value_heads)?;
376
377 let y = if self.use_flash_attn {
378 let q = q.transpose(1, 2)?;
379 let k = k.transpose(1, 2)?;
380 let v = v.transpose(1, 2)?;
381 let softmax_scale = 1f32 / (self.head_dim as f32).sqrt();
382 flash_attn(&q, &k, &v, softmax_scale, seq_len > 1)?.transpose(1, 2)?
383 } else {
384 let in_dtype = q.dtype();
385 let q = q.to_dtype(DType::F32)?;
386 let k = k.to_dtype(DType::F32)?;
387 let v = v.to_dtype(DType::F32)?;
388 let att = (q.matmul(&k.t()?)? / (self.head_dim as f64).sqrt())?;
389 let att = if seq_len == 1 {
390 att
391 } else {
392 let mask = cache.mask(seq_len, index_pos)?.broadcast_as(att.shape())?;
393 masked_fill(&att, &mask, f32::NEG_INFINITY)?
394 };
395 let att = candle_nn::ops::softmax_last_dim(&att)?;
396 att.matmul(&v.contiguous()?)?.to_dtype(in_dtype)?
397 };
398
399 let y = y.transpose(1, 2)?.reshape((
400 b_sz,
401 seq_len,
402 self.num_attention_heads * self.head_dim,
403 ))?;
404 self.o_proj.forward(&y)
405 }
406}
407
408#[derive(Debug, Clone)]
410struct ShortConv {
411 in_proj: Linear,
412 out_proj: Linear,
413 conv_weight: Tensor,
414 l_cache: usize,
415 hidden_size: usize,
416 span: tracing::Span,
417}
418
419impl ShortConv {
420 fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
421 let hidden_size = cfg.hidden_size;
422 let l_cache = cfg.conv_l_cache;
423
424 let in_proj = linear(hidden_size, 3 * hidden_size, vb.pp("in_proj"))?;
426 let out_proj = linear(hidden_size, hidden_size, vb.pp("out_proj"))?;
427
428 let conv_weight = vb.get((hidden_size, 1, l_cache), "conv.weight")?;
430
431 Ok(Self {
432 in_proj,
433 out_proj,
434 conv_weight,
435 l_cache,
436 hidden_size,
437 span: tracing::span!(tracing::Level::TRACE, "shortconv"),
438 })
439 }
440
441 fn forward(&self, x: &Tensor, block_idx: usize, cache: &mut Cache) -> Result<Tensor> {
442 let _enter = self.span.enter();
443 let (b_sz, seq_len, _) = x.dims3()?;
444
445 let bcx = self.in_proj.forward(x)?.transpose(1, 2)?;
447 let b = bcx.narrow(1, 0, self.hidden_size)?;
448 let c = bcx.narrow(1, self.hidden_size, self.hidden_size)?;
449 let x_proj = bcx.narrow(1, 2 * self.hidden_size, self.hidden_size)?;
450
451 let bx = (b * &x_proj)?.contiguous()?;
453
454 let conv_weight = self.conv_weight.squeeze(1)?;
456
457 let conv_out = if seq_len == 1 {
458 let mut state = match &cache.conv_states[block_idx] {
460 Some(s) => s.clone(),
461 None => Tensor::zeros(
462 (b_sz, self.hidden_size, self.l_cache),
463 bx.dtype(),
464 bx.device(),
465 )?,
466 };
467
468 if self.l_cache > 1 {
470 let tail = state.narrow(2, 1, self.l_cache - 1)?;
471 state = Tensor::cat(&[tail, bx.clone()], 2)?;
472 } else {
473 state = bx.clone();
474 }
475
476 if cache.use_kv_cache {
477 cache.conv_states[block_idx] = Some(state.clone());
478 }
479
480 (state * conv_weight.unsqueeze(0)?)?
482 .sum_keepdim(2)?
483 .contiguous()?
484 } else {
485 let conv = Conv1d::new(
487 self.conv_weight.clone(),
488 None,
489 Conv1dConfig {
490 padding: self.l_cache.saturating_sub(1),
491 groups: self.hidden_size,
492 ..Default::default()
493 },
494 );
495 let mut out = conv.forward(&bx)?;
496 out = out.narrow(2, 0, seq_len)?;
497
498 if cache.use_kv_cache && self.l_cache > 0 {
500 let start = seq_len.saturating_sub(self.l_cache);
501 let cache_len = seq_len - start;
502 let mut cache_src = bx.narrow(2, start, cache_len)?;
503 if cache_len < self.l_cache {
504 let pad = self.l_cache - cache_len;
505 let zeros = Tensor::zeros(
506 (b_sz, self.hidden_size, pad),
507 cache_src.dtype(),
508 cache_src.device(),
509 )?;
510 cache_src = Tensor::cat(&[zeros, cache_src], 2)?;
511 }
512 cache.conv_states[block_idx] = Some(cache_src);
513 }
514
515 out
516 };
517
518 let conv_out = (c * &conv_out)?;
520 let conv_out = conv_out.transpose(1, 2)?.contiguous()?;
521 self.out_proj.forward(&conv_out)
522 }
523}
524
525#[derive(Debug, Clone)]
527enum LayerKind {
528 Attention(Box<Attention>),
529 ShortConv(ShortConv),
530}
531
532#[derive(Debug, Clone)]
533struct DecoderLayer {
534 input_layernorm: RmsNorm,
535 post_attention_layernorm: RmsNorm,
536 mlp: Mlp,
537 kind: LayerKind,
538 span: tracing::Span,
539}
540
541impl DecoderLayer {
542 fn new(cfg: &Config, layer_idx: usize, vb: VarBuilder) -> Result<Self> {
543 let input_layernorm = RmsNorm::new(cfg.hidden_size, cfg.norm_eps, vb.pp("operator_norm"))?;
545 let post_attention_layernorm =
546 RmsNorm::new(cfg.hidden_size, cfg.norm_eps, vb.pp("ffn_norm"))?;
547 let mlp = Mlp::new(cfg, vb.pp("feed_forward"))?;
549
550 let layer_type = cfg
551 .layer_types
552 .get(layer_idx)
553 .copied()
554 .unwrap_or(LayerType::FullAttention);
555 let kind = match layer_type {
556 LayerType::FullAttention => {
557 LayerKind::Attention(Box::new(Attention::new(cfg, vb.pp("self_attn"))?))
558 }
559 LayerType::Conv => LayerKind::ShortConv(ShortConv::new(cfg, vb.pp("conv"))?),
560 };
561
562 Ok(Self {
563 input_layernorm,
564 post_attention_layernorm,
565 mlp,
566 kind,
567 span: tracing::span!(tracing::Level::TRACE, "layer"),
568 })
569 }
570
571 fn forward(
572 &self,
573 x: &Tensor,
574 index_pos: usize,
575 block_idx: usize,
576 cache: &mut Cache,
577 ) -> Result<Tensor> {
578 let _enter = self.span.enter();
579 let residual = x;
580 let x = self.input_layernorm.forward(x)?;
581
582 let x = match &self.kind {
583 LayerKind::Attention(attn) => attn.forward(&x, index_pos, block_idx, cache)?,
584 LayerKind::ShortConv(conv) => conv.forward(&x, block_idx, cache)?,
585 };
586
587 let x = (x + residual)?;
588 let residual = &x;
589 let x = self.post_attention_layernorm.forward(&x)?;
590 let x = self.mlp.forward(&x)?;
591 x + residual
592 }
593}
594
595#[derive(Debug, Clone)]
597pub struct Model {
598 embed_tokens: Embedding,
599 layers: Vec<DecoderLayer>,
600 embedding_norm: RmsNorm,
601 lm_head: Linear,
602 dtype: DType,
603}
604
605impl Model {
606 pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
607 let vb_m = vb.pp("model");
608
609 let embed_tokens =
610 Embedding::new(cfg.vocab_size, cfg.hidden_size, vb_m.pp("embed_tokens"))?;
611
612 let mut layers = Vec::with_capacity(cfg.num_hidden_layers);
613 let vb_l = vb_m.pp("layers");
614 for layer_idx in 0..cfg.num_hidden_layers {
615 let layer = DecoderLayer::new(cfg, layer_idx, vb_l.pp(layer_idx))?;
616 layers.push(layer);
617 }
618
619 let embedding_norm =
620 RmsNorm::new(cfg.hidden_size, cfg.norm_eps, vb_m.pp("embedding_norm"))?;
621
622 let lm_head = if cfg.tie_embedding {
623 Linear::from_weights(embed_tokens.embeddings().clone(), None)
624 } else {
625 linear(cfg.hidden_size, cfg.vocab_size, vb.pp("lm_head"))?
626 };
627
628 Ok(Self {
629 embed_tokens,
630 layers,
631 embedding_norm,
632 lm_head,
633 dtype: vb.dtype(),
634 })
635 }
636
637 pub fn forward(
638 &self,
639 input_ids: &Tensor,
640 index_pos: usize,
641 cache: &mut Cache,
642 ) -> Result<Tensor> {
643 let (_, seq_len) = input_ids.dims2()?;
644 let mut hidden_states = self.embed_tokens.forward(input_ids)?;
645
646 for (block_idx, layer) in self.layers.iter().enumerate() {
647 hidden_states = layer.forward(&hidden_states, index_pos, block_idx, cache)?;
648 }
649
650 let hidden_states = self.embedding_norm.forward(&hidden_states)?;
651 let hidden_states = hidden_states.i((.., seq_len - 1, ..))?.contiguous()?;
652 let logits = self.lm_head.forward(&hidden_states)?;
653 logits.to_dtype(DType::F32)
654 }
655
656 pub fn dtype(&self) -> DType {
657 self.dtype
658 }
659}