Skip to main content

candle_transformers/models/
llama.rs

1//! Llama inference implementation.
2//!
3//! See ["LLaMA: Open and Efficient Foundation Language Models"](https://arxiv.org/abs/2302.13971)
4//!
5//! Implementation based on Hugging Face's [transformers](https://github.com/huggingface/transformers/blob/main/src/transformers/models/llama/modeling_llama.py)
6
7use super::with_tracing::{linear_no_bias as linear, Linear, RmsNorm};
8use candle::{DType, Device, IndexOp, Result, Tensor, D};
9use candle_nn::{embedding, Embedding, Module, VarBuilder};
10use std::{collections::HashMap, f32::consts::PI};
11
12pub const DEFAULT_MAX_SEQ_LEN: usize = 4096;
13
14#[derive(Debug, Clone, serde::Deserialize, Default)]
15pub enum Llama3RopeType {
16    #[serde(rename = "llama3")]
17    Llama3,
18    #[default]
19    #[serde(rename = "default")]
20    Default,
21}
22
23#[derive(Debug, Clone, serde::Deserialize, Default)]
24pub struct Llama3RopeConfig {
25    pub factor: f32,
26    pub low_freq_factor: f32,
27    pub high_freq_factor: f32,
28    pub original_max_position_embeddings: usize,
29    pub rope_type: Llama3RopeType,
30}
31#[derive(Debug, Clone, serde::Deserialize)]
32#[serde(untagged)]
33pub enum LlamaEosToks {
34    Single(u32),
35    Multiple(Vec<u32>),
36}
37
38#[derive(Debug, Clone, serde::Deserialize)]
39pub struct LlamaConfig {
40    pub hidden_size: usize,
41    pub intermediate_size: usize,
42    pub vocab_size: usize,
43    pub num_hidden_layers: usize,
44    pub num_attention_heads: usize,
45    pub num_key_value_heads: Option<usize>,
46    pub rms_norm_eps: f64,
47    #[serde(default = "default_rope")]
48    pub rope_theta: f32,
49    pub bos_token_id: Option<u32>,
50    pub eos_token_id: Option<LlamaEosToks>,
51    pub rope_scaling: Option<Llama3RopeConfig>,
52    pub max_position_embeddings: usize,
53    pub tie_word_embeddings: Option<bool>,
54}
55
56impl LlamaConfig {
57    pub fn num_key_value_heads(&self) -> usize {
58        self.num_key_value_heads.unwrap_or(self.num_attention_heads)
59    }
60}
61
62fn default_rope() -> f32 {
63    10_000.0
64}
65
66impl LlamaConfig {
67    pub fn into_config(self, use_flash_attn: bool) -> Config {
68        Config {
69            hidden_size: self.hidden_size,
70            intermediate_size: self.intermediate_size,
71            vocab_size: self.vocab_size,
72            num_hidden_layers: self.num_hidden_layers,
73            num_attention_heads: self.num_attention_heads,
74            num_key_value_heads: self.num_key_value_heads(),
75            rms_norm_eps: self.rms_norm_eps,
76            rope_theta: self.rope_theta,
77            use_flash_attn,
78            bos_token_id: self.bos_token_id,
79            eos_token_id: self.eos_token_id,
80            rope_scaling: self.rope_scaling,
81            max_position_embeddings: self.max_position_embeddings,
82            tie_word_embeddings: self.tie_word_embeddings.unwrap_or(false),
83        }
84    }
85}
86
87#[derive(Debug, Clone)]
88pub struct Config {
89    pub hidden_size: usize,
90    pub intermediate_size: usize,
91    pub vocab_size: usize,
92    pub num_hidden_layers: usize,
93    pub num_attention_heads: usize,
94    pub num_key_value_heads: usize,
95    pub use_flash_attn: bool,
96    pub rms_norm_eps: f64,
97    pub rope_theta: f32,
98    pub bos_token_id: Option<u32>,
99    pub eos_token_id: Option<LlamaEosToks>,
100    pub rope_scaling: Option<Llama3RopeConfig>,
101    pub max_position_embeddings: usize,
102    pub tie_word_embeddings: bool,
103}
104
105impl Config {
106    pub fn config_7b_v1(use_flash_attn: bool) -> Self {
107        Self {
108            hidden_size: 4096,
109            intermediate_size: 11008,
110            vocab_size: 32000,
111            num_hidden_layers: 32,
112            num_attention_heads: 32,
113            num_key_value_heads: 32,
114            use_flash_attn,
115            rms_norm_eps: 1e-6,
116            rope_theta: 10_000.0,
117            bos_token_id: None,
118            eos_token_id: None,
119            rope_scaling: None,
120            max_position_embeddings: DEFAULT_MAX_SEQ_LEN,
121            tie_word_embeddings: false,
122        }
123    }
124
125    pub fn config_7b_v2(use_flash_attn: bool) -> Self {
126        Self {
127            hidden_size: 4096,
128            intermediate_size: 11008,
129            vocab_size: 32000,
130            num_hidden_layers: 32,
131            num_attention_heads: 32,
132            num_key_value_heads: 32,
133            use_flash_attn,
134            rms_norm_eps: 1e-5,
135            rope_theta: 10_000.0,
136            bos_token_id: None,
137            eos_token_id: None,
138            rope_scaling: None,
139            max_position_embeddings: DEFAULT_MAX_SEQ_LEN,
140            tie_word_embeddings: false,
141        }
142    }
143}
144
145#[derive(Debug, Clone)]
146pub struct Cache {
147    masks: HashMap<(usize, usize), Tensor>,
148    pub use_kv_cache: bool,
149    kvs: Vec<Option<(Tensor, Tensor)>>,
150    cos: Tensor,
151    sin: Tensor,
152    device: Device,
153}
154
155fn calculate_default_inv_freq(cfg: &Config) -> Vec<f32> {
156    let head_dim = cfg.hidden_size / cfg.num_attention_heads;
157    (0..head_dim)
158        .step_by(2)
159        .map(|i| 1f32 / cfg.rope_theta.powf(i as f32 / head_dim as f32))
160        .collect()
161}
162
163impl Cache {
164    pub fn new(use_kv_cache: bool, dtype: DType, config: &Config, device: &Device) -> Result<Self> {
165        // precompute freqs_cis
166        let theta = match &config.rope_scaling {
167            None
168            | Some(Llama3RopeConfig {
169                rope_type: Llama3RopeType::Default,
170                ..
171            }) => calculate_default_inv_freq(config),
172            Some(rope_scaling) => {
173                let low_freq_wavelen = rope_scaling.original_max_position_embeddings as f32
174                    / rope_scaling.low_freq_factor;
175                let high_freq_wavelen = rope_scaling.original_max_position_embeddings as f32
176                    / rope_scaling.high_freq_factor;
177
178                calculate_default_inv_freq(config)
179                    .into_iter()
180                    .map(|freq| {
181                        let wavelen = 2. * PI / freq;
182                        if wavelen < high_freq_wavelen {
183                            freq
184                        } else if wavelen > low_freq_wavelen {
185                            freq / rope_scaling.factor
186                        } else {
187                            let smooth = (rope_scaling.original_max_position_embeddings as f32
188                                / wavelen
189                                - rope_scaling.low_freq_factor)
190                                / (rope_scaling.high_freq_factor - rope_scaling.low_freq_factor);
191                            (1. - smooth) * freq / rope_scaling.factor + smooth * freq
192                        }
193                    })
194                    .collect::<Vec<_>>()
195            }
196        };
197
198        let theta = Tensor::new(theta, device)?;
199
200        let idx_theta = Tensor::arange(0, config.max_position_embeddings as u32, device)?
201            .to_dtype(DType::F32)?
202            .reshape((config.max_position_embeddings, 1))?
203            .matmul(&theta.reshape((1, theta.elem_count()))?)?;
204        // This is different from the paper, see:
205        // https://github.com/huggingface/transformers/blob/6112b1c6442aaf7affd2b0676a1cd4eee30c45cf/src/transformers/models/llama/modeling_llama.py#L112
206        let cos = idx_theta.cos()?.to_dtype(dtype)?;
207        let sin = idx_theta.sin()?.to_dtype(dtype)?;
208        Ok(Self {
209            masks: HashMap::new(),
210            use_kv_cache,
211            kvs: vec![None; config.num_hidden_layers],
212            device: device.clone(),
213            cos,
214            sin,
215        })
216    }
217
218    fn mask(&mut self, seq_len: usize, index_pos: usize) -> Result<Tensor> {
219        let kv_len = index_pos + seq_len;
220        if let Some(mask) = self.masks.get(&(seq_len, kv_len)) {
221            Ok(mask.clone())
222        } else {
223            let mask = crate::utils::build_causal_mask(seq_len, index_pos, &self.device)?;
224            self.masks.insert((seq_len, kv_len), mask.clone());
225            Ok(mask)
226        }
227    }
228}
229
230#[derive(Debug, Clone)]
231struct CausalSelfAttention {
232    q_proj: Linear,
233    k_proj: Linear,
234    v_proj: Linear,
235    o_proj: Linear,
236    num_attention_heads: usize,
237    num_key_value_heads: usize,
238    head_dim: usize,
239    use_flash_attn: bool,
240    span: tracing::Span,
241    span_rot: tracing::Span,
242    max_position_embeddings: usize,
243}
244
245#[cfg(feature = "flash-attn")]
246fn flash_attn(
247    q: &Tensor,
248    k: &Tensor,
249    v: &Tensor,
250    softmax_scale: f32,
251    causal: bool,
252) -> Result<Tensor> {
253    candle_flash_attn::flash_attn(q, k, v, softmax_scale, causal)
254}
255
256#[cfg(not(feature = "flash-attn"))]
257fn flash_attn(_: &Tensor, _: &Tensor, _: &Tensor, _: f32, _: bool) -> Result<Tensor> {
258    unimplemented!("compile with '--features flash-attn'")
259}
260
261impl CausalSelfAttention {
262    fn apply_rotary_emb(&self, x: &Tensor, index_pos: usize, cache: &Cache) -> Result<Tensor> {
263        let _enter = self.span_rot.enter();
264        let (_b_sz, _, seq_len, _hidden_size) = x.dims4()?;
265        let cos = cache.cos.narrow(0, index_pos, seq_len)?;
266        let sin = cache.sin.narrow(0, index_pos, seq_len)?;
267        candle_nn::rotary_emb::rope(x, &cos, &sin)
268    }
269
270    fn forward(
271        &self,
272        x: &Tensor,
273        index_pos: usize,
274        block_idx: usize,
275        cache: &mut Cache,
276    ) -> Result<Tensor> {
277        let _enter = self.span.enter();
278        let (b_sz, seq_len, hidden_size) = x.dims3()?;
279        let q = self.q_proj.forward(x)?;
280        let k = self.k_proj.forward(x)?;
281        let v = self.v_proj.forward(x)?;
282
283        let q = q
284            .reshape((b_sz, seq_len, self.num_attention_heads, self.head_dim))?
285            .transpose(1, 2)?
286            .contiguous()?;
287        let k = k
288            .reshape((b_sz, seq_len, self.num_key_value_heads, self.head_dim))?
289            .transpose(1, 2)?
290            .contiguous()?;
291        let mut v = v
292            .reshape((b_sz, seq_len, self.num_key_value_heads, self.head_dim))?
293            .transpose(1, 2)?;
294
295        let q = self.apply_rotary_emb(&q, index_pos, cache)?;
296        let mut k = self.apply_rotary_emb(&k, index_pos, cache)?;
297
298        if cache.use_kv_cache {
299            if let Some((cache_k, cache_v)) = &cache.kvs[block_idx] {
300                k = Tensor::cat(&[cache_k, &k], 2)?.contiguous()?;
301                v = Tensor::cat(&[cache_v, &v], 2)?.contiguous()?;
302                let k_seq_len = k.dims()[1];
303                if k_seq_len > self.max_position_embeddings {
304                    k = k
305                        .narrow(
306                            D::Minus1,
307                            k_seq_len - self.max_position_embeddings,
308                            self.max_position_embeddings,
309                        )?
310                        .contiguous()?
311                }
312                let v_seq_len = v.dims()[1];
313                if v_seq_len > 2 * self.max_position_embeddings {
314                    v = v
315                        .narrow(
316                            D::Minus1,
317                            v_seq_len - self.max_position_embeddings,
318                            self.max_position_embeddings,
319                        )?
320                        .contiguous()?
321                }
322            }
323            cache.kvs[block_idx] = Some((k.clone(), v.clone()))
324        }
325
326        let k = self.repeat_kv(k)?;
327        let v = self.repeat_kv(v)?;
328
329        let y = if self.use_flash_attn {
330            // flash-attn expects (b_sz, seq_len, nheads, head_dim)
331            let q = q.transpose(1, 2)?;
332            let k = k.transpose(1, 2)?;
333            let v = v.transpose(1, 2)?;
334            let softmax_scale = 1f32 / (self.head_dim as f32).sqrt();
335            flash_attn(&q, &k, &v, softmax_scale, seq_len > 1)?.transpose(1, 2)?
336        } else {
337            let in_dtype = q.dtype();
338            let q = q.to_dtype(DType::F32)?;
339            let k = k.to_dtype(DType::F32)?;
340            let v = v.to_dtype(DType::F32)?;
341            let att = (q.matmul(&k.t()?)? / (self.head_dim as f64).sqrt())?;
342            let att = if seq_len == 1 {
343                att
344            } else {
345                let mask = cache.mask(seq_len, index_pos)?.broadcast_as(att.shape())?;
346                masked_fill(&att, &mask, f32::NEG_INFINITY)?
347            };
348
349            let att = candle_nn::ops::softmax_last_dim(&att)?;
350            // Convert to contiguous as matmul doesn't support strided vs for now.
351            att.matmul(&v.contiguous()?)?.to_dtype(in_dtype)?
352        };
353        let y = y.transpose(1, 2)?.reshape(&[b_sz, seq_len, hidden_size])?;
354        let y = self.o_proj.forward(&y)?;
355        Ok(y)
356    }
357
358    fn repeat_kv(&self, x: Tensor) -> Result<Tensor> {
359        crate::utils::repeat_kv(x, self.num_attention_heads / self.num_key_value_heads)
360    }
361
362    fn load(vb: VarBuilder, cfg: &Config) -> Result<Self> {
363        let span = tracing::span!(tracing::Level::TRACE, "attn");
364        let span_rot = tracing::span!(tracing::Level::TRACE, "attn-rot");
365        let size_in = cfg.hidden_size;
366        let size_q = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_attention_heads;
367        let size_kv = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_key_value_heads;
368        let q_proj = linear(size_in, size_q, vb.pp("q_proj"))?;
369        let k_proj = linear(size_in, size_kv, vb.pp("k_proj"))?;
370        let v_proj = linear(size_in, size_kv, vb.pp("v_proj"))?;
371        let o_proj = linear(size_q, size_in, vb.pp("o_proj"))?;
372        Ok(Self {
373            q_proj,
374            k_proj,
375            v_proj,
376            o_proj,
377            num_attention_heads: cfg.num_attention_heads,
378            num_key_value_heads: cfg.num_key_value_heads,
379            head_dim: cfg.hidden_size / cfg.num_attention_heads,
380            use_flash_attn: cfg.use_flash_attn,
381            span,
382            span_rot,
383            max_position_embeddings: cfg.max_position_embeddings,
384        })
385    }
386}
387
388fn masked_fill(on_false: &Tensor, mask: &Tensor, on_true: f32) -> Result<Tensor> {
389    let shape = mask.shape();
390    let on_true = Tensor::new(on_true, on_false.device())?.broadcast_as(shape.dims())?;
391    let m = mask.where_cond(&on_true, on_false)?;
392    Ok(m)
393}
394
395#[derive(Debug, Clone)]
396struct Mlp {
397    c_fc1: Linear,
398    c_fc2: Linear,
399    c_proj: Linear,
400    span: tracing::Span,
401}
402
403impl Mlp {
404    fn forward(&self, x: &Tensor) -> Result<Tensor> {
405        let _enter = self.span.enter();
406        let x = (candle_nn::ops::silu(&self.c_fc1.forward(x)?)? * self.c_fc2.forward(x)?)?;
407        self.c_proj.forward(&x)
408    }
409
410    fn load(vb: VarBuilder, cfg: &Config) -> Result<Self> {
411        let span = tracing::span!(tracing::Level::TRACE, "mlp");
412        let h_size = cfg.hidden_size;
413        let i_size = cfg.intermediate_size;
414        let c_fc1 = linear(h_size, i_size, vb.pp("gate_proj"))?;
415        let c_fc2 = linear(h_size, i_size, vb.pp("up_proj"))?;
416        let c_proj = linear(i_size, h_size, vb.pp("down_proj"))?;
417        Ok(Self {
418            c_fc1,
419            c_fc2,
420            c_proj,
421            span,
422        })
423    }
424}
425
426#[derive(Debug, Clone)]
427struct Block {
428    rms_1: RmsNorm,
429    attn: CausalSelfAttention,
430    rms_2: RmsNorm,
431    mlp: Mlp,
432    span: tracing::Span,
433}
434
435impl Block {
436    fn forward(
437        &self,
438        x: &Tensor,
439        index_pos: usize,
440        block_idx: usize,
441        cache: &mut Cache,
442    ) -> Result<Tensor> {
443        let _enter = self.span.enter();
444        let residual = x;
445        let x = self.rms_1.forward(x)?;
446        let x = (self.attn.forward(&x, index_pos, block_idx, cache)? + residual)?;
447        let residual = &x;
448        let x = (self.mlp.forward(&self.rms_2.forward(&x)?)? + residual)?;
449        Ok(x)
450    }
451
452    fn load(vb: VarBuilder, cfg: &Config) -> Result<Self> {
453        let span = tracing::span!(tracing::Level::TRACE, "block");
454        let attn = CausalSelfAttention::load(vb.pp("self_attn"), cfg)?;
455        let mlp = Mlp::load(vb.pp("mlp"), cfg)?;
456        let rms_1 = RmsNorm::new(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("input_layernorm"))?;
457        let rms_2 = RmsNorm::new(
458            cfg.hidden_size,
459            cfg.rms_norm_eps,
460            vb.pp("post_attention_layernorm"),
461        )?;
462        Ok(Self {
463            rms_1,
464            attn,
465            rms_2,
466            mlp,
467            span,
468        })
469    }
470}
471
472#[derive(Debug, Clone)]
473pub struct Llama {
474    wte: Embedding,
475    blocks: Vec<Block>,
476    ln_f: RmsNorm,
477    lm_head: Linear,
478}
479
480impl Llama {
481    // required by LLaVA
482    pub fn embed(&self, x: &Tensor) -> Result<Tensor> {
483        self.wte.forward(x)
484    }
485    // required by LLaVA
486    pub fn forward_input_embed(
487        &self,
488        input_embed: &Tensor,
489        index_pos: usize,
490        cache: &mut Cache,
491    ) -> Result<Tensor> {
492        let (_, seq_len, _) = input_embed.dims3()?;
493        let mut x = input_embed.clone();
494        for (block_idx, block) in self.blocks.iter().enumerate() {
495            x = block.forward(&x, index_pos, block_idx, cache)?;
496        }
497        let x = self.ln_f.forward(&x)?;
498        let x = x.i((.., seq_len - 1, ..))?.contiguous()?;
499        let logits = self.lm_head.forward(&x)?;
500        logits.to_dtype(DType::F32)
501    }
502
503    pub fn forward(&self, x: &Tensor, index_pos: usize, cache: &mut Cache) -> Result<Tensor> {
504        let (_b_sz, seq_len) = x.dims2()?;
505        let mut x = self.wte.forward(x)?;
506        for (block_idx, block) in self.blocks.iter().enumerate() {
507            x = block.forward(&x, index_pos, block_idx, cache)?;
508        }
509        let x = self.ln_f.forward(&x)?;
510        let x = x.i((.., seq_len - 1, ..))?.contiguous()?;
511        let logits = self.lm_head.forward(&x)?;
512        logits.to_dtype(DType::F32)
513    }
514
515    pub fn load(vb: VarBuilder, cfg: &Config) -> Result<Self> {
516        let wte = embedding(cfg.vocab_size, cfg.hidden_size, vb.pp("model.embed_tokens"))?;
517        let lm_head = if cfg.tie_word_embeddings {
518            Linear::from_weights(wte.embeddings().clone(), None)
519        } else {
520            linear(cfg.hidden_size, cfg.vocab_size, vb.pp("lm_head"))?
521        };
522        let ln_f = RmsNorm::new(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("model.norm"))?;
523        let blocks: Vec<_> = (0..cfg.num_hidden_layers)
524            .map(|i| Block::load(vb.pp(format!("model.layers.{i}")), cfg).unwrap())
525            .collect();
526
527        Ok(Self {
528            wte,
529            blocks,
530            ln_f,
531            lm_head,
532        })
533    }
534}