Skip to main content

candle_transformers/models/gemma4/
text.rs

1//! Gemma 4 text decoder.
2//!
3//! and following the candle gemma3.rs patterns.
4
5use std::sync::Arc;
6
7use candle::{DType, Device, Module, Result, Tensor, D};
8use candle_nn::{linear_b as linear_bias, Activation, Linear, VarBuilder};
9
10use super::config::Gemma4TextConfig;
11
12// ── RmsNorm (Gemma-style with +1 offset) ────────────────────────────────────
13
14#[derive(Debug, Clone)]
15struct RmsNorm {
16    weight: Tensor,
17    eps: f64,
18}
19
20impl RmsNorm {
21    fn new(dim: usize, eps: f64, vb: VarBuilder) -> Result<Self> {
22        let weight = vb.get(dim, "weight")?;
23        Ok(Self { weight, eps })
24    }
25}
26
27impl Module for RmsNorm {
28    fn forward(&self, x: &Tensor) -> Result<Tensor> {
29        let x_dtype = x.dtype();
30        let internal_dtype = match x_dtype {
31            DType::F16 | DType::BF16 => DType::F32,
32            d => d,
33        };
34        let hidden_size = x.dim(D::Minus1)?;
35        let x = x.to_dtype(internal_dtype)?;
36        let norm_x = (x.sqr()?.sum_keepdim(D::Minus1)? / hidden_size as f64)?;
37        let x_normed = x.broadcast_div(&(norm_x + self.eps)?.sqrt()?)?;
38        x_normed
39            .to_dtype(x_dtype)?
40            .broadcast_mul(&(&self.weight + 1.0)?)
41    }
42}
43
44/// Pure RMS normalization without learned weight (used for V norm).
45fn v_norm(v: &Tensor, eps: f64) -> Result<Tensor> {
46    let original_dtype = v.dtype();
47    let v_f32 = v.to_dtype(DType::F32)?;
48    let mean_sq = v_f32.sqr()?.mean_keepdim(D::Minus1)?;
49    let rms = (mean_sq + eps)?.sqrt()?;
50    v_f32.broadcast_div(&rms)?.to_dtype(original_dtype)
51}
52
53// ── RotaryEmbedding (standard, for sliding layers) ──────────────────────────
54
55#[derive(Debug, Clone)]
56struct RotaryEmbedding {
57    sin: Tensor,
58    cos: Tensor,
59}
60
61impl RotaryEmbedding {
62    fn new(
63        dtype: DType,
64        head_dim: usize,
65        rope_theta: f64,
66        max_seq_len: usize,
67        dev: &Device,
68    ) -> Result<Self> {
69        let inv_freq: Vec<_> = (0..head_dim)
70            .step_by(2)
71            .map(|i| 1f32 / rope_theta.powf(i as f64 / head_dim as f64) as f32)
72            .collect();
73        let inv_freq_len = inv_freq.len();
74        let inv_freq = Tensor::from_vec(inv_freq, (1, inv_freq_len), dev)?.to_dtype(dtype)?;
75        let t = Tensor::arange(0u32, max_seq_len as u32, dev)?
76            .to_dtype(dtype)?
77            .reshape((max_seq_len, 1))?;
78        let freqs = t.matmul(&inv_freq)?;
79        Ok(Self {
80            sin: freqs.sin()?,
81            cos: freqs.cos()?,
82        })
83    }
84
85    fn apply_rotary_emb_qkv(
86        &self,
87        q: &Tensor,
88        k: &Tensor,
89        seqlen_offset: usize,
90    ) -> Result<(Tensor, Tensor)> {
91        let (_b_sz, _h, seq_len, _n_embd) = q.dims4()?;
92        let cos = self.cos.narrow(0, seqlen_offset, seq_len)?;
93        let sin = self.sin.narrow(0, seqlen_offset, seq_len)?;
94        let q_embed = candle_nn::rotary_emb::rope(&q.contiguous()?, &cos, &sin)?;
95        let k_embed = candle_nn::rotary_emb::rope(&k.contiguous()?, &cos, &sin)?;
96        Ok((q_embed, k_embed))
97    }
98}
99
100// ── ProportionalRotaryEmbedding (for global/full layers) ────────────────────
101
102#[derive(Debug, Clone)]
103struct ProportionalRotaryEmbedding {
104    sin: Tensor,
105    cos: Tensor,
106}
107
108impl ProportionalRotaryEmbedding {
109    fn new(
110        dtype: DType,
111        head_dim: usize,
112        rope_theta: f64,
113        partial_rotary_factor: f64,
114        max_seq_len: usize,
115        dev: &Device,
116    ) -> Result<Self> {
117        let rope_angles = (partial_rotary_factor * head_dim as f64 / 2.0) as usize;
118        let half_dim = head_dim / 2;
119
120        let mut inv_freq_vec = Vec::with_capacity(half_dim);
121        for i in 0..rope_angles {
122            inv_freq_vec.push(1f32 / (rope_theta as f32).powf((2 * i) as f32 / head_dim as f32));
123        }
124        // Pad with zeros for non-rotated dimensions -> cos=1, sin=0 -> identity
125        inv_freq_vec.extend(std::iter::repeat_n(0f32, half_dim - rope_angles));
126
127        let inv_freq = Tensor::from_vec(inv_freq_vec, (1, half_dim), dev)?;
128        let t = Tensor::arange(0u32, max_seq_len as u32, dev)?
129            .to_dtype(DType::F32)?
130            .reshape((max_seq_len, 1))?;
131        let freqs = t.matmul(&inv_freq)?;
132        let cos = freqs.cos()?.to_dtype(dtype)?;
133        let sin = freqs.sin()?.to_dtype(dtype)?;
134
135        Ok(Self { cos, sin })
136    }
137
138    fn apply_rotary_emb_qkv(
139        &self,
140        q: &Tensor,
141        k: &Tensor,
142        seqlen_offset: usize,
143    ) -> Result<(Tensor, Tensor)> {
144        let (_b_sz, _h, seq_len, _n_embd) = q.dims4()?;
145        let cos = self.cos.narrow(0, seqlen_offset, seq_len)?;
146        let sin = self.sin.narrow(0, seqlen_offset, seq_len)?;
147        let q_embed = candle_nn::rotary_emb::rope(&q.contiguous()?, &cos, &sin)?;
148        let k_embed = candle_nn::rotary_emb::rope(&k.contiguous()?, &cos, &sin)?;
149        Ok((q_embed, k_embed))
150    }
151}
152
153// ── MLP ─────────────────────────────────────────────────────────────────────
154
155#[derive(Debug, Clone)]
156#[allow(clippy::upper_case_acronyms)]
157struct MLP {
158    gate_proj: Linear,
159    up_proj: Linear,
160    down_proj: Linear,
161    act_fn: Activation,
162}
163
164impl MLP {
165    fn new(
166        hidden_size: usize,
167        intermediate_size: usize,
168        act: Activation,
169        bias: bool,
170        vb: VarBuilder,
171    ) -> Result<Self> {
172        let gate_proj = linear_bias(hidden_size, intermediate_size, bias, vb.pp("gate_proj"))?;
173        let up_proj = linear_bias(hidden_size, intermediate_size, bias, vb.pp("up_proj"))?;
174        let down_proj = linear_bias(intermediate_size, hidden_size, bias, vb.pp("down_proj"))?;
175        Ok(Self {
176            gate_proj,
177            up_proj,
178            down_proj,
179            act_fn: act,
180        })
181    }
182}
183
184impl Module for MLP {
185    fn forward(&self, xs: &Tensor) -> Result<Tensor> {
186        let lhs = xs.apply(&self.gate_proj)?.apply(&self.act_fn)?;
187        let rhs = xs.apply(&self.up_proj)?;
188        (lhs * rhs)?.apply(&self.down_proj)
189    }
190}
191
192// ── Flash attention ─────────────────────────────────────────────────────────
193
194#[cfg(feature = "flash-attn")]
195fn flash_attn(
196    q: &Tensor,
197    k: &Tensor,
198    v: &Tensor,
199    softmax_scale: f32,
200    causal: bool,
201) -> Result<Tensor> {
202    candle_flash_attn::flash_attn(q, k, v, softmax_scale, causal)
203}
204
205#[cfg(not(feature = "flash-attn"))]
206fn flash_attn(_: &Tensor, _: &Tensor, _: &Tensor, _: f32, _: bool) -> Result<Tensor> {
207    unimplemented!("compile with '--features flash-attn'")
208}
209
210// ── KvCache ─────────────────────────────────────────────────────────────────
211
212#[derive(Debug, Clone)]
213enum KvCache {
214    Normal(candle_nn::kv_cache::KvCache),
215    Rotating(candle_nn::kv_cache::RotatingKvCache),
216}
217
218// ── Attention ───────────────────────────────────────────────────────────────
219
220#[derive(Debug, Clone)]
221struct Attention {
222    q_proj: Linear,
223    k_proj: Linear,
224    v_proj: Linear,
225    o_proj: Linear,
226    q_norm: RmsNorm,
227    k_norm: RmsNorm,
228    num_heads: usize,
229    num_kv_heads: usize,
230    num_kv_groups: usize,
231    head_dim: usize,
232    rms_norm_eps: f64,
233    is_sliding: bool,
234    rotary_emb_global: Arc<ProportionalRotaryEmbedding>,
235    rotary_emb_local: Arc<RotaryEmbedding>,
236    kv_cache: KvCache,
237    use_flash_attn: bool,
238}
239
240impl Attention {
241    #[allow(clippy::too_many_arguments)]
242    fn new(
243        rotary_emb_global: Arc<ProportionalRotaryEmbedding>,
244        rotary_emb_local: Arc<RotaryEmbedding>,
245        cfg: &Gemma4TextConfig,
246        layer_idx: usize,
247        vb: VarBuilder,
248    ) -> Result<Self> {
249        let hidden_sz = cfg.hidden_size;
250        let num_heads = cfg.num_attention_heads;
251        let bias = cfg.attention_bias;
252        let is_sliding = cfg.is_sliding(layer_idx);
253
254        let (head_dim, num_kv_heads) = if is_sliding {
255            (cfg.head_dim, cfg.num_key_value_heads)
256        } else {
257            let global_kv = cfg
258                .num_global_key_value_heads
259                .unwrap_or(cfg.num_key_value_heads);
260            (cfg.global_head_dim, global_kv)
261        };
262
263        let num_kv_groups = num_heads / num_kv_heads;
264        let q_proj = linear_bias(hidden_sz, num_heads * head_dim, bias, vb.pp("q_proj"))?;
265        let k_proj = linear_bias(hidden_sz, num_kv_heads * head_dim, bias, vb.pp("k_proj"))?;
266        let v_proj = linear_bias(hidden_sz, num_kv_heads * head_dim, bias, vb.pp("v_proj"))?;
267        let o_proj = linear_bias(num_heads * head_dim, hidden_sz, bias, vb.pp("o_proj"))?;
268        let q_norm = RmsNorm::new(head_dim, cfg.rms_norm_eps, vb.pp("q_norm"))?;
269        let k_norm = RmsNorm::new(head_dim, cfg.rms_norm_eps, vb.pp("k_norm"))?;
270
271        let kv_cache = if is_sliding {
272            KvCache::Rotating(candle_nn::kv_cache::RotatingKvCache::new(
273                2,
274                cfg.effective_sliding_window(),
275            ))
276        } else {
277            KvCache::Normal(candle_nn::kv_cache::KvCache::new(
278                2,
279                cfg.max_position_embeddings,
280            ))
281        };
282
283        Ok(Self {
284            q_proj,
285            k_proj,
286            v_proj,
287            o_proj,
288            q_norm,
289            k_norm,
290            num_heads,
291            num_kv_heads,
292            num_kv_groups,
293            head_dim,
294            rms_norm_eps: cfg.rms_norm_eps,
295            is_sliding,
296            rotary_emb_global,
297            rotary_emb_local,
298            kv_cache,
299            use_flash_attn: cfg.use_flash_attn,
300        })
301    }
302
303    fn forward(
304        &mut self,
305        xs: &Tensor,
306        attention_mask: Option<&Tensor>,
307        sliding_attention_mask: Option<&Tensor>,
308        seqlen_offset: usize,
309    ) -> Result<Tensor> {
310        let (b_sz, q_len, _) = xs.dims3()?;
311
312        let mut q = self.q_proj.forward(xs)?;
313        let mut k = self.k_proj.forward(xs)?;
314        let v = self.v_proj.forward(xs)?;
315
316        q = q
317            .reshape((b_sz, q_len, self.num_heads, self.head_dim))?
318            .transpose(1, 2)?;
319        k = k
320            .reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))?
321            .transpose(1, 2)?;
322        let v = v
323            .reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))?
324            .transpose(1, 2)?;
325
326        // Q/K norms
327        q = self.q_norm.forward(&q)?;
328        k = self.k_norm.forward(&k)?;
329        // V norm (RMS without learned weight)
330        let v = v_norm(&v, self.rms_norm_eps)?;
331
332        // Apply RoPE
333        let (q, k) = if self.is_sliding {
334            self.rotary_emb_local
335                .apply_rotary_emb_qkv(&q, &k, seqlen_offset)?
336        } else {
337            self.rotary_emb_global
338                .apply_rotary_emb_qkv(&q, &k, seqlen_offset)?
339        };
340
341        let (k, v) = match &mut self.kv_cache {
342            KvCache::Normal(cache) => cache.append(&k, &v)?,
343            KvCache::Rotating(cache) => cache.append(&k, &v)?,
344        };
345
346        let k = crate::utils::repeat_kv(k, self.num_kv_groups)?.contiguous()?;
347        let v = crate::utils::repeat_kv(v, self.num_kv_groups)?.contiguous()?;
348
349        let mask = if self.is_sliding {
350            sliding_attention_mask
351        } else {
352            attention_mask
353        };
354
355        let attn_output = if self.use_flash_attn {
356            let q = q.transpose(1, 2)?;
357            let k = k.transpose(1, 2)?;
358            let v = v.transpose(1, 2)?;
359            let scale = 1f32 / (self.head_dim as f32).sqrt();
360            flash_attn(&q, &k, &v, scale, mask.is_some())?.transpose(1, 2)?
361        } else {
362            let scale = 1f64 / f64::sqrt(self.head_dim as f64);
363            let attn_weights = (q.matmul(&k.transpose(2, 3)?)? * scale)?;
364
365            let attn_weights = match mask {
366                None => attn_weights,
367                Some(mask) => attn_weights.broadcast_add(mask)?,
368            };
369            let attn_weights = candle_nn::ops::softmax_last_dim(&attn_weights)?;
370            attn_weights.matmul(&v)?
371        };
372        attn_output
373            .transpose(1, 2)?
374            .reshape((b_sz, q_len, ()))?
375            .apply(&self.o_proj)
376    }
377
378    fn clear_kv_cache(&mut self) {
379        match &mut self.kv_cache {
380            KvCache::Normal(c) => c.reset(),
381            KvCache::Rotating(c) => c.reset(),
382        }
383    }
384}
385
386// ── DecoderLayer ────────────────────────────────────────────────────────────
387
388#[derive(Debug, Clone)]
389struct DecoderLayer {
390    self_attn: Attention,
391    mlp: MLP,
392    input_layernorm: RmsNorm,
393    post_attention_layernorm: RmsNorm,
394    pre_feedforward_layernorm: RmsNorm,
395    post_feedforward_layernorm: RmsNorm,
396    #[allow(dead_code)]
397    is_sliding: bool,
398}
399
400impl DecoderLayer {
401    fn new(
402        rotary_emb_global: Arc<ProportionalRotaryEmbedding>,
403        rotary_emb_local: Arc<RotaryEmbedding>,
404        cfg: &Gemma4TextConfig,
405        layer_idx: usize,
406        vb: VarBuilder,
407    ) -> Result<Self> {
408        let is_sliding = cfg.is_sliding(layer_idx);
409        let self_attn = Attention::new(
410            rotary_emb_global,
411            rotary_emb_local,
412            cfg,
413            layer_idx,
414            vb.pp("self_attn"),
415        )?;
416        let mlp = MLP::new(
417            cfg.hidden_size,
418            cfg.intermediate_size,
419            cfg.hidden_activation,
420            false,
421            vb.pp("mlp"),
422        )?;
423        let input_layernorm =
424            RmsNorm::new(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("input_layernorm"))?;
425        let post_attention_layernorm = RmsNorm::new(
426            cfg.hidden_size,
427            cfg.rms_norm_eps,
428            vb.pp("post_attention_layernorm"),
429        )?;
430        let pre_feedforward_layernorm = RmsNorm::new(
431            cfg.hidden_size,
432            cfg.rms_norm_eps,
433            vb.pp("pre_feedforward_layernorm"),
434        )?;
435        let post_feedforward_layernorm = RmsNorm::new(
436            cfg.hidden_size,
437            cfg.rms_norm_eps,
438            vb.pp("post_feedforward_layernorm"),
439        )?;
440        Ok(Self {
441            self_attn,
442            mlp,
443            input_layernorm,
444            post_attention_layernorm,
445            pre_feedforward_layernorm,
446            post_feedforward_layernorm,
447            is_sliding,
448        })
449    }
450
451    fn forward(
452        &mut self,
453        xs: &Tensor,
454        attention_mask: Option<&Tensor>,
455        sliding_attention_mask: Option<&Tensor>,
456        seqlen_offset: usize,
457    ) -> Result<Tensor> {
458        let residual = xs;
459        let xs = self.input_layernorm.forward(xs)?;
460        let xs =
461            self.self_attn
462                .forward(&xs, attention_mask, sliding_attention_mask, seqlen_offset)?;
463        let xs = xs.apply(&self.post_attention_layernorm)?;
464        let xs = (xs + residual)?;
465        let residual = &xs;
466        let xs = xs.apply(&self.pre_feedforward_layernorm)?;
467        let xs = xs.apply(&self.mlp)?;
468        let xs = xs.apply(&self.post_feedforward_layernorm)?;
469        residual + xs
470    }
471
472    fn clear_kv_cache(&mut self) {
473        self.self_attn.clear_kv_cache()
474    }
475}
476
477// ── Causal mask ─────────────────────────────────────────────────────────────
478
479fn prepare_decoder_attention_mask(
480    b_size: usize,
481    tgt_len: usize,
482    seqlen_offset: usize,
483    sliding_window: Option<usize>,
484    dtype: DType,
485    device: &Device,
486) -> Result<Tensor> {
487    let mask: Vec<_> = if let Some(sliding_window) = sliding_window {
488        (0..tgt_len)
489            .flat_map(|i| {
490                (0..tgt_len).map(move |j| {
491                    if i < j || j + sliding_window < i {
492                        f32::NEG_INFINITY
493                    } else {
494                        0.
495                    }
496                })
497            })
498            .collect()
499    } else {
500        (0..tgt_len)
501            .flat_map(|i| (0..tgt_len).map(move |j| if i < j { f32::NEG_INFINITY } else { 0f32 }))
502            .collect()
503    };
504    let mask = Tensor::from_slice(&mask, (tgt_len, tgt_len), device)?;
505    let mask = if seqlen_offset > 0 {
506        let mask0 = Tensor::zeros((tgt_len, seqlen_offset), DType::F32, device)?;
507        Tensor::cat(&[&mask0, &mask], D::Minus1)?
508    } else {
509        mask
510    };
511    mask.expand((b_size, 1, tgt_len, tgt_len + seqlen_offset))?
512        .to_dtype(dtype)
513}
514
515// ── TextModel ───────────────────────────────────────────────────────────────
516
517#[derive(Debug, Clone)]
518pub struct TextModel {
519    embed_tokens: candle_nn::Embedding,
520    layers: Vec<DecoderLayer>,
521    norm: RmsNorm,
522    lm_head: Linear,
523    final_logit_softcapping: Option<f64>,
524    device: Device,
525    dtype: DType,
526    hidden_size: usize,
527    sliding_window: usize,
528}
529
530impl TextModel {
531    pub fn new(cfg: &Gemma4TextConfig, vb: VarBuilder) -> Result<Self> {
532        let vb_m = vb.pp("model");
533        let embed_tokens =
534            candle_nn::embedding(cfg.vocab_size, cfg.hidden_size, vb_m.pp("embed_tokens"))?;
535
536        let rotary_emb_global = Arc::new(ProportionalRotaryEmbedding::new(
537            vb.dtype(),
538            cfg.global_head_dim,
539            cfg.rope_theta,
540            cfg.partial_rotary_factor(),
541            cfg.max_position_embeddings,
542            vb_m.device(),
543        )?);
544        let rotary_emb_local = Arc::new(RotaryEmbedding::new(
545            vb.dtype(),
546            cfg.head_dim,
547            cfg.rope_local_base_freq(),
548            cfg.max_position_embeddings,
549            vb_m.device(),
550        )?);
551
552        let mut layers = Vec::with_capacity(cfg.num_hidden_layers);
553        let vb_l = vb_m.pp("layers");
554        for layer_idx in 0..cfg.num_hidden_layers {
555            let layer = DecoderLayer::new(
556                rotary_emb_global.clone(),
557                rotary_emb_local.clone(),
558                cfg,
559                layer_idx,
560                vb_l.pp(layer_idx),
561            )?;
562            layers.push(layer)
563        }
564        let norm = RmsNorm::new(cfg.hidden_size, cfg.rms_norm_eps, vb_m.pp("norm"))?;
565        let lm_head = if cfg.tie_word_embeddings {
566            Linear::new(embed_tokens.embeddings().clone(), None)
567        } else {
568            candle_nn::linear_no_bias(cfg.hidden_size, cfg.vocab_size, vb.pp("lm_head"))?
569        };
570        Ok(Self {
571            embed_tokens,
572            layers,
573            norm,
574            lm_head,
575            final_logit_softcapping: cfg.final_logit_softcapping,
576            device: vb.device().clone(),
577            dtype: vb.dtype(),
578            hidden_size: cfg.hidden_size,
579            sliding_window: cfg.sliding_window,
580        })
581    }
582
583    fn create_attention_masks(
584        &self,
585        batch_size: usize,
586        seq_len: usize,
587        seqlen_offset: usize,
588    ) -> Result<(Option<Tensor>, Option<Tensor>)> {
589        if seq_len <= 1 {
590            return Ok((None, None));
591        }
592        let mask = prepare_decoder_attention_mask(
593            batch_size,
594            seq_len,
595            seqlen_offset,
596            None,
597            self.dtype,
598            &self.device,
599        )?;
600        let sliding_mask = prepare_decoder_attention_mask(
601            batch_size,
602            seq_len,
603            seqlen_offset,
604            Some(self.sliding_window),
605            self.dtype,
606            &self.device,
607        )?;
608        Ok((Some(mask), Some(sliding_mask)))
609    }
610
611    pub fn embed_tokens(&self, input_ids: &Tensor) -> Result<Tensor> {
612        let xs = self.embed_tokens.forward(input_ids)?;
613        xs * (self.hidden_size as f64).sqrt()
614    }
615
616    pub fn forward(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
617        let (b_size, seq_len) = input_ids.dims2()?;
618        let xs = self.embed_tokens(input_ids)?;
619        self.forward_embeds(&xs, seqlen_offset, b_size, seq_len)
620    }
621
622    pub fn forward_embeds(
623        &mut self,
624        xs: &Tensor,
625        seqlen_offset: usize,
626        batch_size: usize,
627        seq_len: usize,
628    ) -> Result<Tensor> {
629        let (attention_mask, sliding_attention_mask) =
630            self.create_attention_masks(batch_size, seq_len, seqlen_offset)?;
631
632        let mut xs = xs.clone();
633        for layer in self.layers.iter_mut() {
634            xs = layer.forward(
635                &xs,
636                attention_mask.as_ref(),
637                sliding_attention_mask.as_ref(),
638                seqlen_offset,
639            )?
640        }
641        let logits = xs
642            .narrow(1, seq_len - 1, 1)?
643            .apply(&self.norm)?
644            .apply(&self.lm_head)?;
645        match self.final_logit_softcapping {
646            None => Ok(logits),
647            Some(sc) => Ok(((logits / sc)?.tanh()? * sc)?),
648        }
649    }
650
651    pub fn clear_kv_cache(&mut self) {
652        for layer in self.layers.iter_mut() {
653            layer.clear_kv_cache()
654        }
655    }
656}