Skip to main content

candle_transformers/models/
glm4_new.rs

1use crate::models::glm4::EosTokenId;
2use crate::{
3    models::with_tracing::{linear_b, linear_no_bias, Linear, RmsNorm},
4    utils::repeat_kv,
5};
6use candle::{DType, Device, IndexOp, Module, Result, Tensor, D};
7use candle_nn::{kv_cache::KvCache, Activation, VarBuilder};
8use std::sync::Arc;
9
10#[derive(Debug, Clone, serde::Deserialize)]
11pub struct Config {
12    pub vocab_size: usize,
13    pub hidden_size: usize,
14    pub intermediate_size: usize,
15    pub num_hidden_layers: usize,
16    pub num_attention_heads: usize,
17    pub head_dim: Option<usize>,
18    pub partial_rotary_factor: Option<f32>,
19    pub attention_bias: Option<bool>,
20    pub num_key_value_heads: usize,
21    pub max_position_embeddings: usize,
22    pub sliding_window: Option<usize>,
23    pub tie_word_embeddings: bool,
24    pub rope_theta: f64,
25    pub rms_norm_eps: f64,
26    pub hidden_act: Activation,
27    pub eos_token_id: Option<EosTokenId>,
28}
29
30#[derive(Debug, Clone)]
31pub(crate) struct RotaryEmbedding {
32    sin: Tensor,
33    cos: Tensor,
34    rotary_dim: usize,
35}
36
37impl RotaryEmbedding {
38    pub(crate) fn new(dtype: DType, cfg: &Config, dev: &Device) -> Result<Self> {
39        let dim = cfg
40            .head_dim
41            .unwrap_or(cfg.hidden_size / cfg.num_attention_heads);
42        let rotary_dim = if let Some(factor) = cfg.partial_rotary_factor {
43            (factor * dim as f32) as usize
44        } else {
45            dim
46        };
47        let max_seq_len = cfg.max_position_embeddings;
48        let inv_freq: Vec<_> = (0..rotary_dim)
49            .step_by(2)
50            .map(|i| 1f32 / cfg.rope_theta.powf(i as f64 / rotary_dim as f64) as f32)
51            .collect();
52        let inv_freq_len = inv_freq.len();
53        let inv_freq = Tensor::from_vec(inv_freq, (1, inv_freq_len), dev)?.to_dtype(dtype)?;
54        let t = Tensor::arange(0u32, max_seq_len as u32, dev)?
55            .to_dtype(dtype)?
56            .reshape((max_seq_len, 1))?;
57        let freqs = t.matmul(&inv_freq)?;
58        Ok(Self {
59            sin: freqs.sin()?,
60            cos: freqs.cos()?,
61            rotary_dim,
62        })
63    }
64
65    pub(crate) fn apply(&self, xs: &Tensor, offset: usize) -> Result<Tensor> {
66        let (_, _, seq_len, _) = xs.dims4()?;
67        let (s, e) = (offset, offset + seq_len);
68        let cos = self.cos.i((s..e, ..))?.contiguous()?;
69        let sin = self.sin.i((s..e, ..))?.contiguous()?;
70        let xs_rot = xs
71            .i((0, .., .., ..self.rotary_dim))?
72            .unsqueeze(0)?
73            .contiguous()?;
74        let xs_pass = xs.i((0, .., .., self.rotary_dim..))?.unsqueeze(0)?;
75        let xs_rot = candle_nn::rotary_emb::rope_i(&xs_rot, &cos, &sin).unwrap();
76        Tensor::cat(&[&xs_rot, &xs_pass], D::Minus1)?.contiguous()
77    }
78}
79
80#[derive(Debug, Clone)]
81pub(crate) struct Mlp {
82    gate_up_proj: Linear,
83    down_proj: Linear,
84    act_fn: Activation,
85}
86
87impl Mlp {
88    pub(crate) fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
89        Ok(Self {
90            gate_up_proj: linear_no_bias(
91                cfg.hidden_size,
92                cfg.intermediate_size * 2,
93                vb.pp("gate_up_proj"),
94            )?,
95            down_proj: linear_no_bias(cfg.intermediate_size, cfg.hidden_size, vb.pp("down_proj"))?,
96            act_fn: cfg.hidden_act,
97        })
98    }
99}
100
101impl Module for Mlp {
102    fn forward(&self, x: &Tensor) -> Result<Tensor> {
103        let w = self.gate_up_proj.forward(x)?;
104        let dim = w.dims().len() - 1;
105        let gate = w.narrow(dim, 0, w.dim(dim)? / 2)?.contiguous()?;
106        let gate = gate.apply(&self.act_fn)?;
107        let up_states = w
108            .narrow(dim, w.dim(dim)? / 2, w.dim(dim)? / 2)?
109            .contiguous()?;
110        self.down_proj.forward(&(gate * up_states)?)
111    }
112}
113
114#[derive(Debug, Clone)]
115pub(crate) struct Attention {
116    q_proj: Linear,
117    k_proj: Linear,
118    v_proj: Linear,
119    o_proj: Linear,
120    num_heads: usize,
121    num_kv_heads: usize,
122    num_kv_groups: usize,
123    head_dim: usize,
124    hidden_size: usize,
125    rotary_emb: Arc<RotaryEmbedding>,
126    kv_cache: KvCache,
127}
128
129impl Attention {
130    pub(crate) fn new(
131        cfg: &Config,
132        rotary_emb: Arc<RotaryEmbedding>,
133        vb: VarBuilder,
134    ) -> Result<Self> {
135        let head_dim = cfg
136            .head_dim
137            .unwrap_or(cfg.hidden_size / cfg.num_attention_heads);
138        let num_heads = cfg.num_attention_heads;
139        let num_kv_heads = cfg.num_key_value_heads;
140        let num_kv_groups = num_heads / num_kv_heads;
141
142        let q_proj = linear_b(
143            cfg.hidden_size,
144            num_heads * head_dim,
145            cfg.attention_bias.unwrap_or(false),
146            vb.pp("q_proj"),
147        )?;
148        let k_proj = linear_b(
149            cfg.hidden_size,
150            num_kv_heads * head_dim,
151            cfg.attention_bias.unwrap_or(false),
152            vb.pp("k_proj"),
153        )?;
154        let v_proj = linear_b(
155            cfg.hidden_size,
156            num_kv_heads * head_dim,
157            cfg.attention_bias.unwrap_or(false),
158            vb.pp("v_proj"),
159        )?;
160        let o_proj = linear_b(
161            num_heads * head_dim,
162            cfg.hidden_size,
163            false,
164            vb.pp("o_proj"),
165        )?;
166
167        // Necessary because the hidden_size in the config isn't always accurate
168        let hidden_size = head_dim * cfg.num_attention_heads;
169
170        // Initialize KV cache with 512 tokens capacity to reduce initial memory allocation.
171        // The cache will grow in chunks of 512 tokens when needed.
172        let kv_cache = KvCache::new(2, 512);
173
174        Ok(Self {
175            q_proj,
176            k_proj,
177            v_proj,
178            o_proj,
179            num_heads,
180            num_kv_heads,
181            num_kv_groups,
182            head_dim,
183            hidden_size,
184            rotary_emb,
185            kv_cache,
186        })
187    }
188
189    pub(crate) fn forward(
190        &mut self,
191        x: &Tensor,
192        attn_mask: Option<&Tensor>,
193        offset: usize,
194    ) -> Result<Tensor> {
195        let (b, l, _) = x.dims3()?;
196
197        let q = self.q_proj.forward(x)?;
198        let k = self.k_proj.forward(x)?;
199        let v = self.v_proj.forward(x)?;
200
201        let q = q
202            .reshape((b, l, self.num_heads, self.head_dim))?
203            .transpose(1, 2)?;
204        let k = k
205            .reshape((b, l, self.num_kv_heads, self.head_dim))?
206            .transpose(1, 2)?;
207        let v = v
208            .reshape((b, l, self.num_kv_heads, self.head_dim))?
209            .transpose(1, 2)?;
210
211        let q = self.rotary_emb.apply(&q, offset)?;
212        let k = self.rotary_emb.apply(&k, offset)?;
213
214        let (k, v) = self.kv_cache.append(&k.contiguous()?, &v.contiguous()?)?;
215
216        let k = repeat_kv(k, self.num_kv_groups)?;
217        let v = repeat_kv(v, self.num_kv_groups)?;
218
219        let scale = 1.0 / (self.head_dim as f64).sqrt();
220        let mut scores = (q.matmul(&k.transpose(2, 3)?)? * scale)?;
221        if let Some(m) = attn_mask {
222            scores = scores.broadcast_add(m)?;
223        }
224        let probs = candle_nn::ops::softmax_last_dim(&scores)?;
225        let ctx = probs.matmul(&v)?;
226
227        ctx.transpose(1, 2)?
228            .reshape((b, l, self.hidden_size))?
229            .apply(&self.o_proj)
230    }
231
232    pub(crate) fn clear_kv_cache(&mut self) {
233        self.kv_cache.reset();
234    }
235}
236
237#[derive(Debug, Clone)]
238struct DecoderLayer {
239    self_attn: Attention,
240    mlp: Mlp,
241    input_layernorm: RmsNorm,
242    post_attention_layernorm: RmsNorm,
243    post_mlp_layernorm: RmsNorm,
244    post_self_attn_layernorm: RmsNorm,
245}
246
247impl DecoderLayer {
248    fn new(cfg: &Config, rotary: Arc<RotaryEmbedding>, vb: VarBuilder) -> Result<Self> {
249        let self_attn = Attention::new(cfg, rotary, vb.pp("self_attn"))?;
250        let mlp = Mlp::new(cfg, vb.pp("mlp"))?;
251
252        let input_layernorm =
253            RmsNorm::new(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("input_layernorm"))?;
254        let post_attention_layernorm = RmsNorm::new(
255            cfg.hidden_size,
256            cfg.rms_norm_eps,
257            vb.pp("post_attention_layernorm"),
258        )?;
259        let post_self_attn_layernorm = RmsNorm::new(
260            cfg.hidden_size,
261            cfg.rms_norm_eps,
262            vb.pp("post_self_attn_layernorm"),
263        )?;
264        let post_mlp_layernorm = RmsNorm::new(
265            cfg.hidden_size,
266            cfg.rms_norm_eps,
267            vb.pp("post_mlp_layernorm"),
268        )?;
269
270        Ok(Self {
271            self_attn,
272            mlp,
273            input_layernorm,
274            post_attention_layernorm,
275            post_self_attn_layernorm,
276            post_mlp_layernorm,
277        })
278    }
279
280    fn forward(&mut self, xs: &Tensor, mask: Option<&Tensor>, offset: usize) -> Result<Tensor> {
281        let residual = xs;
282        let hidden_states = self.input_layernorm.forward(xs)?;
283        let hidden_states = self.self_attn.forward(&hidden_states, mask, offset)?;
284        let hidden_states = self.post_self_attn_layernorm.forward(&hidden_states)?;
285        let hidden_states = (residual + hidden_states)?;
286        let residual = &hidden_states;
287        let hidden_states = self.post_attention_layernorm.forward(&hidden_states)?;
288        let hidden_states = self.mlp.forward(&hidden_states)?;
289        let hidden_states = self.post_mlp_layernorm.forward(&hidden_states)?;
290        residual + hidden_states
291    }
292
293    fn clear_kv_cache(&mut self) {
294        self.self_attn.clear_kv_cache();
295    }
296}
297
298#[derive(Debug, Clone)]
299pub struct Model {
300    embed_tokens: candle_nn::Embedding,
301    layers: Vec<DecoderLayer>,
302    norm: RmsNorm,
303    device: Device,
304    dtype: DType,
305}
306
307impl Model {
308    pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
309        let embed_tokens =
310            candle_nn::embedding(cfg.vocab_size, cfg.hidden_size, vb.pp("model.embed_tokens"))?;
311        let rotary = Arc::new(RotaryEmbedding::new(vb.dtype(), cfg, vb.device())?);
312        let mut layers = Vec::with_capacity(cfg.num_hidden_layers);
313        let vb_l = vb.pp("model.layers");
314        for i in 0..cfg.num_hidden_layers {
315            layers.push(DecoderLayer::new(cfg, rotary.clone(), vb_l.pp(i))?);
316        }
317        Ok(Self {
318            embed_tokens,
319            layers,
320            norm: RmsNorm::new(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("model.norm"))?,
321            device: vb.device().clone(),
322            dtype: vb.dtype(),
323        })
324    }
325
326    fn clear_kv_cache(&mut self) {
327        for l in &mut self.layers {
328            l.clear_kv_cache();
329        }
330    }
331
332    fn causal_mask(
333        &self,
334        b: usize,
335        tgt: usize,
336        offset: usize,
337        sw: Option<usize>,
338    ) -> Result<Tensor> {
339        let minf = f32::NEG_INFINITY;
340        let mask: Vec<_> = (0..tgt)
341            .flat_map(|i| {
342                (0..(tgt + offset)).map(move |j| {
343                    let past_ok = j <= i + offset;
344                    let sw_ok = match sw {
345                        Some(w) => (i + offset) as i64 - j as i64 <= w as i64,
346                        None => true,
347                    };
348                    if past_ok && sw_ok {
349                        0.
350                    } else {
351                        minf
352                    }
353                })
354            })
355            .collect();
356        Tensor::from_slice(&mask, (b, 1, tgt, tgt + offset), &self.device)?.to_dtype(self.dtype)
357    }
358
359    pub fn forward(&mut self, input: &Tensor, offset: usize) -> Result<Tensor> {
360        let (b, l) = input.dims2()?;
361        let mut h = self.embed_tokens.forward(input)?;
362
363        let causal = if l == 1 {
364            None
365        } else {
366            Some(self.causal_mask(b, l, offset, None)?)
367        };
368
369        for layer in &mut self.layers {
370            h = layer.forward(&h, causal.as_ref(), offset)?;
371        }
372        self.norm.forward(&h)
373    }
374}
375
376#[derive(Debug, Clone)]
377pub struct ModelForCausalLM {
378    base: Model,
379    lm_head: Linear,
380}
381
382impl ModelForCausalLM {
383    pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
384        let base = Model::new(cfg, vb.clone())?;
385        let lm_head = if cfg.tie_word_embeddings {
386            Linear::from_weights(base.embed_tokens.embeddings().clone(), None)
387        } else {
388            linear_no_bias(cfg.hidden_size, cfg.vocab_size, vb.pp("lm_head"))?
389        };
390        Ok(Self { base, lm_head })
391    }
392
393    pub fn forward(&mut self, input: &Tensor, offset: usize) -> Result<Tensor> {
394        let (_, l) = input.dims2()?;
395        self.base
396            .forward(input, offset)?
397            .narrow(1, l - 1, 1)?
398            .apply(&self.lm_head)
399    }
400
401    pub fn clear_kv_cache(&mut self) {
402        self.base.clear_kv_cache();
403    }
404}