1use crate::model_config;
11use super::traits::*;
12use anyhow::Result;
13use serde::{Serialize, Deserialize};
14
15model_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 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
55pub struct MistralModelV2 {
57 config: MistralConfig,
58 device: Device,
59
60 embed_tokens: Tensor,
62 layers: Vec<MistralLayer>,
63 norm: Tensor,
64 lm_head: Tensor,
65}
66
67pub struct MistralLayer {
69 self_attn: MistralAttention,
70 mlp: MistralMLP,
71 input_layernorm: Tensor,
72 post_attention_layernorm: Tensor,
73}
74
75pub 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
88pub 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 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 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 if let Some(lm_head_weights) = weights.get("lm_head.weight") {
154 model.lm_head = ops_fn::transpose(lm_head_weights)?;
155 }
156
157 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 let mut hidden_states = ops_fn::embedding(input_ids, &self.embed_tokens)?;
170
171 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 hidden_states = ops_fn::rms_norm(&hidden_states, &self.norm, self.config.rms_norm_eps)?;
183
184 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 let normed = ops_fn::rms_norm(hidden_states, &self.input_layernorm, 1e-5)?;
347
348 let attn_output = self.self_attn.forward(&normed, attention_mask, rope_theta, sliding_window)?;
350
351 let hidden_states = ops_fn::add(hidden_states, &attn_output)?;
353
354 let normed = ops_fn::rms_norm(&hidden_states, &self.post_attention_layernorm, 1e-5)?;
356
357 let mlp_output = self.mlp.forward(&normed)?;
359
360 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 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 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 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
414fn 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 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 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 let (q_with_rope, k_with_rope) = apply_rope(&q_reshaped, &k_reshaped, seq_len, self.head_dim, rope_theta)?;
527
528 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 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 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 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 let attention_weights = candle_nn::ops::softmax_last_dim(&masked_scores)?;
573
574 let v_contiguous = v_expanded.contiguous()?;
576 let attn_output = attention_weights.matmul(&v_contiguous)?;
577
578 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 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}