Skip to main content

candle_transformers/models/
granitemoehybrid.rs

1//! GraniteMoeHybrid 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 GraniteMoeHybridRopeType {
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 GraniteMoeHybridRopeConfig {
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: GraniteMoeHybridRopeType,
29}
30
31#[derive(Debug, Clone, serde::Deserialize)]
32pub struct GraniteMoeHybridConfig {
33    pub hidden_size: usize,
34    pub intermediate_size: usize,
35    pub vocab_size: usize,
36    pub num_hidden_layers: usize,
37    pub num_attention_heads: usize,
38    pub num_key_value_heads: Option<usize>,
39    pub rms_norm_eps: f64,
40    #[serde(default = "default_rope")]
41    pub rope_theta: f32,
42    pub bos_token_id: Option<u32>,
43    pub eos_token_id: Option<u32>,
44    pub rope_scaling: Option<GraniteMoeHybridRopeConfig>,
45    pub max_position_embeddings: usize,
46    #[serde(default)]
47    pub layer_types: Vec<GraniteMoeHybridLayerType>,
48    #[serde(default = "default_one")]
49    pub attention_multiplier: f32,
50    #[serde(default = "default_one")]
51    pub embedding_multiplier: f32,
52    #[serde(default = "default_one")]
53    pub residual_multiplier: f32,
54    #[serde(default = "default_one")]
55    pub logits_scaling: f32,
56    #[serde(default)]
57    pub shared_intermediate_size: Option<usize>,
58}
59
60impl GraniteMoeHybridConfig {
61    pub fn num_key_value_heads(&self) -> usize {
62        self.num_key_value_heads.unwrap_or(self.num_attention_heads)
63    }
64}
65
66fn default_rope() -> f32 {
67    10_000.0
68}
69
70fn default_one() -> f32 {
71    1.0
72}
73
74#[derive(Debug, Clone, serde::Deserialize, Default)]
75#[serde(rename_all = "lowercase")]
76pub enum GraniteMoeHybridLayerType {
77    #[default]
78    Attention,
79    Mamba,
80}
81
82impl GraniteMoeHybridConfig {
83    pub fn into_config(self, use_flash_attn: bool) -> GraniteMoeHybridInternalConfig {
84        let layer_types = if self.layer_types.is_empty() {
85            vec![GraniteMoeHybridLayerType::Attention; self.num_hidden_layers]
86        } else {
87            self.layer_types.clone()
88        };
89        let shared_intermediate_size = self
90            .shared_intermediate_size
91            .unwrap_or(self.intermediate_size);
92        GraniteMoeHybridInternalConfig {
93            hidden_size: self.hidden_size,
94            intermediate_size: self.intermediate_size,
95            shared_intermediate_size,
96            vocab_size: self.vocab_size,
97            num_hidden_layers: self.num_hidden_layers,
98            num_attention_heads: self.num_attention_heads,
99            num_key_value_heads: self.num_key_value_heads(),
100            use_flash_attn,
101            rms_norm_eps: self.rms_norm_eps,
102            rope_theta: self.rope_theta,
103            bos_token_id: self.bos_token_id,
104            eos_token_id: self.eos_token_id,
105            rope_scaling: self.rope_scaling,
106            max_position_embeddings: self.max_position_embeddings,
107            layer_types,
108            attention_multiplier: self.attention_multiplier,
109            embedding_multiplier: self.embedding_multiplier,
110            residual_multiplier: self.residual_multiplier,
111            logits_scaling: self.logits_scaling,
112        }
113    }
114}
115
116#[derive(Debug, Clone)]
117pub struct GraniteMoeHybridInternalConfig {
118    pub hidden_size: usize,
119    pub intermediate_size: usize,
120    pub shared_intermediate_size: usize,
121    pub vocab_size: usize,
122    pub num_hidden_layers: usize,
123    pub num_attention_heads: usize,
124    pub num_key_value_heads: usize,
125    pub use_flash_attn: bool,
126    pub rms_norm_eps: f64,
127    pub rope_theta: f32,
128    pub bos_token_id: Option<u32>,
129    pub eos_token_id: Option<u32>,
130    pub rope_scaling: Option<GraniteMoeHybridRopeConfig>,
131    pub max_position_embeddings: usize,
132    pub layer_types: Vec<GraniteMoeHybridLayerType>,
133    pub attention_multiplier: f32,
134    pub embedding_multiplier: f32,
135    pub residual_multiplier: f32,
136    pub logits_scaling: f32,
137}
138
139#[derive(Debug, Clone)]
140pub struct GraniteMoeHybridCache {
141    masks: HashMap<(usize, usize), Tensor>,
142    pub use_kv_cache: bool,
143    kvs: Vec<Option<(Tensor, Tensor)>>,
144    cos: Tensor,
145    sin: Tensor,
146    device: Device,
147}
148
149fn calculate_default_inv_freq(cfg: &GraniteMoeHybridInternalConfig) -> Vec<f32> {
150    let head_dim = cfg.hidden_size / cfg.num_attention_heads;
151    (0..head_dim)
152        .step_by(2)
153        .map(|i| 1f32 / cfg.rope_theta.powf(i as f32 / head_dim as f32))
154        .collect()
155}
156
157impl GraniteMoeHybridCache {
158    pub fn new(
159        use_kv_cache: bool,
160        dtype: DType,
161        config: &GraniteMoeHybridInternalConfig,
162        device: &Device,
163    ) -> Result<Self> {
164        // precompute freqs_cis
165        let theta = match &config.rope_scaling {
166            None
167            | Some(GraniteMoeHybridRopeConfig {
168                rope_type: GraniteMoeHybridRopeType::Default,
169                ..
170            }) => calculate_default_inv_freq(config),
171            Some(rope_scaling) => {
172                let low_freq_wavelen = rope_scaling.original_max_position_embeddings as f32
173                    / rope_scaling.low_freq_factor;
174                let high_freq_wavelen = rope_scaling.original_max_position_embeddings as f32
175                    / rope_scaling.high_freq_factor;
176
177                calculate_default_inv_freq(config)
178                    .into_iter()
179                    .map(|freq| {
180                        let wavelen = 2. * PI / freq;
181                        if wavelen < high_freq_wavelen {
182                            freq
183                        } else if wavelen > low_freq_wavelen {
184                            freq / rope_scaling.factor
185                        } else {
186                            let smooth = (rope_scaling.original_max_position_embeddings as f32
187                                / wavelen
188                                - rope_scaling.low_freq_factor)
189                                / (rope_scaling.high_freq_factor - rope_scaling.low_freq_factor);
190                            (1. - smooth) * freq / rope_scaling.factor + smooth * freq
191                        }
192                    })
193                    .collect::<Vec<_>>()
194            }
195        };
196
197        let theta = Tensor::new(theta, device)?;
198
199        let idx_theta = Tensor::arange(0, config.max_position_embeddings as u32, device)?
200            .to_dtype(DType::F32)?
201            .reshape((config.max_position_embeddings, 1))?
202            .matmul(&theta.reshape((1, theta.elem_count()))?)?;
203        let cos = idx_theta.cos()?.to_dtype(dtype)?;
204        let sin = idx_theta.sin()?.to_dtype(dtype)?;
205        Ok(Self {
206            masks: HashMap::new(),
207            use_kv_cache,
208            kvs: vec![None; config.num_hidden_layers],
209            device: device.clone(),
210            cos,
211            sin,
212        })
213    }
214
215    fn mask(&mut self, seq_len: usize, index_pos: usize) -> Result<Tensor> {
216        let kv_len = index_pos + seq_len;
217        if let Some(mask) = self.masks.get(&(seq_len, kv_len)) {
218            Ok(mask.clone())
219        } else {
220            let mask = crate::utils::build_causal_mask(seq_len, index_pos, &self.device)?;
221            self.masks.insert((seq_len, kv_len), mask.clone());
222            Ok(mask)
223        }
224    }
225}
226
227#[derive(Debug, Clone)]
228struct CausalSelfAttention {
229    q_proj: Linear,
230    k_proj: Linear,
231    v_proj: Linear,
232    o_proj: Linear,
233    num_attention_heads: usize,
234    num_key_value_heads: usize,
235    head_dim: usize,
236    use_flash_attn: bool,
237    span: tracing::Span,
238    span_rot: tracing::Span,
239    max_position_embeddings: usize,
240    attention_multiplier: f32,
241}
242
243#[cfg(feature = "flash-attn")]
244fn flash_attn(
245    q: &Tensor,
246    k: &Tensor,
247    v: &Tensor,
248    softmax_scale: f32,
249    causal: bool,
250) -> Result<Tensor> {
251    candle_flash_attn::flash_attn(q, k, v, softmax_scale, causal)
252}
253
254#[cfg(not(feature = "flash-attn"))]
255fn flash_attn(_: &Tensor, _: &Tensor, _: &Tensor, _: f32, _: bool) -> Result<Tensor> {
256    unimplemented!("compile with '--features flash-attn'")
257}
258
259impl CausalSelfAttention {
260    fn apply_rotary_emb(
261        &self,
262        x: &Tensor,
263        index_pos: usize,
264        cache: &GraniteMoeHybridCache,
265    ) -> Result<Tensor> {
266        let _enter = self.span_rot.enter();
267        let (_b_sz, _, seq_len, _hidden_size) = x.dims4()?;
268        let cos = cache.cos.narrow(0, index_pos, seq_len)?;
269        let sin = cache.sin.narrow(0, index_pos, seq_len)?;
270        candle_nn::rotary_emb::rope(x, &cos, &sin)
271    }
272
273    fn forward(
274        &self,
275        x: &Tensor,
276        index_pos: usize,
277        block_idx: usize,
278        cache: &mut GraniteMoeHybridCache,
279    ) -> Result<Tensor> {
280        let _enter = self.span.enter();
281        let (b_sz, seq_len, hidden_size) = x.dims3()?;
282        let q = self.q_proj.forward(x)?;
283        let k = self.k_proj.forward(x)?;
284        let v = self.v_proj.forward(x)?;
285
286        let q = q
287            .reshape((b_sz, seq_len, self.num_attention_heads, self.head_dim))?
288            .transpose(1, 2)?
289            .contiguous()?;
290        let k = k
291            .reshape((b_sz, seq_len, self.num_key_value_heads, self.head_dim))?
292            .transpose(1, 2)?
293            .contiguous()?;
294        let mut v = v
295            .reshape((b_sz, seq_len, self.num_key_value_heads, self.head_dim))?
296            .transpose(1, 2)?;
297
298        let q = self.apply_rotary_emb(&q, index_pos, cache)?;
299        let mut k = self.apply_rotary_emb(&k, index_pos, cache)?;
300
301        if cache.use_kv_cache {
302            if let Some((cache_k, cache_v)) = &cache.kvs[block_idx] {
303                k = Tensor::cat(&[cache_k, &k], 2)?.contiguous()?;
304                v = Tensor::cat(&[cache_v, &v], 2)?.contiguous()?;
305                let k_seq_len = k.dims()[1];
306                if k_seq_len > self.max_position_embeddings {
307                    k = k
308                        .narrow(
309                            D::Minus1,
310                            k_seq_len - self.max_position_embeddings,
311                            self.max_position_embeddings,
312                        )?
313                        .contiguous()?
314                }
315                let v_seq_len = v.dims()[1];
316                if v_seq_len > 2 * self.max_position_embeddings {
317                    v = v
318                        .narrow(
319                            D::Minus1,
320                            v_seq_len - self.max_position_embeddings,
321                            self.max_position_embeddings,
322                        )?
323                        .contiguous()?
324                }
325            }
326            cache.kvs[block_idx] = Some((k.clone(), v.clone()))
327        }
328
329        let k = self.repeat_kv(k)?;
330        let v = self.repeat_kv(v)?;
331
332        let y = if self.use_flash_attn {
333            // flash-attn expects (b_sz, seq_len, nheads, head_dim)
334            let q = q.transpose(1, 2)?;
335            let k = k.transpose(1, 2)?;
336            let v = v.transpose(1, 2)?;
337            flash_attn(&q, &k, &v, self.attention_multiplier, seq_len > 1)?.transpose(1, 2)?
338        } else {
339            let in_dtype = q.dtype();
340            let q = q.to_dtype(DType::F32)?;
341            let k = k.to_dtype(DType::F32)?;
342            let v = v.to_dtype(DType::F32)?;
343            let att = q
344                .matmul(&k.t()?)?
345                .affine(self.attention_multiplier as f64, 0.)?;
346            let att = if seq_len == 1 {
347                att
348            } else {
349                let mask = cache.mask(seq_len, index_pos)?.broadcast_as(att.shape())?;
350                masked_fill(&att, &mask, f32::NEG_INFINITY)?
351            };
352            let att = candle_nn::ops::softmax(&att, D::Minus1)?;
353            // Convert to contiguous as matmul doesn't support strided vs for now.
354            att.matmul(&v.contiguous()?)?.to_dtype(in_dtype)?
355        };
356        let y = y.transpose(1, 2)?.reshape(&[b_sz, seq_len, hidden_size])?;
357        let y = self.o_proj.forward(&y)?;
358        Ok(y)
359    }
360
361    fn repeat_kv(&self, x: Tensor) -> Result<Tensor> {
362        crate::utils::repeat_kv(x, self.num_attention_heads / self.num_key_value_heads)
363    }
364
365    fn load(vb: VarBuilder, cfg: &GraniteMoeHybridInternalConfig) -> Result<Self> {
366        let span = tracing::span!(tracing::Level::TRACE, "attn");
367        let span_rot = tracing::span!(tracing::Level::TRACE, "attn-rot");
368        let size_in = cfg.hidden_size;
369        let size_q = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_attention_heads;
370        let size_kv = (cfg.hidden_size / cfg.num_attention_heads) * cfg.num_key_value_heads;
371        let q_proj = linear(size_in, size_q, vb.pp("q_proj"))?;
372        let k_proj = linear(size_in, size_kv, vb.pp("k_proj"))?;
373        let v_proj = linear(size_in, size_kv, vb.pp("v_proj"))?;
374        let o_proj = linear(size_q, size_in, vb.pp("o_proj"))?;
375        Ok(Self {
376            q_proj,
377            k_proj,
378            v_proj,
379            o_proj,
380            num_attention_heads: cfg.num_attention_heads,
381            num_key_value_heads: cfg.num_key_value_heads,
382            head_dim: cfg.hidden_size / cfg.num_attention_heads,
383            use_flash_attn: cfg.use_flash_attn,
384            span,
385            span_rot,
386            max_position_embeddings: cfg.max_position_embeddings,
387            attention_multiplier: cfg.attention_multiplier,
388        })
389    }
390}
391
392/// Utility function to fill elements of a tensor based on a boolean mask.
393fn masked_fill(on_false: &Tensor, mask: &Tensor, on_true: f32) -> Result<Tensor> {
394    let shape = mask.shape();
395    let on_true = Tensor::new(on_true, on_false.device())?.broadcast_as(shape.dims())?;
396    let m = mask.where_cond(&on_true, on_false)?;
397    Ok(m)
398}
399
400// A simple feed forward network with a gated activation
401// (GeLU, SiLU, etc.). The goal is to add non-linearity and
402// increase the model's capacity to learn complex patterns.
403#[derive(Debug, Clone)]
404struct MultiLayerPercepton {
405    input_linear: Linear,
406    output_linear: Linear,
407    span: tracing::Span,
408}
409
410impl MultiLayerPercepton {
411    fn forward(&self, x: &Tensor) -> Result<Tensor> {
412        let _enter = self.span.enter();
413        let projected = self.input_linear.forward(x)?;
414        let chunks = projected.chunk(2, D::Minus1)?;
415        let (left, right) = (&chunks[0], &chunks[1]);
416        let gated = (candle_nn::ops::silu(left)? * right)?;
417        self.output_linear.forward(&gated)
418    }
419
420    fn load(vb: VarBuilder, cfg: &GraniteMoeHybridInternalConfig) -> Result<Self> {
421        let span = tracing::span!(tracing::Level::TRACE, "mlp");
422        let h_size = cfg.hidden_size;
423        let inter_size = cfg.shared_intermediate_size;
424        let input_linear = linear(h_size, inter_size * 2, vb.pp("shared_mlp.input_linear"))?;
425        let output_linear = linear(inter_size, h_size, vb.pp("shared_mlp.output_linear"))?;
426        Ok(Self {
427            input_linear,
428            output_linear,
429            span,
430        })
431    }
432}
433
434// A Block is a actually a Transformer layer, consisting of
435// a self-attention mechanism followed by a feed-forward neural network (MLP).
436#[derive(Debug, Clone)]
437struct Block {
438    rms_1: RmsNorm,
439    attn: CausalSelfAttention,
440    rms_2: RmsNorm,
441    multi_layer_percepton: MultiLayerPercepton,
442    span: tracing::Span,
443    residual_scale: f32,
444}
445
446impl Block {
447    fn forward(
448        &self,
449        x: &Tensor,
450        index_pos: usize,
451        block_idx: usize,
452        cache: &mut GraniteMoeHybridCache,
453    ) -> Result<Tensor> {
454        let _enter = self.span.enter();
455        let residual = x;
456        let x = self.rms_1.forward(x)?;
457        let attn = self.attn.forward(&x, index_pos, block_idx, cache)?;
458        let attn = scale_tensor(attn, self.residual_scale)?;
459        let x = (attn + residual)?;
460        let residual = &x;
461        let multi_layer_percepton_out = self
462            .multi_layer_percepton
463            .forward(&self.rms_2.forward(&x)?)?;
464        let multi_layer_percepton_out =
465            scale_tensor(multi_layer_percepton_out, self.residual_scale)?;
466        let x = (multi_layer_percepton_out + residual)?;
467        Ok(x)
468    }
469
470    fn load(vb: VarBuilder, cfg: &GraniteMoeHybridInternalConfig) -> Result<Self> {
471        let span = tracing::span!(tracing::Level::TRACE, "block");
472        let attn = CausalSelfAttention::load(vb.pp("self_attn"), cfg)?;
473        let multi_layer_percepton = MultiLayerPercepton::load(vb.clone(), cfg)?;
474        let rms_1 = RmsNorm::new(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("input_layernorm"))?;
475        let rms_2 = RmsNorm::new(
476            cfg.hidden_size,
477            cfg.rms_norm_eps,
478            vb.pp("post_attention_layernorm"),
479        )?;
480        Ok(Self {
481            rms_1,
482            attn,
483            rms_2,
484            multi_layer_percepton,
485            span,
486            residual_scale: cfg.residual_multiplier,
487        })
488    }
489}
490
491#[derive(Debug, Clone)]
492pub struct GraniteMoeHybrid {
493    word_token_embedding: Embedding,
494    blocks: Vec<Block>,
495    ln_f: RmsNorm,
496    logits_scale: f32,
497    embedding_scale: f32,
498}
499
500impl GraniteMoeHybrid {
501    pub fn forward(
502        &self,
503        x: &Tensor,
504        index_pos: usize,
505        cache: &mut GraniteMoeHybridCache,
506    ) -> Result<Tensor> {
507        let (_b_sz, seq_len) = x.dims2()?;
508        let x = self.word_token_embedding.forward(x)?;
509        let x = scale_tensor(x, self.embedding_scale)?;
510        let x = self
511            .blocks
512            .iter()
513            .enumerate()
514            .try_fold(x, |x, (block_idx, block)| {
515                block.forward(&x, index_pos, block_idx, cache)
516            })?;
517        // Final normalization
518        let x = self.ln_f.forward(&x)?;
519        let x = x.i((.., seq_len - 1, ..))?.contiguous()?;
520        // Project to vocabulary size
521        let logits = x.matmul(&self.word_token_embedding.embeddings().t()?)?;
522        let logits = logits.to_dtype(DType::F32)?;
523        // Scale the logits if needed (that's also different from Granite 1)
524        let scaled_logits = if (self.logits_scale - 1.0).abs() < f32::EPSILON {
525            logits
526        } else {
527            logits.affine(self.logits_scale as f64, 0.)?
528        };
529
530        Ok(scaled_logits)
531    }
532
533    pub fn load(vb: VarBuilder, cfg: &GraniteMoeHybridInternalConfig) -> Result<Self> {
534        let wte = embedding(cfg.vocab_size, cfg.hidden_size, vb.pp("model.embed_tokens"))?;
535        let ln_f = RmsNorm::new(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("model.norm"))?;
536        if cfg.layer_types.len() != cfg.num_hidden_layers {
537            candle::bail!(
538                "layer_types length {} does not match num_hidden_layers {}",
539                cfg.layer_types.len(),
540                cfg.num_hidden_layers
541            );
542        }
543        let blocks = cfg
544            .layer_types
545            .iter()
546            .enumerate()
547            .map(|(idx, layer_ty)| match layer_ty {
548                GraniteMoeHybridLayerType::Attention => {
549                    Block::load(vb.pp(format!("model.layers.{idx}")), cfg)
550                }
551                GraniteMoeHybridLayerType::Mamba => {
552                    // TODO: Not supprting Mamba layers (blocks) for now,
553                    // so we only iterate over attention layers.
554                    candle::bail!(
555                        "mamba layers are not yet supported in GraniteMoeHybrid inference"
556                    )
557                }
558            })
559            .collect::<Result<Vec<_>>>()?;
560
561        Ok(Self {
562            word_token_embedding: wte,
563            blocks,
564            ln_f,
565            logits_scale: if cfg.logits_scaling == 0.0 {
566                1.0
567            } else {
568                1.0 / cfg.logits_scaling
569            },
570            embedding_scale: cfg.embedding_multiplier,
571        })
572    }
573}
574
575fn scale_tensor(tensor: Tensor, scale: f32) -> Result<Tensor> {
576    if (scale - 1.0).abs() < f32::EPSILON {
577        Ok(tensor)
578    } else {
579        tensor.affine(scale as f64, 0.)
580    }
581}