Skip to main content

candle_transformers/models/
granite.rs

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