Skip to main content

runtime/models_v2/
mistral.rs

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