1use 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#[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
44fn 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#[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#[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 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#[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#[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#[derive(Debug, Clone)]
213enum KvCache {
214 Normal(candle_nn::kv_cache::KvCache),
215 Rotating(candle_nn::kv_cache::RotatingKvCache),
216}
217
218#[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 = self.q_norm.forward(&q)?;
328 k = self.k_norm.forward(&k)?;
329 let v = v_norm(&v, self.rms_norm_eps)?;
331
332 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#[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
477fn 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#[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}