Skip to main content

runtime/models_v2/
recurrent_gemma.rs

1//! RecurrentGemma Model V2 - Griffin Architecture with Linear Recurrence
2//!
3//! This implements the RecurrentGemma (Griffin) architecture which features:
4//! - Interleaved local attention and linear recurrence layers
5//! - Real-gated Linear Recurrent Unit (RG-LRU)
6//! - Local sliding window attention for global context
7//! - Efficient O(1) state per token during generation
8//!
9//! Supports: RecurrentGemma-2B, RecurrentGemma-9B
10
11use crate::model_config;
12use super::traits::*;
13use anyhow::Result;
14use serde::{Serialize, Deserialize};
15
16/// RecurrentGemma model configuration
17model_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,  // Local attention window
29    lru_width: usize = 0,                 // 0 = auto: hidden_size
30    recurrent_block_ratio: usize = 2,     // 1 attention per N blocks
31    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        // Recurrent layers are all except those at positions divisible by recurrent_block_ratio
64        layer_idx % self.recurrent_block_ratio != 0
65    }
66}
67
68/// Main RecurrentGemma model
69pub 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
78/// Layer type for RecurrentGemma
79pub enum RecurrentGemmaLayerType {
80    Attention(GriffinAttention),
81    Recurrent(GriffinRecurrent),
82}
83
84/// RecurrentGemma layer
85pub 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
95/// Griffin local attention block
96pub 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
108/// Griffin Real-Gated Linear Recurrent Unit (RG-LRU)
109pub struct GriffinRecurrent {
110    // Linear projections
111    linear_x: Tensor,    // [hidden_size, lru_width]
112    linear_y: Tensor,    // [hidden_size, lru_width]
113
114    // Recurrence parameters
115    a_param: Tensor,     // [lru_width] - learnable decay parameter
116    input_gate: Tensor,  // [hidden_size, lru_width]
117    output_proj: Tensor, // [lru_width, hidden_size]
118
119    lru_width: usize,
120}
121
122/// Griffin MLP
123pub struct GriffinMLP {
124    gate_proj: Tensor,
125    up_proj: Tensor,
126    down_proj: Tensor,
127}
128
129/// RecurrentGemma state for generation
130#[derive(Clone)]
131pub struct RecurrentGemmaState {
132    /// LRU state [batch, lru_width]
133    pub lru_state: Tensor,
134    /// Conv state for temporal convolution [batch, lru_width, conv_len]
135    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                // Gemma-style: multiply embeddings by sqrt(hidden_size)
198                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        // Initialize states for recurrent layers
231        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        // Process with full forward for prompt
247        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        // Generation loop
256        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        // Pre-norm for temporal block
385        let normed = ops_fn::rms_norm(hidden_states, &self.input_layernorm, self.config.rms_norm_eps)?;
386
387        // Apply temporal mixing (attention or recurrent)
388        let temporal_out = match &self.layer_type {
389            RecurrentGemmaLayerType::Attention(attn) => attn.forward(&normed)?,
390            RecurrentGemmaLayerType::Recurrent(rec) => rec.forward(&normed)?,
391        };
392
393        // Post-norm and residual
394        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        // MLP block
398        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        // Project Q, K, V
499        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        // Reshape for attention
508        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        // GQA expansion
519        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        // Attention scores
535        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        // Apply local + causal mask
543        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                    // Causal: can't see future
549                    // Local: can only see within window
550                    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        // Reshape back
567        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        // Project x and y
629        let x = ops_fn::matmul(hidden_states, &self.linear_x)?;
630        let y = ops_fn::matmul(hidden_states, &self.linear_y)?;
631
632        // Compute recurrence
633        let x_candle = x.to_candle()?;
634        let y_candle = y.to_candle()?;
635        let a = self.a_param.to_candle()?.neg()?.exp()?; // Decay factor
636
637        // Process sequence
638        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            // h = a * h + (1 - a) * x
646            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            // output = y * h (gating)
650            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        // hidden_states: [batch, 1, hidden_size] or [batch, hidden_size]
662        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        // Project
670        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        // Recurrence update
677        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        // Output
686        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        // Ensure output is 3D
691        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        // Layer 0: 0 % 3 == 0 -> Attention
793        // Layer 1: 1 % 3 != 0 -> Recurrent
794        // Layer 2: 2 % 3 != 0 -> Recurrent
795        // Layer 3: 3 % 3 == 0 -> Attention
796        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}