Skip to main content

candle_transformers/models/
glm4.rs

1//! GLM-4 inference implementation.
2//!
3//! An open bilingual language model with 130B parameters.
4//!
5//! Based on implementation from [ChatGLM-6B](https://github.com/THUDM/ChatGLM-6B)
6
7use crate::models::with_tracing::{linear_b as linear, Linear};
8use candle::{DType, Device, IndexOp, Module, Result, Tensor, D};
9use candle_nn::VarBuilder;
10use serde::de::{self, Deserializer, Visitor};
11use serde::Deserialize;
12use std::fmt;
13
14#[derive(Debug, Clone)]
15pub enum EosTokenId {
16    Single(u32),
17    Multiple(Vec<u32>),
18}
19
20impl<'de> Deserialize<'de> for EosTokenId {
21    fn deserialize<D>(deserializer: D) -> std::result::Result<Self, D::Error>
22    where
23        D: Deserializer<'de>,
24    {
25        struct EosTokenIdVisitor;
26
27        impl<'de> Visitor<'de> for EosTokenIdVisitor {
28            type Value = EosTokenId;
29
30            fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
31                formatter.write_str("an integer or a list of integers")
32            }
33
34            fn visit_u64<E>(self, value: u64) -> std::result::Result<Self::Value, E>
35            where
36                E: de::Error,
37            {
38                if value <= u32::MAX as u64 {
39                    Ok(EosTokenId::Single(value as u32))
40                } else {
41                    Err(de::Error::custom("value too large for u32"))
42                }
43            }
44
45            fn visit_seq<A>(self, mut seq: A) -> std::result::Result<Self::Value, A::Error>
46            where
47                A: serde::de::SeqAccess<'de>,
48            {
49                let mut values = Vec::new();
50                while let Some(value) = seq.next_element::<u32>()? {
51                    values.push(value);
52                }
53                Ok(EosTokenId::Multiple(values))
54            }
55        }
56
57        deserializer.deserialize_any(EosTokenIdVisitor)
58    }
59}
60
61fn default_one() -> usize {
62    1
63}
64
65#[derive(Debug, Clone, serde::Deserialize)]
66pub struct Config {
67    pub num_layers: usize,
68    pub padded_vocab_size: usize,
69    pub hidden_size: usize,
70    pub ffn_hidden_size: usize,
71    pub kv_channels: usize,
72    pub num_attention_heads: usize,
73    pub seq_length: usize,
74    pub layernorm_epsilon: f64,
75    pub rmsnorm: bool,
76    pub apply_residual_connection_post_layernorm: bool,
77    pub post_layer_norm: bool,
78    pub add_bias_linear: bool,
79    pub add_qkv_bias: bool,
80    pub bias_dropout_fusion: bool,
81    pub multi_query_attention: bool,
82    pub multi_query_group_num: usize,
83    pub apply_query_key_layer_scaling: bool,
84    pub attention_softmax_in_fp32: bool,
85    pub fp32_residual_connection: bool,
86    #[serde(default = "default_one")]
87    pub rope_ratio: usize,
88    pub eos_token_id: Option<EosTokenId>,
89}
90
91#[derive(Debug, Clone)]
92struct RotaryEmbedding {
93    cache: Tensor,
94}
95
96impl RotaryEmbedding {
97    fn new(cfg: &Config, dtype: DType, dev: &Device) -> Result<Self> {
98        let rotary_dim = cfg.kv_channels;
99        let n_elem = rotary_dim / 2;
100        let base = 10_000f64 * cfg.rope_ratio as f64;
101        let inv_freq: Vec<_> = (0..n_elem)
102            .step_by(2)
103            .map(|i| 1f32 / base.powf(i as f64 / n_elem as f64) as f32)
104            .collect();
105        let inv_freq_len = inv_freq.len();
106        let inv_freq = Tensor::from_vec(inv_freq, (1, inv_freq_len), dev)?.to_dtype(dtype)?;
107        let t = Tensor::arange(0u32, cfg.seq_length as u32, dev)?
108            .to_dtype(dtype)?
109            .reshape((cfg.seq_length, 1))?;
110        let freqs = t.matmul(&inv_freq)?;
111        let cache = Tensor::stack(&[&freqs.cos()?, &freqs.sin()?], D::Minus1)?;
112        Ok(Self { cache })
113    }
114
115    fn apply(&self, xs: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
116        let (seqlen, _b, np, _hn) = xs.dims4()?;
117        let cache = self.cache.narrow(0, seqlen_offset, seqlen)?;
118        let rot_dim = cache.dim(D::Minus2)? * 2;
119        let (xs, xs_pass) = (
120            xs.narrow(D::Minus1, 0, rot_dim)?,
121            xs.narrow(D::Minus1, rot_dim, rot_dim)?,
122        );
123        let xshaped = xs.reshape((seqlen, (), np, rot_dim / 2, 2))?;
124        let cache = cache.reshape((seqlen, (), 1, rot_dim / 2, 2))?;
125        let (xshaped0, xshaped1) = (
126            xshaped.i((.., .., .., .., 0))?,
127            xshaped.i((.., .., .., .., 1))?,
128        );
129        let (cache0, cache1) = (cache.i((.., .., .., .., 0))?, cache.i((.., .., .., .., 1))?);
130        let xs_out = Tensor::stack(
131            &[
132                (xshaped0.broadcast_mul(&cache0)? - xshaped1.broadcast_mul(&cache1)?)?,
133                (xshaped1.broadcast_mul(&cache0)? + xshaped0.broadcast_mul(&cache1)?)?,
134            ],
135            D::Minus1,
136        )?;
137        let xs_out = xs_out.flatten_from(3)?;
138        Tensor::cat(&[xs_out, xs_pass], D::Minus1)
139    }
140}
141
142#[derive(Debug, Clone)]
143struct CoreAttention {
144    coeff: Option<f64>,
145    norm_factor: f64,
146    dtype: DType,
147}
148
149fn masked_fill(on_false: &Tensor, mask: &Tensor, on_true: f32, dtype: DType) -> Result<Tensor> {
150    let shape = mask.shape();
151    let on_true = Tensor::new(on_true, on_false.device())?.broadcast_as(shape.dims())?;
152    let m = mask.where_cond(&on_true.to_dtype(dtype)?, on_false)?;
153    Ok(m)
154}
155
156impl CoreAttention {
157    fn new(layer_number: usize, cfg: &Config, dtype: DType) -> Result<Self> {
158        let norm_factor = (cfg.kv_channels as f64).sqrt();
159        let (norm_factor, coeff) = if cfg.apply_query_key_layer_scaling {
160            let coeff = f64::max(1.0, layer_number as f64);
161            (norm_factor * coeff, Some(coeff))
162        } else {
163            (norm_factor, None)
164        };
165        Ok(Self {
166            coeff,
167            norm_factor,
168            dtype,
169        })
170    }
171
172    fn forward(
173        &self,
174        query_layer: &Tensor,
175        key_layer: &Tensor,
176        value_layer: &Tensor,
177        attention_mask: &Option<Tensor>,
178    ) -> Result<Tensor> {
179        let output_size = (
180            query_layer.dim(1)?, // b
181            query_layer.dim(2)?, // np
182            query_layer.dim(0)?, // sq
183            key_layer.dim(0)?,   // sk
184        );
185        let query_layer =
186            query_layer.reshape((output_size.2, output_size.0 * output_size.1, ()))?;
187        let key_layer = key_layer.reshape((output_size.3, output_size.0 * output_size.1, ()))?;
188        let matmul_result = Tensor::matmul(
189            &query_layer.transpose(0, 1)?.contiguous()?,
190            &key_layer.transpose(0, 1)?.transpose(1, 2)?.contiguous()?,
191        )?;
192        let matmul_result = (matmul_result / self.norm_factor)?.reshape(output_size)?;
193        let matmul_result = match self.coeff {
194            None => matmul_result,
195            Some(coeff) => (matmul_result * coeff)?,
196        };
197        let attention_scores = match attention_mask {
198            Some(mask) => masked_fill(
199                &matmul_result,
200                &mask.broadcast_left((matmul_result.dim(0)?, matmul_result.dim(1)?))?,
201                f32::NEG_INFINITY,
202                self.dtype,
203            )?,
204            None => matmul_result,
205        };
206        let attention_probs = candle_nn::ops::softmax_last_dim(&attention_scores)?;
207
208        let output_size = (
209            value_layer.dim(1)?,
210            value_layer.dim(2)?,
211            query_layer.dim(0)?,
212            value_layer.dim(3)?,
213        );
214        let value_layer =
215            value_layer.reshape((value_layer.dim(0)?, output_size.0 * output_size.1, ()))?;
216        let attention_probs =
217            attention_probs.reshape((output_size.0 * output_size.1, output_size.2, ()))?;
218        let context_layer = Tensor::matmul(
219            &attention_probs.contiguous()?,
220            &value_layer.transpose(0, 1)?.contiguous()?,
221        )?;
222        let context_layer = context_layer.reshape(output_size)?;
223        let context_layer = context_layer.permute((2, 0, 1, 3))?.contiguous()?;
224        context_layer.flatten_from(D::Minus2)
225    }
226}
227
228#[derive(Debug, Clone)]
229struct SelfAttention {
230    query_key_value: Linear,
231    core_attention: CoreAttention,
232    dense: Linear,
233    multi_query_attention: bool,
234    num_attention_heads_per_partition: usize,
235    num_multi_query_groups_per_partition: usize,
236    hidden_size_per_attention_head: usize,
237    kv_cache: Option<(Tensor, Tensor)>,
238}
239
240impl SelfAttention {
241    fn new(layer_number: usize, cfg: &Config, vb: VarBuilder) -> Result<Self> {
242        let projection_size = cfg.kv_channels * cfg.num_attention_heads;
243        let hidden_size_per_attention_head = projection_size / cfg.num_attention_heads;
244        let qkv_hidden_size = if cfg.multi_query_attention {
245            projection_size + 2 * hidden_size_per_attention_head * cfg.multi_query_group_num
246        } else {
247            3 * projection_size
248        };
249        let query_key_value = linear(
250            cfg.hidden_size,
251            qkv_hidden_size,
252            cfg.add_bias_linear || cfg.add_qkv_bias,
253            vb.pp("query_key_value"),
254        )?;
255        let core_attention = CoreAttention::new(layer_number, cfg, vb.dtype())?;
256        let dense = linear(
257            cfg.hidden_size,
258            cfg.hidden_size,
259            cfg.add_bias_linear,
260            vb.pp("dense"),
261        )?;
262        Ok(Self {
263            query_key_value,
264            core_attention,
265            dense,
266            multi_query_attention: cfg.multi_query_attention,
267            num_attention_heads_per_partition: cfg.num_attention_heads,
268            num_multi_query_groups_per_partition: cfg.multi_query_group_num,
269            hidden_size_per_attention_head: cfg.kv_channels,
270            kv_cache: None,
271        })
272    }
273
274    fn reset_kv_cache(&mut self) {
275        self.kv_cache = None
276    }
277
278    fn forward(
279        &mut self,
280        xs: &Tensor,
281        attention_mask: &Option<Tensor>,
282        rotary_emb: &RotaryEmbedding,
283    ) -> Result<Tensor> {
284        let mixed_x_layer = xs.apply(&self.query_key_value)?;
285        if !self.multi_query_attention {
286            candle::bail!("only multi_query_attention=true is supported")
287        }
288        let hpa = self.hidden_size_per_attention_head;
289        let query_layer =
290            mixed_x_layer.narrow(D::Minus1, 0, self.num_attention_heads_per_partition * hpa)?;
291        let key_layer = mixed_x_layer.narrow(
292            D::Minus1,
293            self.num_attention_heads_per_partition * hpa,
294            self.num_multi_query_groups_per_partition * hpa,
295        )?;
296        let value_layer = mixed_x_layer.narrow(
297            D::Minus1,
298            self.num_attention_heads_per_partition * hpa
299                + self.num_multi_query_groups_per_partition * hpa,
300            self.num_multi_query_groups_per_partition * hpa,
301        )?;
302        let query_layer = query_layer.reshape((
303            query_layer.dim(0)?,
304            query_layer.dim(1)?,
305            self.num_attention_heads_per_partition,
306            hpa,
307        ))?;
308        let key_layer = key_layer.reshape((
309            key_layer.dim(0)?,
310            key_layer.dim(1)?,
311            self.num_multi_query_groups_per_partition,
312            hpa,
313        ))?;
314        let value_layer = value_layer.reshape((
315            value_layer.dim(0)?,
316            value_layer.dim(1)?,
317            self.num_multi_query_groups_per_partition,
318            hpa,
319        ))?;
320
321        // Rotary embeddings.
322        let seqlen_offset = match &self.kv_cache {
323            None => 0,
324            Some((prev_k, _)) => prev_k.dim(0)?,
325        };
326        let query_layer = rotary_emb.apply(&query_layer, seqlen_offset)?;
327        let key_layer = rotary_emb.apply(&key_layer, seqlen_offset)?;
328
329        // KV cache.
330        let (key_layer, value_layer) = match &self.kv_cache {
331            None => (key_layer, value_layer),
332            Some((prev_k, prev_v)) => {
333                let k = Tensor::cat(&[prev_k, &key_layer], 0)?;
334                let v = Tensor::cat(&[prev_v, &value_layer], 0)?;
335                (k, v)
336            }
337        };
338        self.kv_cache = Some((key_layer.clone(), value_layer.clone()));
339
340        // Repeat KV.
341        let ratio =
342            self.num_attention_heads_per_partition / self.num_multi_query_groups_per_partition;
343        let key_layer = {
344            let (d0, d1, d2, d3) = key_layer.dims4()?;
345            key_layer
346                .unsqueeze(D::Minus2)?
347                .expand((d0, d1, d2, ratio, d3))?
348                .reshape((
349                    d0,
350                    d1,
351                    self.num_attention_heads_per_partition,
352                    self.hidden_size_per_attention_head,
353                ))?
354        };
355        let value_layer = {
356            let (d0, d1, d2, d3) = value_layer.dims4()?;
357            value_layer
358                .unsqueeze(D::Minus2)?
359                .expand((d0, d1, d2, ratio, d3))?
360                .reshape((
361                    d0,
362                    d1,
363                    self.num_attention_heads_per_partition,
364                    self.hidden_size_per_attention_head,
365                ))?
366        };
367
368        let context_layer =
369            self.core_attention
370                .forward(&query_layer, &key_layer, &value_layer, attention_mask)?;
371        let output = context_layer.apply(&self.dense)?;
372        Ok(output)
373    }
374}
375
376#[allow(clippy::upper_case_acronyms)]
377#[derive(Debug, Clone)]
378struct MLP {
379    dense_h_to_4h: Linear,
380    dense_4h_to_h: Linear,
381}
382
383impl MLP {
384    fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
385        let dense_h_to_4h = linear(
386            cfg.hidden_size,
387            cfg.ffn_hidden_size * 2,
388            cfg.add_bias_linear,
389            vb.pp("dense_h_to_4h"),
390        )?;
391        let dense_4h_to_h = linear(
392            cfg.ffn_hidden_size,
393            cfg.hidden_size,
394            cfg.add_bias_linear,
395            vb.pp("dense_4h_to_h"),
396        )?;
397        Ok(Self {
398            dense_4h_to_h,
399            dense_h_to_4h,
400        })
401    }
402}
403
404impl Module for MLP {
405    fn forward(&self, xs: &Tensor) -> Result<Tensor> {
406        xs.apply(&self.dense_h_to_4h)?
407            .apply(&candle_nn::Activation::Swiglu)?
408            .apply(&self.dense_4h_to_h)
409    }
410}
411
412#[derive(Debug, Clone)]
413struct Block {
414    input_layernorm: candle_nn::LayerNorm,
415    self_attention: SelfAttention,
416    post_attention_layernorm: candle_nn::LayerNorm,
417    mlp: MLP,
418    apply_residual_connection_post_layernorm: bool,
419}
420
421impl Block {
422    fn new(layer_number: usize, cfg: &Config, vb: VarBuilder) -> Result<Self> {
423        let input_layernorm = if cfg.rmsnorm {
424            candle_nn::rms_norm(
425                cfg.hidden_size,
426                cfg.layernorm_epsilon,
427                vb.pp("input_layernorm"),
428            )?
429            .into_inner()
430        } else {
431            candle_nn::layer_norm(
432                cfg.hidden_size,
433                cfg.layernorm_epsilon,
434                vb.pp("input_layernorm"),
435            )?
436        };
437        let post_attention_layernorm = if cfg.rmsnorm {
438            candle_nn::rms_norm(
439                cfg.hidden_size,
440                cfg.layernorm_epsilon,
441                vb.pp("post_attention_layernorm"),
442            )?
443            .into_inner()
444        } else {
445            candle_nn::layer_norm(
446                cfg.hidden_size,
447                cfg.layernorm_epsilon,
448                vb.pp("post_attention_layernorm"),
449            )?
450        };
451        let self_attention = SelfAttention::new(layer_number, cfg, vb.pp("self_attention"))?;
452        let mlp = MLP::new(cfg, vb.pp("mlp"))?;
453        Ok(Self {
454            input_layernorm,
455            self_attention,
456            post_attention_layernorm,
457            mlp,
458            apply_residual_connection_post_layernorm: cfg.apply_residual_connection_post_layernorm,
459        })
460    }
461
462    fn reset_kv_cache(&mut self) {
463        self.self_attention.reset_kv_cache()
464    }
465
466    fn forward(
467        &mut self,
468        xs: &Tensor,
469        attention_mask: &Option<Tensor>,
470        rotary_emb: &RotaryEmbedding,
471    ) -> Result<Tensor> {
472        let layernorm_output = xs.apply(&self.input_layernorm)?;
473        let attention_output =
474            self.self_attention
475                .forward(&layernorm_output, attention_mask, rotary_emb)?;
476        let residual = if self.apply_residual_connection_post_layernorm {
477            &layernorm_output
478        } else {
479            xs
480        };
481        let layernorm_input = (residual + attention_output)?;
482        let layernorm_output = layernorm_input.apply(&self.post_attention_layernorm)?;
483        let mlp_output = layernorm_output.apply(&self.mlp)?;
484        let residual = if self.apply_residual_connection_post_layernorm {
485            &layernorm_output
486        } else {
487            &layernorm_input
488        };
489        mlp_output + residual
490    }
491}
492
493#[derive(Debug, Clone)]
494struct Transformer {
495    layers: Vec<Block>,
496    final_layernorm: Option<candle_nn::LayerNorm>,
497    rotary_emb: RotaryEmbedding,
498}
499
500impl Transformer {
501    fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
502        let vb_l = vb.pp("layers");
503        let mut layers = Vec::with_capacity(cfg.num_layers);
504        for layer_index in 0..cfg.num_layers {
505            let block = Block::new(layer_index + 1, cfg, vb_l.pp(layer_index))?;
506            layers.push(block)
507        }
508        let final_layernorm = if cfg.post_layer_norm {
509            let ln = if cfg.rmsnorm {
510                candle_nn::rms_norm(
511                    cfg.hidden_size,
512                    cfg.layernorm_epsilon,
513                    vb.pp("final_layernorm"),
514                )?
515                .into_inner()
516            } else {
517                candle_nn::layer_norm(
518                    cfg.hidden_size,
519                    cfg.layernorm_epsilon,
520                    vb.pp("final_layernorm"),
521                )?
522            };
523            Some(ln)
524        } else {
525            None
526        };
527        let rotary_emb = RotaryEmbedding::new(cfg, vb.dtype(), vb.device())?;
528        Ok(Self {
529            layers,
530            final_layernorm,
531            rotary_emb,
532        })
533    }
534
535    fn reset_kv_cache(&mut self) {
536        for block in self.layers.iter_mut() {
537            block.reset_kv_cache()
538        }
539    }
540
541    fn forward(&mut self, xs: &Tensor, attention_mask: &Option<Tensor>) -> Result<Tensor> {
542        let mut xs = xs.clone();
543        for block in self.layers.iter_mut() {
544            xs = block.forward(&xs, attention_mask, &self.rotary_emb)?
545        }
546        match self.final_layernorm.as_ref() {
547            None => Ok(xs),
548            Some(ln) => xs.apply(ln),
549        }
550    }
551}
552
553#[derive(Debug, Clone)]
554struct Embedding {
555    word_embeddings: candle_nn::Embedding,
556    fp32_residual_connection: bool,
557}
558
559impl Embedding {
560    fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
561        let word_embeddings = candle_nn::embedding(
562            cfg.padded_vocab_size,
563            cfg.hidden_size,
564            vb.pp("word_embeddings"),
565        )?;
566        Ok(Self {
567            word_embeddings,
568            fp32_residual_connection: cfg.fp32_residual_connection,
569        })
570    }
571}
572
573impl Module for Embedding {
574    fn forward(&self, xs: &Tensor) -> Result<Tensor> {
575        let xs = self.word_embeddings.forward(xs)?.transpose(0, 1)?; // b,s,h -> s,b,h
576        if self.fp32_residual_connection {
577            xs.to_dtype(candle::DType::F32)
578        } else {
579            xs.contiguous()
580        }
581    }
582}
583
584#[derive(Debug, Clone)]
585pub struct Model {
586    embedding: Embedding,
587    encoder: Transformer,
588    output_layer: Linear,
589}
590
591fn get_mask(size: usize, device: &Device) -> Result<Tensor> {
592    let mask: Vec<_> = (0..size)
593        .flat_map(|i| (0..size).map(move |j| u8::from(j > i)))
594        .collect();
595    Tensor::from_slice(&mask, (size, size), device)
596}
597
598impl Model {
599    pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
600        let vb = vb.pp("transformer");
601        let embedding = Embedding::new(cfg, vb.pp("embedding"))?;
602        let encoder = Transformer::new(cfg, vb.pp("encoder"))?;
603        let output_layer = linear(
604            cfg.hidden_size,
605            cfg.padded_vocab_size,
606            false,
607            vb.pp("output_layer"),
608        )?;
609
610        Ok(Self {
611            embedding,
612            encoder,
613            output_layer,
614        })
615    }
616
617    pub fn reset_kv_cache(&mut self) {
618        self.encoder.reset_kv_cache()
619    }
620
621    pub fn forward(&mut self, xs: &Tensor) -> Result<Tensor> {
622        let (_b_size, seq_len) = xs.dims2()?;
623        let input_embeds = xs.apply(&self.embedding)?;
624        let attention_mask = if seq_len <= 1 {
625            None
626        } else {
627            Some(get_mask(seq_len, xs.device())?)
628        };
629        let xs = self.encoder.forward(&input_embeds, &attention_mask)?;
630        let lm_logits = xs.i(seq_len - 1)?.apply(&self.output_layer)?;
631        Ok(lm_logits)
632    }
633}