1use crate::model_config;
12use super::traits::*;
13use anyhow::Result;
14use serde::{Serialize, Deserialize};
15
16model_config!(RecurrentGemmaConfig {
18 vocab_size: usize = 256000,
19 hidden_size: usize = 2560,
20 num_hidden_layers: usize = 26,
21 intermediate_size: usize = 7680,
22 num_attention_heads: usize = 10,
23 num_key_value_heads: usize = 1,
24 head_dim: usize = 256,
25 max_position_embeddings: usize = 8192,
26 rms_norm_eps: f32 = 1e-6,
27 rope_theta: f32 = 10000.0,
28 attention_window_size: usize = 2048, lru_width: usize = 0, recurrent_block_ratio: usize = 2, tie_word_embeddings: bool = true,
32 pad_token_id: i64 = 0,
33 bos_token_id: i64 = 2,
34 eos_token_id: i64 = 1,
35});
36
37impl RecurrentGemmaConfig {
38 pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
39 Self {
40 vocab_size: gguf.vocab_size,
41 hidden_size: gguf.hidden_size,
42 num_hidden_layers: gguf.num_hidden_layers,
43 intermediate_size: gguf.intermediate_size,
44 num_attention_heads: gguf.num_attention_heads,
45 num_key_value_heads: gguf.num_key_value_heads,
46 head_dim: gguf.hidden_size / gguf.num_attention_heads,
47 rms_norm_eps: gguf.rms_norm_eps,
48 rope_theta: gguf.rope_theta,
49 max_position_embeddings: gguf.max_position_embeddings,
50 ..Default::default()
51 }
52 }
53
54 pub fn effective_lru_width(&self) -> usize {
55 if self.lru_width > 0 {
56 self.lru_width
57 } else {
58 self.hidden_size
59 }
60 }
61
62 pub fn is_recurrent_layer(&self, layer_idx: usize) -> bool {
63 layer_idx % self.recurrent_block_ratio != 0
65 }
66}
67
68pub struct RecurrentGemmaModelV2 {
70 config: RecurrentGemmaConfig,
71 device: Device,
72 embed_tokens: Tensor,
73 layers: Vec<RecurrentGemmaLayer>,
74 norm: Tensor,
75 lm_head: Tensor,
76}
77
78pub enum RecurrentGemmaLayerType {
80 Attention(GriffinAttention),
81 Recurrent(GriffinRecurrent),
82}
83
84pub struct RecurrentGemmaLayer {
86 layer_type: RecurrentGemmaLayerType,
87 mlp: GriffinMLP,
88 input_layernorm: Tensor,
89 pre_feedforward_layernorm: Tensor,
90 post_attention_layernorm: Tensor,
91 post_feedforward_layernorm: Tensor,
92 config: RecurrentGemmaConfig,
93}
94
95pub struct GriffinAttention {
97 q_proj: Tensor,
98 k_proj: Tensor,
99 v_proj: Tensor,
100 o_proj: Tensor,
101 num_heads: usize,
102 num_key_value_heads: usize,
103 head_dim: usize,
104 scale: f32,
105 window_size: usize,
106}
107
108pub struct GriffinRecurrent {
110 linear_x: Tensor, linear_y: Tensor, a_param: Tensor, input_gate: Tensor, output_proj: Tensor, lru_width: usize,
120}
121
122pub struct GriffinMLP {
124 gate_proj: Tensor,
125 up_proj: Tensor,
126 down_proj: Tensor,
127}
128
129#[derive(Clone)]
131pub struct RecurrentGemmaState {
132 pub lru_state: Tensor,
134 pub conv_state: Option<Tensor>,
136}
137
138impl Model for RecurrentGemmaModelV2 {
139 type Config = RecurrentGemmaConfig;
140
141 fn new(config: RecurrentGemmaConfig) -> Result<Self> {
142 let device = Device::CPU;
143
144 let embed_tokens = ops_fn::zeros(&[config.vocab_size, config.hidden_size], DataType::Float32, &device)?;
145 let norm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, &device)?;
146
147 let lm_head = if config.tie_word_embeddings {
148 embed_tokens.clone()
149 } else {
150 ops_fn::zeros(&[config.hidden_size, config.vocab_size], DataType::Float32, &device)?
151 };
152
153 let mut layers = Vec::with_capacity(config.num_hidden_layers);
154 for i in 0..config.num_hidden_layers {
155 layers.push(RecurrentGemmaLayer::new(&config, i, &device)?);
156 }
157
158 Ok(Self {
159 config,
160 device,
161 embed_tokens,
162 layers,
163 norm,
164 lm_head,
165 })
166 }
167
168 fn from_weights(config: RecurrentGemmaConfig, weights: ModelWeights) -> Result<Self> {
169 let mut model = Self::new(config)?;
170
171 if let Some(w) = weights.get("model.embed_tokens.weight") {
172 model.embed_tokens = w.clone();
173 }
174
175 if let Some(w) = weights.get("model.norm.weight") {
176 model.norm = w.clone();
177 }
178
179 if !model.config.tie_word_embeddings {
180 if let Some(w) = weights.get("lm_head.weight") {
181 model.lm_head = ops_fn::transpose(w)?;
182 }
183 }
184
185 for (i, layer) in model.layers.iter_mut().enumerate() {
186 layer.load_weights(&weights, i)?;
187 }
188
189 Ok(model)
190 }
191
192 fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
193 match inputs {
194 ModelInputs::Text { input_ids, .. } => {
195 let mut hidden_states = ops_fn::embedding(input_ids, &self.embed_tokens)?;
196
197 let scale = (self.config.hidden_size as f32).sqrt();
199 hidden_states = ops_fn::scale(&hidden_states, scale)?;
200
201 for layer in &self.layers {
202 hidden_states = layer.forward(&hidden_states)?;
203 }
204
205 hidden_states = ops_fn::rms_norm(&hidden_states, &self.norm, self.config.rms_norm_eps)?;
206
207 let logits = if self.config.tie_word_embeddings {
208 let embed_t = ops_fn::transpose(&self.embed_tokens)?;
209 ops_fn::matmul(&hidden_states, &embed_t)?
210 } else {
211 ops_fn::matmul(&hidden_states, &self.lm_head)?
212 };
213
214 Ok(ModelOutputs::Logits {
215 logits,
216 hidden_states: None,
217 })
218 }
219 _ => Err(anyhow::anyhow!("RecurrentGemma only supports text inputs")),
220 }
221 }
222
223 fn generate(&self, prompt: &str, config: &GenerationConfig) -> Result<String> {
224 use crate::tokenizer::Tokenizer;
225 use rand::Rng;
226
227 let tokenizer = Tokenizer::new();
228 let mut tokens: Vec<u32> = tokenizer.encode(prompt);
229
230 let batch_size = 1;
232 let lru_width = self.config.effective_lru_width();
233 let mut layer_states: Vec<Option<RecurrentGemmaState>> = Vec::new();
234
235 for i in 0..self.config.num_hidden_layers {
236 if self.config.is_recurrent_layer(i) {
237 layer_states.push(Some(RecurrentGemmaState {
238 lru_state: ops_fn::zeros(&[batch_size, lru_width], DataType::Float32, &self.device)?,
239 conv_state: None,
240 }));
241 } else {
242 layer_states.push(None);
243 }
244 }
245
246 let input_ids = Tensor::from_i64_slice(
248 &tokens.iter().map(|&t| t as i64).collect::<Vec<_>>(),
249 &[1, tokens.len()],
250 &self.device
251 )?;
252 let inputs = ModelInputs::text(input_ids);
253 let _ = self.forward(&inputs)?;
254
255 for _ in 0..config.max_new_tokens {
257 let last_token = *tokens.last().unwrap_or(&0);
258 let input_tensor = Tensor::from_i64_slice(&[last_token as i64], &[1, 1], &self.device)?;
259
260 let mut hidden = ops_fn::embedding(&input_tensor, &self.embed_tokens)?;
261 let scale = (self.config.hidden_size as f32).sqrt();
262 hidden = ops_fn::scale(&hidden, scale)?;
263
264 for (i, layer) in self.layers.iter().enumerate() {
265 hidden = layer.forward_with_state(&hidden, layer_states[i].as_mut())?;
266 }
267
268 hidden = ops_fn::rms_norm(&hidden, &self.norm, self.config.rms_norm_eps)?;
269
270 let logits = if self.config.tie_word_embeddings {
271 let embed_t = ops_fn::transpose(&self.embed_tokens)?;
272 ops_fn::matmul(&hidden, &embed_t)?
273 } else {
274 ops_fn::matmul(&hidden, &self.lm_head)?
275 };
276
277 let logits_candle = logits.to_candle()?;
278 let last_logits = logits_candle.squeeze(0)?.squeeze(0)?;
279 let logits_vec: Vec<f32> = last_logits.to_vec1()?;
280
281 let next_token = if config.do_sample && config.temperature > 0.0 {
282 let scaled: Vec<f32> = logits_vec.iter().map(|&x| x / config.temperature).collect();
283 let max_val = scaled.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
284 let exp_sum: f32 = scaled.iter().map(|&x| (x - max_val).exp()).sum();
285 let probs: Vec<f32> = scaled.iter().map(|&x| (x - max_val).exp() / exp_sum).collect();
286
287 let mut rng = rand::thread_rng();
288 let random_val: f32 = rng.gen();
289 let mut cumulative = 0.0;
290 let mut sampled = 0u32;
291
292 for (idx, &prob) in probs.iter().enumerate() {
293 cumulative += prob;
294 if random_val <= cumulative {
295 sampled = idx as u32;
296 break;
297 }
298 }
299 sampled
300 } else {
301 logits_vec.iter()
302 .enumerate()
303 .max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap())
304 .map(|(idx, _)| idx as u32)
305 .unwrap_or(0)
306 };
307
308 if next_token == config.eos_token_id {
309 break;
310 }
311
312 tokens.push(next_token);
313 }
314
315 Ok(tokenizer.decode(&tokens))
316 }
317
318 fn config(&self) -> &Self::Config { &self.config }
319
320 fn memory_requirements(&self) -> MemoryRequirements {
321 let param_size = (
322 self.config.vocab_size * self.config.hidden_size +
323 self.config.num_hidden_layers * (
324 4 * self.config.hidden_size * self.config.hidden_size +
325 3 * self.config.hidden_size * self.config.intermediate_size
326 )
327 ) * 4;
328
329 let lru_width = self.config.effective_lru_width();
330 let num_recurrent = self.config.num_hidden_layers * (self.config.recurrent_block_ratio - 1) / self.config.recurrent_block_ratio;
331 let state_size = num_recurrent * lru_width * 4;
332
333 MemoryRequirements {
334 gpu_memory: param_size,
335 cpu_memory: param_size / 4,
336 kv_cache_memory: state_size,
337 peak_memory: param_size + param_size / 2,
338 }
339 }
340
341 fn to_device(&mut self, device: &Device) -> Result<()> {
342 self.device = device.clone();
343 self.embed_tokens = self.embed_tokens.to_device(device)?;
344 self.norm = self.norm.to_device(device)?;
345 if !self.config.tie_word_embeddings {
346 self.lm_head = self.lm_head.to_device(device)?;
347 }
348 for layer in &mut self.layers {
349 layer.to_device(device)?;
350 }
351 Ok(())
352 }
353}
354
355impl RecurrentGemmaLayer {
356 fn new(config: &RecurrentGemmaConfig, layer_idx: usize, device: &Device) -> Result<Self> {
357 let layer_type = if config.is_recurrent_layer(layer_idx) {
358 RecurrentGemmaLayerType::Recurrent(GriffinRecurrent::new(config, device)?)
359 } else {
360 RecurrentGemmaLayerType::Attention(GriffinAttention::new(config, device)?)
361 };
362
363 let mlp = GriffinMLP::new(config, device)?;
364
365 let input_layernorm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
366 let pre_feedforward_layernorm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
367 let post_attention_layernorm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
368 let post_feedforward_layernorm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
369
370 Ok(Self {
371 layer_type,
372 mlp,
373 input_layernorm,
374 pre_feedforward_layernorm,
375 post_attention_layernorm,
376 post_feedforward_layernorm,
377 config: config.clone(),
378 })
379 }
380
381 fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
382 let residual = hidden_states.clone();
383
384 let normed = ops_fn::rms_norm(hidden_states, &self.input_layernorm, self.config.rms_norm_eps)?;
386
387 let temporal_out = match &self.layer_type {
389 RecurrentGemmaLayerType::Attention(attn) => attn.forward(&normed)?,
390 RecurrentGemmaLayerType::Recurrent(rec) => rec.forward(&normed)?,
391 };
392
393 let temporal_out = ops_fn::rms_norm(&temporal_out, &self.post_attention_layernorm, self.config.rms_norm_eps)?;
395 let hidden_states = ops_fn::add(&residual, &temporal_out)?;
396
397 let residual = hidden_states.clone();
399 let normed = ops_fn::rms_norm(&hidden_states, &self.pre_feedforward_layernorm, self.config.rms_norm_eps)?;
400 let mlp_out = self.mlp.forward(&normed)?;
401 let mlp_out = ops_fn::rms_norm(&mlp_out, &self.post_feedforward_layernorm, self.config.rms_norm_eps)?;
402
403 ops_fn::add(&residual, &mlp_out)
404 }
405
406 fn forward_with_state(&self, hidden_states: &Tensor, state: Option<&mut RecurrentGemmaState>) -> Result<Tensor> {
407 let residual = hidden_states.clone();
408
409 let normed = ops_fn::rms_norm(hidden_states, &self.input_layernorm, self.config.rms_norm_eps)?;
410
411 let temporal_out = match (&self.layer_type, state) {
412 (RecurrentGemmaLayerType::Attention(attn), _) => attn.forward(&normed)?,
413 (RecurrentGemmaLayerType::Recurrent(rec), Some(s)) => rec.forward_with_state(&normed, s)?,
414 (RecurrentGemmaLayerType::Recurrent(rec), None) => rec.forward(&normed)?,
415 };
416
417 let temporal_out = ops_fn::rms_norm(&temporal_out, &self.post_attention_layernorm, self.config.rms_norm_eps)?;
418 let hidden_states = ops_fn::add(&residual, &temporal_out)?;
419
420 let residual = hidden_states.clone();
421 let normed = ops_fn::rms_norm(&hidden_states, &self.pre_feedforward_layernorm, self.config.rms_norm_eps)?;
422 let mlp_out = self.mlp.forward(&normed)?;
423 let mlp_out = ops_fn::rms_norm(&mlp_out, &self.post_feedforward_layernorm, self.config.rms_norm_eps)?;
424
425 ops_fn::add(&residual, &mlp_out)
426 }
427
428 fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
429 let prefix = format!("model.layers.{}", layer_idx);
430
431 if let Some(w) = weights.get(&format!("{}.input_layernorm.weight", prefix)) {
432 self.input_layernorm = w.clone();
433 }
434 if let Some(w) = weights.get(&format!("{}.pre_feedforward_layernorm.weight", prefix)) {
435 self.pre_feedforward_layernorm = w.clone();
436 }
437 if let Some(w) = weights.get(&format!("{}.post_attention_layernorm.weight", prefix)) {
438 self.post_attention_layernorm = w.clone();
439 }
440 if let Some(w) = weights.get(&format!("{}.post_feedforward_layernorm.weight", prefix)) {
441 self.post_feedforward_layernorm = w.clone();
442 }
443
444 match &mut self.layer_type {
445 RecurrentGemmaLayerType::Attention(attn) => attn.load_weights(weights, layer_idx)?,
446 RecurrentGemmaLayerType::Recurrent(rec) => rec.load_weights(weights, layer_idx)?,
447 }
448
449 self.mlp.load_weights(weights, layer_idx)?;
450
451 Ok(())
452 }
453
454 fn to_device(&mut self, device: &Device) -> Result<()> {
455 self.input_layernorm = self.input_layernorm.to_device(device)?;
456 self.pre_feedforward_layernorm = self.pre_feedforward_layernorm.to_device(device)?;
457 self.post_attention_layernorm = self.post_attention_layernorm.to_device(device)?;
458 self.post_feedforward_layernorm = self.post_feedforward_layernorm.to_device(device)?;
459
460 match &mut self.layer_type {
461 RecurrentGemmaLayerType::Attention(attn) => attn.to_device(device)?,
462 RecurrentGemmaLayerType::Recurrent(rec) => rec.to_device(device)?,
463 }
464
465 self.mlp.to_device(device)?;
466 Ok(())
467 }
468}
469
470impl GriffinAttention {
471 fn new(config: &RecurrentGemmaConfig, device: &Device) -> Result<Self> {
472 let num_heads = config.num_attention_heads;
473 let num_key_value_heads = config.num_key_value_heads;
474 let head_dim = config.head_dim;
475
476 let q_proj = ops_fn::zeros(&[config.hidden_size, num_heads * head_dim], DataType::Float32, device)?;
477 let k_proj = ops_fn::zeros(&[config.hidden_size, num_key_value_heads * head_dim], DataType::Float32, device)?;
478 let v_proj = ops_fn::zeros(&[config.hidden_size, num_key_value_heads * head_dim], DataType::Float32, device)?;
479 let o_proj = ops_fn::zeros(&[num_heads * head_dim, config.hidden_size], DataType::Float32, device)?;
480
481 Ok(Self {
482 q_proj,
483 k_proj,
484 v_proj,
485 o_proj,
486 num_heads,
487 num_key_value_heads,
488 head_dim,
489 scale: (head_dim as f32).powf(-0.5),
490 window_size: config.attention_window_size,
491 })
492 }
493
494 fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
495 let shape = hidden_states.shape();
496 let (batch_size, seq_len, _) = (shape[0], shape[1], shape[2]);
497
498 let q = ops_fn::matmul(hidden_states, &self.q_proj)?;
500 let k = ops_fn::matmul(hidden_states, &self.k_proj)?;
501 let v = ops_fn::matmul(hidden_states, &self.v_proj)?;
502
503 let q_candle = q.to_candle()?;
504 let k_candle = k.to_candle()?;
505 let v_candle = v.to_candle()?;
506
507 let q_reshaped = q_candle
509 .reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?
510 .transpose(1, 2)?;
511 let k_reshaped = k_candle
512 .reshape(&[batch_size, seq_len, self.num_key_value_heads, self.head_dim])?
513 .transpose(1, 2)?;
514 let v_reshaped = v_candle
515 .reshape(&[batch_size, seq_len, self.num_key_value_heads, self.head_dim])?
516 .transpose(1, 2)?;
517
518 let num_groups = self.num_heads / self.num_key_value_heads;
520 let (k_expanded, v_expanded) = if num_groups > 1 {
521 let k_rep = k_reshaped
522 .unsqueeze(2)?
523 .broadcast_as(&[batch_size, self.num_key_value_heads, num_groups, seq_len, self.head_dim])?
524 .reshape(&[batch_size, self.num_heads, seq_len, self.head_dim])?;
525 let v_rep = v_reshaped
526 .unsqueeze(2)?
527 .broadcast_as(&[batch_size, self.num_key_value_heads, num_groups, seq_len, self.head_dim])?
528 .reshape(&[batch_size, self.num_heads, seq_len, self.head_dim])?;
529 (k_rep, v_rep)
530 } else {
531 (k_reshaped, v_reshaped)
532 };
533
534 let k_t = k_expanded.transpose(2, 3)?;
536 let q_cont = q_reshaped.contiguous()?;
537 let k_cont = k_t.contiguous()?;
538
539 let scores = q_cont.matmul(&k_cont)?;
540 let scaled_scores = (scores * (self.scale as f64))?;
541
542 let device = scaled_scores.device();
544 let mask = {
545 let mut mask_data = vec![0.0f32; seq_len * seq_len];
546 for i in 0..seq_len {
547 for j in 0..seq_len {
548 let is_causal_ok = j <= i;
551 let is_local_ok = i.saturating_sub(self.window_size) <= j;
552 if !is_causal_ok || !is_local_ok {
553 mask_data[i * seq_len + j] = f32::NEG_INFINITY;
554 }
555 }
556 }
557 candle_core::Tensor::from_vec(mask_data, &[1, 1, seq_len, seq_len], device)?
558 };
559
560 let masked_scores = scaled_scores.broadcast_add(&mask)?;
561 let attention_weights = candle_nn::ops::softmax_last_dim(&masked_scores)?;
562
563 let v_cont = v_expanded.contiguous()?;
564 let attn_output = attention_weights.matmul(&v_cont)?;
565
566 let attn_output = attn_output
568 .transpose(1, 2)?
569 .reshape(&[batch_size, seq_len, self.num_heads * self.head_dim])?;
570
571 let attn_output = Tensor::from_candle(attn_output);
572 ops_fn::matmul(&attn_output, &self.o_proj)
573 }
574
575 fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
576 let prefix = format!("model.layers.{}.temporal_block", layer_idx);
577
578 if let Some(w) = weights.get(&format!("{}.q_proj.weight", prefix)) {
579 self.q_proj = ops_fn::transpose(w)?;
580 }
581 if let Some(w) = weights.get(&format!("{}.k_proj.weight", prefix)) {
582 self.k_proj = ops_fn::transpose(w)?;
583 }
584 if let Some(w) = weights.get(&format!("{}.v_proj.weight", prefix)) {
585 self.v_proj = ops_fn::transpose(w)?;
586 }
587 if let Some(w) = weights.get(&format!("{}.o_proj.weight", prefix)) {
588 self.o_proj = ops_fn::transpose(w)?;
589 }
590
591 Ok(())
592 }
593
594 fn to_device(&mut self, device: &Device) -> Result<()> {
595 self.q_proj = self.q_proj.to_device(device)?;
596 self.k_proj = self.k_proj.to_device(device)?;
597 self.v_proj = self.v_proj.to_device(device)?;
598 self.o_proj = self.o_proj.to_device(device)?;
599 Ok(())
600 }
601}
602
603impl GriffinRecurrent {
604 fn new(config: &RecurrentGemmaConfig, device: &Device) -> Result<Self> {
605 let hidden_size = config.hidden_size;
606 let lru_width = config.effective_lru_width();
607
608 let linear_x = ops_fn::zeros(&[hidden_size, lru_width], DataType::Float32, device)?;
609 let linear_y = ops_fn::zeros(&[hidden_size, lru_width], DataType::Float32, device)?;
610 let a_param = ops_fn::zeros(&[lru_width], DataType::Float32, device)?;
611 let input_gate = ops_fn::zeros(&[hidden_size, lru_width], DataType::Float32, device)?;
612 let output_proj = ops_fn::zeros(&[lru_width, hidden_size], DataType::Float32, device)?;
613
614 Ok(Self {
615 linear_x,
616 linear_y,
617 a_param,
618 input_gate,
619 output_proj,
620 lru_width,
621 })
622 }
623
624 fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
625 let shape = hidden_states.shape();
626 let (batch_size, seq_len, _) = (shape[0], shape[1], shape[2]);
627
628 let x = ops_fn::matmul(hidden_states, &self.linear_x)?;
630 let y = ops_fn::matmul(hidden_states, &self.linear_y)?;
631
632 let x_candle = x.to_candle()?;
634 let y_candle = y.to_candle()?;
635 let a = self.a_param.to_candle()?.neg()?.exp()?; let mut h = candle_core::Tensor::zeros(&[batch_size, self.lru_width], candle_core::DType::F32, x_candle.device())?;
639 let mut outputs = Vec::new();
640
641 for t in 0..seq_len {
642 let x_t = x_candle.narrow(1, t, 1)?.squeeze(1)?;
643 let y_t = y_candle.narrow(1, t, 1)?.squeeze(1)?;
644
645 let one_minus_a = candle_core::Tensor::ones_like(&a)?.sub(&a)?;
647 h = a.broadcast_mul(&h)?.add(&one_minus_a.broadcast_mul(&x_t)?)?;
648
649 let out_t = y_t.broadcast_mul(&h)?;
651 outputs.push(out_t);
652 }
653
654 let output = candle_core::Tensor::stack(&outputs, 1)?;
655 let output = Tensor::from_candle(output);
656
657 ops_fn::matmul(&output, &self.output_proj)
658 }
659
660 fn forward_with_state(&self, hidden_states: &Tensor, state: &mut RecurrentGemmaState) -> Result<Tensor> {
661 let x_candle = hidden_states.to_candle()?;
663 let x = if x_candle.dims().len() == 3 {
664 x_candle.squeeze(1)?
665 } else {
666 x_candle.clone()
667 };
668
669 let linear_x = self.linear_x.to_candle()?;
671 let linear_y = self.linear_y.to_candle()?;
672
673 let x_proj = x.matmul(&linear_x)?;
674 let y_proj = x.matmul(&linear_y)?;
675
676 let a = self.a_param.to_candle()?.neg()?.exp()?;
678 let one_minus_a = candle_core::Tensor::ones_like(&a)?.sub(&a)?;
679
680 let h_prev = state.lru_state.to_candle()?;
681 let h_new = a.broadcast_mul(&h_prev)?.add(&one_minus_a.broadcast_mul(&x_proj)?)?;
682
683 state.lru_state = Tensor::from_candle(h_new.clone());
684
685 let out = y_proj.broadcast_mul(&h_new)?;
687 let output_proj = self.output_proj.to_candle()?;
688 let output = out.matmul(&output_proj)?;
689
690 let output = if output.dims().len() == 2 {
692 output.unsqueeze(1)?
693 } else {
694 output
695 };
696
697 Ok(Tensor::from_candle(output))
698 }
699
700 fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
701 let prefix = format!("model.layers.{}.temporal_block", layer_idx);
702
703 if let Some(w) = weights.get(&format!("{}.linear_x.weight", prefix)) {
704 self.linear_x = ops_fn::transpose(w)?;
705 }
706 if let Some(w) = weights.get(&format!("{}.linear_y.weight", prefix)) {
707 self.linear_y = ops_fn::transpose(w)?;
708 }
709 if let Some(w) = weights.get(&format!("{}.a_param", prefix)) {
710 self.a_param = w.clone();
711 }
712 if let Some(w) = weights.get(&format!("{}.input_gate.weight", prefix)) {
713 self.input_gate = ops_fn::transpose(w)?;
714 }
715 if let Some(w) = weights.get(&format!("{}.output_proj.weight", prefix)) {
716 self.output_proj = ops_fn::transpose(w)?;
717 }
718
719 Ok(())
720 }
721
722 fn to_device(&mut self, device: &Device) -> Result<()> {
723 self.linear_x = self.linear_x.to_device(device)?;
724 self.linear_y = self.linear_y.to_device(device)?;
725 self.a_param = self.a_param.to_device(device)?;
726 self.input_gate = self.input_gate.to_device(device)?;
727 self.output_proj = self.output_proj.to_device(device)?;
728 Ok(())
729 }
730}
731
732impl GriffinMLP {
733 fn new(config: &RecurrentGemmaConfig, device: &Device) -> Result<Self> {
734 let gate_proj = ops_fn::zeros(&[config.hidden_size, config.intermediate_size], DataType::Float32, device)?;
735 let up_proj = ops_fn::zeros(&[config.hidden_size, config.intermediate_size], DataType::Float32, device)?;
736 let down_proj = ops_fn::zeros(&[config.intermediate_size, config.hidden_size], DataType::Float32, device)?;
737
738 Ok(Self { gate_proj, up_proj, down_proj })
739 }
740
741 fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
742 let gate = ops_fn::matmul(hidden_states, &self.gate_proj)?;
743 let gate = ops_fn::gelu(&gate)?;
744 let up = ops_fn::matmul(hidden_states, &self.up_proj)?;
745 let hidden = ops_fn::mul(&gate, &up)?;
746 ops_fn::matmul(&hidden, &self.down_proj)
747 }
748
749 fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
750 let prefix = format!("model.layers.{}.mlp", layer_idx);
751
752 if let Some(w) = weights.get(&format!("{}.gate_proj.weight", prefix)) {
753 self.gate_proj = ops_fn::transpose(w)?;
754 }
755 if let Some(w) = weights.get(&format!("{}.up_proj.weight", prefix)) {
756 self.up_proj = ops_fn::transpose(w)?;
757 }
758 if let Some(w) = weights.get(&format!("{}.down_proj.weight", prefix)) {
759 self.down_proj = ops_fn::transpose(w)?;
760 }
761
762 Ok(())
763 }
764
765 fn to_device(&mut self, device: &Device) -> Result<()> {
766 self.gate_proj = self.gate_proj.to_device(device)?;
767 self.up_proj = self.up_proj.to_device(device)?;
768 self.down_proj = self.down_proj.to_device(device)?;
769 Ok(())
770 }
771}
772
773#[cfg(test)]
774mod tests {
775 use super::*;
776
777 #[test]
778 fn test_recurrent_gemma_config() {
779 let config = RecurrentGemmaConfig::default();
780 assert_eq!(config.vocab_size, 256000);
781 assert_eq!(config.hidden_size, 2560);
782 assert_eq!(config.recurrent_block_ratio, 2);
783 }
784
785 #[test]
786 fn test_layer_type_selection() {
787 let config = RecurrentGemmaConfig {
788 recurrent_block_ratio: 3,
789 ..Default::default()
790 };
791
792 assert!(!config.is_recurrent_layer(0));
797 assert!(config.is_recurrent_layer(1));
798 assert!(config.is_recurrent_layer(2));
799 assert!(!config.is_recurrent_layer(3));
800 }
801
802 #[test]
803 fn test_recurrent_gemma_model_creation() {
804 let config = RecurrentGemmaConfig {
805 vocab_size: 1000,
806 hidden_size: 64,
807 intermediate_size: 256,
808 num_hidden_layers: 4,
809 num_attention_heads: 4,
810 num_key_value_heads: 2,
811 head_dim: 16,
812 recurrent_block_ratio: 2,
813 ..Default::default()
814 };
815
816 let model = RecurrentGemmaModelV2::new(config).unwrap();
817 assert_eq!(model.config().vocab_size(), 1000);
818 assert_eq!(model.config().hidden_size(), 64);
819 assert_eq!(model.config().num_layers(), 4);
820 }
821
822 #[test]
823 fn test_recurrent_gemma_forward_pass() {
824 let config = RecurrentGemmaConfig {
825 vocab_size: 100,
826 hidden_size: 32,
827 intermediate_size: 128,
828 num_hidden_layers: 2,
829 num_attention_heads: 2,
830 num_key_value_heads: 1,
831 head_dim: 16,
832 recurrent_block_ratio: 2,
833 ..Default::default()
834 };
835
836 let model = RecurrentGemmaModelV2::new(config).unwrap();
837 let input_ids = ops_fn::zeros(&[1, 4], DataType::Int64, &Device::CPU).unwrap();
838 let inputs = ModelInputs::text(input_ids);
839
840 let outputs = model.forward(&inputs).unwrap();
841 match outputs {
842 ModelOutputs::Logits { logits, .. } => {
843 assert_eq!(logits.shape(), &[1, 4, 100]);
844 }
845 _ => panic!("Expected logits output"),
846 }
847 }
848}