Skip to main content

runtime/models_v2/
deepseek.rs

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