Skip to main content

runtime/models_v2/
mixtral.rs

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