Skip to main content

runtime/models_v2/
gemma.rs

1//! Gemma Model V2 - Clean implementation using solid abstractions
2//!
3//! This implements the Gemma architecture which is used by 10+ models including:
4//! - Gemma-2B, Gemma-7B, Gemma2-9B, Gemma2-27B, CodeGemma, PaliGemma
5//! - Uses unified Tensor type from tensor_core
6//! - Implements Model trait from model_core
7
8use crate::model_config;
9use super::traits::*;
10use anyhow::Result;
11use serde::{Serialize, Deserialize};
12
13/// Gemma model configuration using the model_config macro
14model_config!(GemmaConfig {
15    vocab_size: usize = 256000,
16    hidden_size: usize = 3072,
17    intermediate_size: usize = 24576,
18    num_hidden_layers: usize = 28,
19    num_attention_heads: usize = 16,
20    num_key_value_heads: usize = 16,
21    hidden_act: String = "gelu".to_string(),
22    max_position_embeddings: usize = 8192,
23    initializer_range: f32 = 0.02,
24    rms_norm_eps: f32 = 1e-6,
25    use_cache: bool = true,
26    pad_token_id: i64 = 0,
27    bos_token_id: i64 = 2,
28    eos_token_id: i64 = 1,
29    tie_word_embeddings: bool = true,
30    rope_theta: f32 = 10000.0,
31    attention_bias: bool = false,
32    head_dim: usize = 256,
33});
34
35impl GemmaConfig {
36    /// Create GemmaConfig from GGUF model configuration
37    pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
38        Self {
39            vocab_size: gguf.vocab_size,
40            hidden_size: gguf.hidden_size,
41            intermediate_size: gguf.intermediate_size,
42            num_hidden_layers: gguf.num_hidden_layers,
43            num_attention_heads: gguf.num_attention_heads,
44            num_key_value_heads: gguf.num_key_value_heads,
45            rms_norm_eps: gguf.rms_norm_eps,
46            rope_theta: gguf.rope_theta,
47            max_position_embeddings: gguf.max_position_embeddings,
48            head_dim: gguf.head_dim,
49            ..Default::default()
50        }
51    }
52}
53
54/// Main Gemma model implementation
55pub struct GemmaModelV2 {
56    config: GemmaConfig,
57    device: Device,
58
59    // Model components using unified Tensor type
60    embed_tokens: Tensor,
61    layers: Vec<GemmaLayer>,
62    norm: Tensor,
63    lm_head: Tensor,
64}
65
66/// Gemma transformer layer
67pub struct GemmaLayer {
68    self_attn: GemmaAttention,
69    mlp: GemmaMLP,
70    input_layernorm: Tensor,
71    post_attention_layernorm: Tensor,
72}
73
74/// Gemma attention mechanism
75pub struct GemmaAttention {
76    q_proj: Tensor,
77    k_proj: Tensor,
78    v_proj: Tensor,
79    o_proj: Tensor,
80    num_heads: usize,
81    num_key_value_heads: usize,
82    head_dim: usize,
83    scale: f32,
84}
85
86/// Gemma MLP (feed-forward network)
87pub struct GemmaMLP {
88    gate_proj: Tensor,
89    up_proj: Tensor,
90    down_proj: Tensor,
91    hidden_act: String,
92}
93
94impl Model for GemmaModelV2 {
95    type Config = GemmaConfig;
96
97    fn new(config: GemmaConfig) -> Result<Self> {
98        let device = Device::CPU;
99
100        let embed_tokens = ops_fn::zeros(
101            &[config.vocab_size, config.hidden_size],
102            DataType::Float32,
103            &device
104        )?;
105
106        let norm = ops_fn::zeros(
107            &[config.hidden_size],
108            DataType::Float32,
109            &device
110        )?;
111
112        // Gemma typically uses tied embeddings
113        let lm_head = if config.tie_word_embeddings {
114            embed_tokens.clone()
115        } else {
116            ops_fn::zeros(
117                &[config.hidden_size, config.vocab_size],
118                DataType::Float32,
119                &device
120            )?
121        };
122
123        // Create transformer layers
124        let mut layers = Vec::with_capacity(config.num_hidden_layers);
125        for _ in 0..config.num_hidden_layers {
126            layers.push(GemmaLayer::new(&config, &device)?);
127        }
128
129        Ok(Self {
130            config,
131            device,
132            embed_tokens,
133            layers,
134            norm,
135            lm_head,
136        })
137    }
138
139    fn from_weights(config: GemmaConfig, weights: ModelWeights) -> Result<Self> {
140        let mut model = Self::new(config)?;
141
142        // Load weights from the unified weight container
143        if let Some(embed_weights) = weights.get("model.embed_tokens.weight") {
144            model.embed_tokens = embed_weights.clone();
145        }
146
147        if let Some(norm_weights) = weights.get("model.norm.weight") {
148            model.norm = norm_weights.clone();
149        }
150
151        // For tied embeddings, lm_head uses embed_tokens transposed
152        // For separate lm_head, transpose for matmul
153        if !model.config.tie_word_embeddings {
154            if let Some(lm_head_weights) = weights.get("lm_head.weight") {
155                model.lm_head = ops_fn::transpose(lm_head_weights)?;
156            }
157        }
158
159        // Load layer weights
160        for (i, layer) in model.layers.iter_mut().enumerate() {
161            layer.load_weights(&weights, i)?;
162        }
163
164        Ok(model)
165    }
166
167    fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
168        match inputs {
169            ModelInputs::Text { input_ids, attention_mask, .. } => {
170                // 1. Token embedding
171                let mut hidden_states = ops_fn::embedding(input_ids, &self.embed_tokens)?;
172
173                // 2. Gemma scales embeddings by sqrt(hidden_size)
174                let scale = (self.config.hidden_size as f32).sqrt();
175                hidden_states = ops_fn::scale(&hidden_states, scale)?;
176
177                // 3. Apply transformer layers (with RoPE)
178                for layer in &self.layers {
179                    hidden_states = layer.forward(&hidden_states, attention_mask.as_ref(), self.config.rope_theta, self.config.head_dim)?;
180                }
181
182                // 4. Final layer norm
183                hidden_states = ops_fn::layer_norm(&hidden_states, &self.norm, None, self.config.rms_norm_eps)?;
184
185                // 5. Language modeling head
186                // For tied embeddings, we need to transpose embed_tokens for matmul
187                let logits = if self.config.tie_word_embeddings {
188                    let lm_head_t = ops_fn::transpose(&self.embed_tokens)?;
189                    ops_fn::matmul(&hidden_states, &lm_head_t)?
190                } else {
191                    ops_fn::matmul(&hidden_states, &self.lm_head)?
192                };
193
194                Ok(ModelOutputs::Logits {
195                    logits,
196                    hidden_states: None,
197                })
198            }
199            ModelInputs::Multimodal { input_ids, .. } => {
200                let text_inputs = ModelInputs::Text {
201                    input_ids: input_ids.clone(),
202                    attention_mask: None,
203                    position_ids: None,
204                };
205                self.forward(&text_inputs)
206            }
207            _ => Err(anyhow::anyhow!("Gemma model only supports text and multimodal inputs")),
208        }
209    }
210
211    fn generate(&self, prompt: &str, config: &GenerationConfig) -> Result<String> {
212        use crate::tokenizer::Tokenizer;
213        use rand::Rng;
214
215        let tokenizer = Tokenizer::new();
216        let mut tokens: Vec<u32> = tokenizer.encode(prompt);
217
218        for _ in 0..config.max_new_tokens {
219            let tokens_i64: Vec<i64> = tokens.iter().map(|&t| t as i64).collect();
220            let input_tensor = Tensor::from_i64_slice(&tokens_i64, &[1, tokens.len()], &self.device)?;
221
222            let inputs = ModelInputs::Text {
223                input_ids: input_tensor,
224                attention_mask: None,
225                position_ids: None,
226            };
227
228            let outputs = self.forward(&inputs)?;
229
230            let logits = match outputs {
231                ModelOutputs::Logits { logits, .. } => logits,
232                _ => return Err(anyhow::anyhow!("Expected logits output")),
233            };
234
235            let logits_candle = logits.to_candle()?;
236            let shape = logits_candle.dims();
237
238            let last_logits = if shape.len() == 3 {
239                let seq_len = shape[1];
240                logits_candle
241                    .narrow(1, seq_len - 1, 1)?
242                    .squeeze(1)?
243                    .squeeze(0)?
244            } else {
245                let seq_len = shape[0];
246                logits_candle
247                    .narrow(0, seq_len - 1, 1)?
248                    .squeeze(0)?
249            };
250
251            let logits_vec: Vec<f32> = last_logits.to_vec1()?;
252
253            let next_token = if config.do_sample && config.temperature > 0.0 {
254                let scaled: Vec<f32> = logits_vec.iter()
255                    .map(|&x| x / config.temperature)
256                    .collect();
257
258                let max_val = scaled.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
259                let exp_sum: f32 = scaled.iter().map(|&x| (x - max_val).exp()).sum();
260                let probs: Vec<f32> = scaled.iter()
261                    .map(|&x| (x - max_val).exp() / exp_sum)
262                    .collect();
263
264                let mut rng = rand::thread_rng();
265                let random_val: f32 = rng.gen();
266                let mut cumulative = 0.0;
267                let mut sampled = 0u32;
268
269                for (idx, &prob) in probs.iter().enumerate() {
270                    cumulative += prob;
271                    if random_val <= cumulative {
272                        sampled = idx as u32;
273                        break;
274                    }
275                }
276                sampled
277            } else {
278                let mut max_idx = 0;
279                let mut max_val = logits_vec[0];
280                for (idx, &val) in logits_vec.iter().enumerate() {
281                    if val > max_val {
282                        max_val = val;
283                        max_idx = idx;
284                    }
285                }
286                max_idx as u32
287            };
288
289            if next_token == config.eos_token_id {
290                break;
291            }
292
293            tokens.push(next_token);
294        }
295
296        Ok(tokenizer.decode(&tokens))
297    }
298
299    fn config(&self) -> &Self::Config {
300        &self.config
301    }
302
303    fn memory_requirements(&self) -> MemoryRequirements {
304        let param_size = self.config.vocab_size * self.config.hidden_size +
305                        self.config.num_hidden_layers * (
306                            4 * self.config.hidden_size * self.config.hidden_size +
307                            3 * self.config.hidden_size * self.config.intermediate_size
308                        );
309
310        let param_bytes = param_size * 4;
311        let kv_cache_bytes = 2 * self.config.num_hidden_layers *
312                           self.config.max_position_embeddings *
313                           self.config.hidden_size * 4;
314
315        MemoryRequirements {
316            gpu_memory: param_bytes,
317            cpu_memory: param_bytes / 4,
318            kv_cache_memory: kv_cache_bytes,
319            peak_memory: param_bytes + kv_cache_bytes,
320        }
321    }
322
323    fn to_device(&mut self, device: &Device) -> Result<()> {
324        self.embed_tokens = self.embed_tokens.to_device(device)?;
325        self.norm = self.norm.to_device(device)?;
326        self.lm_head = self.lm_head.to_device(device)?;
327
328        for layer in &mut self.layers {
329            layer.to_device(device)?;
330        }
331
332        self.device = device.clone();
333        Ok(())
334    }
335}
336
337impl GemmaLayer {
338    fn new(config: &GemmaConfig, device: &Device) -> Result<Self> {
339        let self_attn = GemmaAttention::new(config, device)?;
340        let mlp = GemmaMLP::new(config, device)?;
341
342        let input_layernorm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
343        let post_attention_layernorm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
344
345        Ok(Self {
346            self_attn,
347            mlp,
348            input_layernorm,
349            post_attention_layernorm,
350        })
351    }
352
353    fn forward(&self, hidden_states: &Tensor, attention_mask: Option<&Tensor>, rope_theta: f32, head_dim: usize) -> Result<Tensor> {
354        // 1. Pre-attention layer norm
355        let normed = ops_fn::layer_norm(hidden_states, &self.input_layernorm, None, 1e-6)?;
356
357        // 2. Self attention (with RoPE)
358        let attn_output = self.self_attn.forward(&normed, attention_mask, rope_theta, head_dim)?;
359
360        // 3. Residual connection
361        let hidden_states = ops_fn::add(hidden_states, &attn_output)?;
362
363        // 4. Pre-MLP layer norm
364        let normed = ops_fn::layer_norm(&hidden_states, &self.post_attention_layernorm, None, 1e-6)?;
365
366        // 5. MLP
367        let mlp_output = self.mlp.forward(&normed)?;
368
369        // 6. Residual connection
370        let output = ops_fn::add(&hidden_states, &mlp_output)?;
371
372        Ok(output)
373    }
374
375    fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
376        let prefix = format!("model.layers.{}", layer_idx);
377
378        // Load attention weights (transpose for matmul)
379        if let Some(q_proj) = weights.get(&format!("{}.self_attn.q_proj.weight", prefix)) {
380            self.self_attn.q_proj = ops_fn::transpose(q_proj)?;
381        }
382        if let Some(k_proj) = weights.get(&format!("{}.self_attn.k_proj.weight", prefix)) {
383            self.self_attn.k_proj = ops_fn::transpose(k_proj)?;
384        }
385        if let Some(v_proj) = weights.get(&format!("{}.self_attn.v_proj.weight", prefix)) {
386            self.self_attn.v_proj = ops_fn::transpose(v_proj)?;
387        }
388        if let Some(o_proj) = weights.get(&format!("{}.self_attn.o_proj.weight", prefix)) {
389            self.self_attn.o_proj = ops_fn::transpose(o_proj)?;
390        }
391
392        // Load MLP weights (transpose for matmul)
393        if let Some(gate_proj) = weights.get(&format!("{}.mlp.gate_proj.weight", prefix)) {
394            self.mlp.gate_proj = ops_fn::transpose(gate_proj)?;
395        }
396        if let Some(up_proj) = weights.get(&format!("{}.mlp.up_proj.weight", prefix)) {
397            self.mlp.up_proj = ops_fn::transpose(up_proj)?;
398        }
399        if let Some(down_proj) = weights.get(&format!("{}.mlp.down_proj.weight", prefix)) {
400            self.mlp.down_proj = ops_fn::transpose(down_proj)?;
401        }
402
403        // Load layer norm weights
404        if let Some(input_ln) = weights.get(&format!("{}.input_layernorm.weight", prefix)) {
405            self.input_layernorm = input_ln.clone();
406        }
407        if let Some(post_ln) = weights.get(&format!("{}.post_attention_layernorm.weight", prefix)) {
408            self.post_attention_layernorm = post_ln.clone();
409        }
410
411        Ok(())
412    }
413
414    fn to_device(&mut self, device: &Device) -> Result<()> {
415        self.self_attn.to_device(device)?;
416        self.mlp.to_device(device)?;
417        self.input_layernorm = self.input_layernorm.to_device(device)?;
418        self.post_attention_layernorm = self.post_attention_layernorm.to_device(device)?;
419        Ok(())
420    }
421}
422
423/// Apply Rotary Position Embedding (RoPE) to Q and K tensors
424fn apply_rope(
425    q: &candle_core::Tensor,
426    k: &candle_core::Tensor,
427    seq_len: usize,
428    head_dim: usize,
429    rope_theta: f32,
430) -> Result<(candle_core::Tensor, candle_core::Tensor)> {
431    let device = q.device();
432
433    let half_dim = head_dim / 2;
434    let inv_freq: Vec<f32> = (0..half_dim)
435        .map(|i| 1.0 / rope_theta.powf((2 * i) as f32 / head_dim as f32))
436        .collect();
437
438    let positions: Vec<f32> = (0..seq_len).map(|p| p as f32).collect();
439
440    let mut angles = Vec::with_capacity(seq_len * half_dim);
441    for pos in &positions {
442        for freq in &inv_freq {
443            angles.push(pos * freq);
444        }
445    }
446
447    let angles_tensor = candle_core::Tensor::from_vec(angles, &[seq_len, half_dim], device)?;
448
449    let cos = angles_tensor.cos()?;
450    let sin = angles_tensor.sin()?;
451
452    let cos = cos.unsqueeze(0)?.unsqueeze(0)?;
453    let sin = sin.unsqueeze(0)?.unsqueeze(0)?;
454
455    let q_half1 = q.narrow(3, 0, half_dim)?;
456    let q_half2 = q.narrow(3, half_dim, half_dim)?;
457    let k_half1 = k.narrow(3, 0, half_dim)?;
458    let k_half2 = k.narrow(3, half_dim, half_dim)?;
459
460    let q_rot1 = (q_half1.broadcast_mul(&cos)? - q_half2.broadcast_mul(&sin)?)?;
461    let q_rot2 = (q_half1.broadcast_mul(&sin)? + q_half2.broadcast_mul(&cos)?)?;
462    let k_rot1 = (k_half1.broadcast_mul(&cos)? - k_half2.broadcast_mul(&sin)?)?;
463    let k_rot2 = (k_half1.broadcast_mul(&sin)? + k_half2.broadcast_mul(&cos)?)?;
464
465    let q_rotated = candle_core::Tensor::cat(&[&q_rot1, &q_rot2], 3)?;
466    let k_rotated = candle_core::Tensor::cat(&[&k_rot1, &k_rot2], 3)?;
467
468    Ok((q_rotated, k_rotated))
469}
470
471impl GemmaAttention {
472    fn new(config: &GemmaConfig, device: &Device) -> Result<Self> {
473        let num_heads = config.num_attention_heads;
474        let num_key_value_heads = config.num_key_value_heads;
475        let head_dim = config.head_dim;
476        let scale = 1.0 / (head_dim as f32).sqrt();
477
478        let q_proj = ops_fn::zeros(&[config.hidden_size, num_heads * head_dim], DataType::Float32, device)?;
479        let k_proj = ops_fn::zeros(&[config.hidden_size, num_key_value_heads * head_dim], DataType::Float32, device)?;
480        let v_proj = ops_fn::zeros(&[config.hidden_size, num_key_value_heads * head_dim], DataType::Float32, device)?;
481        let o_proj = ops_fn::zeros(&[num_heads * head_dim, config.hidden_size], DataType::Float32, device)?;
482
483        Ok(Self {
484            q_proj,
485            k_proj,
486            v_proj,
487            o_proj,
488            num_heads,
489            num_key_value_heads,
490            head_dim,
491            scale,
492        })
493    }
494
495    fn forward(&self, hidden_states: &Tensor, _attention_mask: Option<&Tensor>, rope_theta: f32, head_dim: usize) -> Result<Tensor> {
496        let shape = hidden_states.shape();
497        let (batch_size, seq_len, _hidden_size) = if shape.len() == 3 {
498            (shape[0], shape[1], shape[2])
499        } else if shape.len() == 2 {
500            (1, shape[0], shape[1])
501        } else {
502            return Err(anyhow::anyhow!("Invalid hidden_states shape: {:?}", shape));
503        };
504
505        let query_states = ops_fn::matmul(hidden_states, &self.q_proj)?;
506        let key_states = ops_fn::matmul(hidden_states, &self.k_proj)?;
507        let value_states = ops_fn::matmul(hidden_states, &self.v_proj)?;
508
509        let q_candle = query_states.to_candle()?;
510        let k_candle = key_states.to_candle()?;
511        let v_candle = value_states.to_candle()?;
512
513        let q_reshaped = q_candle
514            .reshape(&[batch_size, seq_len, self.num_heads, head_dim])?
515            .transpose(1, 2)?;
516
517        let k_reshaped = k_candle
518            .reshape(&[batch_size, seq_len, self.num_key_value_heads, head_dim])?
519            .transpose(1, 2)?;
520
521        let v_reshaped = v_candle
522            .reshape(&[batch_size, seq_len, self.num_key_value_heads, head_dim])?
523            .transpose(1, 2)?;
524
525        let (q_with_rope, k_with_rope) = apply_rope(&q_reshaped, &k_reshaped, seq_len, head_dim, rope_theta)?;
526
527        let num_groups = self.num_heads / self.num_key_value_heads;
528        let (k_expanded, v_expanded) = if num_groups > 1 {
529            let k_rep = k_with_rope
530                .unsqueeze(2)?
531                .broadcast_as(&[batch_size, self.num_key_value_heads, num_groups, seq_len, head_dim])?
532                .reshape(&[batch_size, self.num_heads, seq_len, head_dim])?;
533            let v_rep = v_reshaped
534                .unsqueeze(2)?
535                .broadcast_as(&[batch_size, self.num_key_value_heads, num_groups, seq_len, head_dim])?
536                .reshape(&[batch_size, self.num_heads, seq_len, head_dim])?;
537            (k_rep, v_rep)
538        } else {
539            (k_with_rope, v_reshaped)
540        };
541
542        let k_t = k_expanded.transpose(2, 3)?;
543
544        let q_contiguous = q_with_rope.contiguous()?;
545        let k_contiguous = k_t.contiguous()?;
546
547        let scores = q_contiguous.matmul(&k_contiguous)?;
548        let scaled_scores = (scores * (self.scale as f64))?;
549
550        let device = scaled_scores.device();
551        let causal_mask = {
552            let mut mask_data = vec![0.0f32; seq_len * seq_len];
553            for i in 0..seq_len {
554                for j in 0..seq_len {
555                    if j > i {
556                        mask_data[i * seq_len + j] = f32::NEG_INFINITY;
557                    }
558                }
559            }
560            candle_core::Tensor::from_vec(mask_data, &[1, 1, seq_len, seq_len], device)?
561        };
562
563        let masked_scores = scaled_scores.broadcast_add(&causal_mask)?;
564
565        let attention_weights = candle_nn::ops::softmax_last_dim(&masked_scores)?;
566
567        let v_contiguous = v_expanded.contiguous()?;
568        let attn_output = attention_weights.matmul(&v_contiguous)?;
569
570        let attn_output = attn_output
571            .transpose(1, 2)?
572            .reshape(&[batch_size, seq_len, self.num_heads * head_dim])?;
573
574        let attn_output = Tensor::from_candle(attn_output);
575
576        let output = ops_fn::matmul(&attn_output, &self.o_proj)?;
577
578        Ok(output)
579    }
580
581    fn to_device(&mut self, device: &Device) -> Result<()> {
582        self.q_proj = self.q_proj.to_device(device)?;
583        self.k_proj = self.k_proj.to_device(device)?;
584        self.v_proj = self.v_proj.to_device(device)?;
585        self.o_proj = self.o_proj.to_device(device)?;
586        Ok(())
587    }
588}
589
590impl GemmaMLP {
591    fn new(config: &GemmaConfig, device: &Device) -> Result<Self> {
592        let gate_proj = ops_fn::zeros(&[config.hidden_size, config.intermediate_size], DataType::Float32, device)?;
593        let up_proj = ops_fn::zeros(&[config.hidden_size, config.intermediate_size], DataType::Float32, device)?;
594        let down_proj = ops_fn::zeros(&[config.intermediate_size, config.hidden_size], DataType::Float32, device)?;
595
596        Ok(Self {
597            gate_proj,
598            up_proj,
599            down_proj,
600            hidden_act: config.hidden_act.clone(),
601        })
602    }
603
604    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
605        let gate_output = ops_fn::matmul(hidden_states, &self.gate_proj)?;
606        let up_output = ops_fn::matmul(hidden_states, &self.up_proj)?;
607
608        let gate_activated = match self.hidden_act.as_str() {
609            "gelu" | "gelu_new" => ops_fn::gelu(&gate_output)?,
610            "silu" | "swish" => ops_fn::silu(&gate_output)?,
611            _ => return Err(anyhow::anyhow!("Unsupported activation: {}", self.hidden_act)),
612        };
613
614        let gated = ops_fn::mul(&gate_activated, &up_output)?;
615
616        let output = ops_fn::matmul(&gated, &self.down_proj)?;
617
618        Ok(output)
619    }
620
621    fn to_device(&mut self, device: &Device) -> Result<()> {
622        self.gate_proj = self.gate_proj.to_device(device)?;
623        self.up_proj = self.up_proj.to_device(device)?;
624        self.down_proj = self.down_proj.to_device(device)?;
625        Ok(())
626    }
627}
628
629#[cfg(test)]
630mod tests {
631    use super::*;
632
633    #[test]
634    fn test_gemma_model_creation() {
635        let config = GemmaConfig {
636            vocab_size: 1000,
637            hidden_size: 128,
638            intermediate_size: 512,
639            num_hidden_layers: 2,
640            num_attention_heads: 8,
641            num_key_value_heads: 8,
642            head_dim: 16,
643            ..Default::default()
644        };
645
646        let model = GemmaModelV2::new(config).unwrap();
647        assert_eq!(model.config().vocab_size(), 1000);
648        assert_eq!(model.config().hidden_size(), 128);
649        assert_eq!(model.config().num_layers(), 2);
650    }
651
652    #[test]
653    fn test_gemma_forward_pass() {
654        let config = GemmaConfig {
655            vocab_size: 100,
656            hidden_size: 64,
657            intermediate_size: 256,
658            num_hidden_layers: 1,
659            num_attention_heads: 4,
660            num_key_value_heads: 4,
661            head_dim: 16,
662            ..Default::default()
663        };
664
665        let model = GemmaModelV2::new(config).unwrap();
666        let input_ids = ops_fn::zeros(&[2, 8], DataType::Int64, &Device::CPU).unwrap();
667        let inputs = ModelInputs::text(input_ids);
668
669        let outputs = model.forward(&inputs).unwrap();
670        match outputs {
671            ModelOutputs::Logits { logits, .. } => {
672                assert_eq!(logits.shape(), &[2, 8, 100]);
673            }
674            _ => panic!("Expected logits output"),
675        }
676    }
677
678    #[test]
679    fn test_gemma_generation() {
680        let config = GemmaConfig {
681            vocab_size: 256,
682            hidden_size: 64,
683            intermediate_size: 256,
684            num_hidden_layers: 1,
685            num_attention_heads: 4,
686            num_key_value_heads: 4,
687            head_dim: 16,
688            ..Default::default()
689        };
690        let model = GemmaModelV2::new(config).unwrap();
691        let gen_config = GenerationConfig {
692            max_new_tokens: 5,
693            ..Default::default()
694        };
695
696        let output = model.generate("Hello", &gen_config).unwrap();
697        assert!(!output.is_empty());
698    }
699}