Skip to main content

runtime/models_v2/
deepseek_moe.rs

1//! DeepSeek-MoE Model V2 - Clean implementation
2//!
3//! DeepSeek-MoE architecture features:
4//! - Fine-grained MoE with many experts (64-160)
5//! - Shared experts that are always active
6//! - Top-K routing with load balancing
7//! - RoPE embeddings
8
9use crate::model_config;
10use super::traits::*;
11use anyhow::Result;
12use serde::{Serialize, Deserialize};
13
14model_config!(DeepSeekMoEConfig {
15    vocab_size: usize = 102400,
16    hidden_size: usize = 2048,
17    intermediate_size: usize = 10944,
18    moe_intermediate_size: usize = 1408,
19    num_hidden_layers: usize = 28,
20    num_attention_heads: usize = 16,
21    num_key_value_heads: usize = 16,
22    hidden_act: String = "silu".to_string(),
23    max_position_embeddings: usize = 4096,
24    rms_norm_eps: f32 = 1e-6,
25    use_cache: bool = true,
26    pad_token_id: i64 = 0,
27    bos_token_id: i64 = 1,
28    eos_token_id: i64 = 2,
29    tie_word_embeddings: bool = false,
30    rope_theta: f32 = 10000.0,
31    num_experts: usize = 64,
32    num_experts_per_tok: usize = 6,
33    num_shared_experts: usize = 2,
34    first_k_dense_replace: usize = 1,
35});
36
37impl DeepSeekMoEConfig {
38    pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
39        Self {
40            vocab_size: gguf.vocab_size,
41            hidden_size: gguf.hidden_size,
42            intermediate_size: gguf.intermediate_size,
43            num_hidden_layers: gguf.num_hidden_layers,
44            num_attention_heads: gguf.num_attention_heads,
45            num_key_value_heads: gguf.num_key_value_heads,
46            max_position_embeddings: gguf.max_position_embeddings,
47            rope_theta: gguf.rope_theta,
48            ..Default::default()
49        }
50    }
51}
52
53pub struct DeepSeekMoEModelV2 {
54    config: DeepSeekMoEConfig,
55    device: Device,
56    embed_tokens: Tensor,
57    layers: Vec<DeepSeekMoELayer>,
58    norm: Tensor,
59    lm_head: Tensor,
60}
61
62pub struct DeepSeekMoELayer {
63    self_attn: DeepSeekAttention,
64    mlp: DeepSeekMoEBlock,
65    input_layernorm: Tensor,
66    post_attention_layernorm: Tensor,
67}
68
69pub struct DeepSeekAttention {
70    q_proj: Tensor,
71    k_proj: Tensor,
72    v_proj: Tensor,
73    o_proj: Tensor,
74    num_heads: usize,
75    num_kv_heads: usize,
76    head_dim: usize,
77    scale: f32,
78}
79
80pub struct DeepSeekMoEBlock {
81    router: Tensor,
82    shared_experts: Vec<DeepSeekExpert>,
83    routed_experts: Vec<DeepSeekExpert>,
84    num_experts_per_tok: usize,
85}
86
87pub struct DeepSeekExpert {
88    gate_proj: Tensor,
89    up_proj: Tensor,
90    down_proj: Tensor,
91}
92
93fn apply_rope(
94    q: &candle_core::Tensor,
95    k: &candle_core::Tensor,
96    seq_len: usize,
97    head_dim: usize,
98    rope_theta: f32,
99) -> Result<(candle_core::Tensor, candle_core::Tensor)> {
100    let device = q.device();
101    let half_dim = head_dim / 2;
102    let inv_freq: Vec<f32> = (0..half_dim)
103        .map(|i| 1.0 / rope_theta.powf((2 * i) as f32 / head_dim as f32))
104        .collect();
105
106    let positions: Vec<f32> = (0..seq_len).map(|p| p as f32).collect();
107    let mut angles = Vec::with_capacity(seq_len * half_dim);
108    for pos in &positions {
109        for freq in &inv_freq {
110            angles.push(pos * freq);
111        }
112    }
113
114    let angles_tensor = candle_core::Tensor::from_vec(angles, &[seq_len, half_dim], device)?;
115    let cos = angles_tensor.cos()?.unsqueeze(0)?.unsqueeze(0)?;
116    let sin = angles_tensor.sin()?.unsqueeze(0)?.unsqueeze(0)?;
117
118    let q_half1 = q.narrow(3, 0, half_dim)?;
119    let q_half2 = q.narrow(3, half_dim, half_dim)?;
120    let k_half1 = k.narrow(3, 0, half_dim)?;
121    let k_half2 = k.narrow(3, half_dim, half_dim)?;
122
123    let q_rot1 = (q_half1.broadcast_mul(&cos)? - q_half2.broadcast_mul(&sin)?)?;
124    let q_rot2 = (q_half1.broadcast_mul(&sin)? + q_half2.broadcast_mul(&cos)?)?;
125    let k_rot1 = (k_half1.broadcast_mul(&cos)? - k_half2.broadcast_mul(&sin)?)?;
126    let k_rot2 = (k_half1.broadcast_mul(&sin)? + k_half2.broadcast_mul(&cos)?)?;
127
128    Ok((
129        candle_core::Tensor::cat(&[&q_rot1, &q_rot2], 3)?,
130        candle_core::Tensor::cat(&[&k_rot1, &k_rot2], 3)?
131    ))
132}
133
134impl Model for DeepSeekMoEModelV2 {
135    type Config = DeepSeekMoEConfig;
136
137    fn new(config: DeepSeekMoEConfig) -> Result<Self> {
138        let device = Device::CPU;
139        let embed_tokens = ops_fn::zeros(&[config.vocab_size, config.hidden_size], DataType::Float32, &device)?;
140        let norm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, &device)?;
141        let lm_head = ops_fn::zeros(&[config.vocab_size, config.hidden_size], DataType::Float32, &device)?;
142
143        let mut layers = Vec::with_capacity(config.num_hidden_layers);
144        for i in 0..config.num_hidden_layers {
145            layers.push(DeepSeekMoELayer::new(&config, &device, i)?);
146        }
147
148        Ok(Self { config, device, embed_tokens, layers, norm, lm_head })
149    }
150
151    fn from_weights(config: DeepSeekMoEConfig, weights: ModelWeights) -> Result<Self> {
152        let mut model = Self::new(config)?;
153        if let Some(w) = weights.get("model.embed_tokens.weight") { model.embed_tokens = w.clone(); }
154        if let Some(w) = weights.get("model.norm.weight") { model.norm = w.clone(); }
155        if let Some(w) = weights.get("lm_head.weight") { model.lm_head = w.clone(); }
156        for (i, layer) in model.layers.iter_mut().enumerate() { layer.load_weights(&weights, i)?; }
157        Ok(model)
158    }
159
160    fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
161        match inputs {
162            ModelInputs::Text { input_ids, .. } => {
163                let seq_len = input_ids.shape()[1];
164                let mut hidden = ops_fn::embedding(input_ids, &self.embed_tokens)?;
165
166                for layer in &self.layers {
167                    hidden = layer.forward(&hidden, seq_len, self.config.rope_theta)?;
168                }
169
170                hidden = ops_fn::rms_norm(&hidden, &self.norm, self.config.rms_norm_eps)?;
171                let logits = ops_fn::matmul(&hidden, &ops_fn::transpose(&self.lm_head)?)?;
172
173                Ok(ModelOutputs::Logits { logits, hidden_states: None })
174            }
175            _ => Err(anyhow::anyhow!("DeepSeek-MoE only supports text inputs")),
176        }
177    }
178
179    fn generate(&self, prompt: &str, config: &GenerationConfig) -> Result<String> {
180        use crate::tokenizer::Tokenizer;
181        use rand::Rng;
182        let tokenizer = Tokenizer::new();
183        let mut tokens: Vec<u32> = tokenizer.encode(prompt);
184        for _ in 0..config.max_new_tokens {
185            let tokens_i64: Vec<i64> = tokens.iter().map(|&t| t as i64).collect();
186            let input = Tensor::from_i64_slice(&tokens_i64, &[1, tokens.len()], &self.device)?;
187            let outputs = self.forward(&ModelInputs::text(input))?;
188            let logits = match outputs { ModelOutputs::Logits { logits, .. } => logits, _ => return Err(anyhow::anyhow!("Expected logits")) };
189            let logits_candle = logits.to_candle()?;
190            let last = logits_candle.narrow(1, logits_candle.dims()[1] - 1, 1)?.squeeze(1)?.squeeze(0)?;
191            let logits_vec: Vec<f32> = last.to_vec1()?;
192            let next = if config.do_sample && config.temperature > 0.0 {
193                let scaled: Vec<f32> = logits_vec.iter().map(|&x| x / config.temperature).collect();
194                let max_v = scaled.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
195                let exp_sum: f32 = scaled.iter().map(|&x| (x - max_v).exp()).sum();
196                let probs: Vec<f32> = scaled.iter().map(|&x| (x - max_v).exp() / exp_sum).collect();
197                let mut rng = rand::thread_rng();
198                let r: f32 = rng.gen();
199                let mut cum = 0.0;
200                let mut s = 0u32;
201                for (i, &p) in probs.iter().enumerate() { cum += p; if r <= cum { s = i as u32; break; } }
202                s
203            } else {
204                logits_vec.iter().enumerate().max_by(|a, b| a.1.partial_cmp(b.1).unwrap()).map(|(i, _)| i as u32).unwrap_or(0)
205            };
206            if next == config.eos_token_id { break; }
207            tokens.push(next);
208        }
209        Ok(tokenizer.decode(&tokens))
210    }
211
212    fn config(&self) -> &Self::Config { &self.config }
213    fn memory_requirements(&self) -> MemoryRequirements {
214        let p = self.config.vocab_size * self.config.hidden_size + self.config.num_hidden_layers * 8 * self.config.hidden_size.pow(2);
215        MemoryRequirements { gpu_memory: p * 4, cpu_memory: p, kv_cache_memory: 2 * self.config.num_hidden_layers * self.config.max_position_embeddings * self.config.hidden_size * 4, peak_memory: p * 5 }
216    }
217    fn to_device(&mut self, device: &Device) -> Result<()> {
218        self.embed_tokens = self.embed_tokens.to_device(device)?;
219        self.norm = self.norm.to_device(device)?;
220        self.lm_head = self.lm_head.to_device(device)?;
221        for layer in &mut self.layers { layer.to_device(device)?; }
222        self.device = device.clone();
223        Ok(())
224    }
225}
226
227impl DeepSeekMoELayer {
228    fn new(config: &DeepSeekMoEConfig, device: &Device, layer_idx: usize) -> Result<Self> {
229        Ok(Self {
230            self_attn: DeepSeekAttention::new(config, device)?,
231            mlp: DeepSeekMoEBlock::new(config, device, layer_idx)?,
232            input_layernorm: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
233            post_attention_layernorm: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
234        })
235    }
236
237    fn forward(&self, hidden_states: &Tensor, seq_len: usize, rope_theta: f32) -> Result<Tensor> {
238        let residual = hidden_states.clone();
239        let h = ops_fn::rms_norm(hidden_states, &self.input_layernorm, 1e-6)?;
240        let attn_out = self.self_attn.forward(&h, seq_len, rope_theta)?;
241        let h = ops_fn::add(&residual, &attn_out)?;
242
243        let residual = h.clone();
244        let h = ops_fn::rms_norm(&h, &self.post_attention_layernorm, 1e-6)?;
245        let mlp_out = self.mlp.forward(&h)?;
246        ops_fn::add(&residual, &mlp_out)
247    }
248
249    fn load_weights(&mut self, weights: &ModelWeights, idx: usize) -> Result<()> {
250        let p = format!("model.layers.{}", idx);
251        if let Some(w) = weights.get(&format!("{}.self_attn.q_proj.weight", p)) { self.self_attn.q_proj = ops_fn::transpose(w)?; }
252        if let Some(w) = weights.get(&format!("{}.self_attn.k_proj.weight", p)) { self.self_attn.k_proj = ops_fn::transpose(w)?; }
253        if let Some(w) = weights.get(&format!("{}.self_attn.v_proj.weight", p)) { self.self_attn.v_proj = ops_fn::transpose(w)?; }
254        if let Some(w) = weights.get(&format!("{}.self_attn.o_proj.weight", p)) { self.self_attn.o_proj = ops_fn::transpose(w)?; }
255        if let Some(w) = weights.get(&format!("{}.input_layernorm.weight", p)) { self.input_layernorm = w.clone(); }
256        if let Some(w) = weights.get(&format!("{}.post_attention_layernorm.weight", p)) { self.post_attention_layernorm = w.clone(); }
257        Ok(())
258    }
259
260    fn to_device(&mut self, device: &Device) -> Result<()> {
261        self.self_attn.to_device(device)?;
262        self.mlp.to_device(device)?;
263        self.input_layernorm = self.input_layernorm.to_device(device)?;
264        self.post_attention_layernorm = self.post_attention_layernorm.to_device(device)?;
265        Ok(())
266    }
267}
268
269impl DeepSeekAttention {
270    fn new(config: &DeepSeekMoEConfig, device: &Device) -> Result<Self> {
271        let head_dim = config.hidden_size / config.num_attention_heads;
272        Ok(Self {
273            q_proj: ops_fn::zeros(&[config.hidden_size, config.num_attention_heads * head_dim], DataType::Float32, device)?,
274            k_proj: ops_fn::zeros(&[config.hidden_size, config.num_key_value_heads * head_dim], DataType::Float32, device)?,
275            v_proj: ops_fn::zeros(&[config.hidden_size, config.num_key_value_heads * head_dim], DataType::Float32, device)?,
276            o_proj: ops_fn::zeros(&[config.num_attention_heads * head_dim, config.hidden_size], DataType::Float32, device)?,
277            num_heads: config.num_attention_heads,
278            num_kv_heads: config.num_key_value_heads,
279            head_dim,
280            scale: 1.0 / (head_dim as f32).sqrt(),
281        })
282    }
283
284    fn forward(&self, hidden_states: &Tensor, seq_len: usize, rope_theta: f32) -> Result<Tensor> {
285        let shape = hidden_states.shape();
286        let batch = shape[0];
287
288        let q = ops_fn::matmul(hidden_states, &self.q_proj)?.to_candle()?;
289        let k = ops_fn::matmul(hidden_states, &self.k_proj)?.to_candle()?;
290        let v = ops_fn::matmul(hidden_states, &self.v_proj)?.to_candle()?;
291
292        let q = q.reshape(&[batch, seq_len, self.num_heads, self.head_dim])?.transpose(1, 2)?;
293        let k = k.reshape(&[batch, seq_len, self.num_kv_heads, self.head_dim])?.transpose(1, 2)?;
294        let v = v.reshape(&[batch, seq_len, self.num_kv_heads, self.head_dim])?.transpose(1, 2)?;
295
296        let (q, k) = apply_rope(&q, &k, seq_len, self.head_dim, rope_theta)?;
297
298        let num_groups = self.num_heads / self.num_kv_heads;
299        let (k, v) = if num_groups > 1 {
300            let k_exp = k.unsqueeze(2)?.broadcast_as(&[batch, self.num_kv_heads, num_groups, seq_len, self.head_dim])?.reshape(&[batch, self.num_heads, seq_len, self.head_dim])?;
301            let v_exp = v.unsqueeze(2)?.broadcast_as(&[batch, self.num_kv_heads, num_groups, seq_len, self.head_dim])?.reshape(&[batch, self.num_heads, seq_len, self.head_dim])?;
302            (k_exp, v_exp)
303        } else {
304            (k, v)
305        };
306
307        let q = q.contiguous()?;
308        let k_t = k.transpose(2, 3)?.contiguous()?;
309        let scores = (q.matmul(&k_t)? * (self.scale as f64))?;
310
311        let device = scores.device();
312        let mask = {
313            let mut m = vec![0.0f32; seq_len * seq_len];
314            for i in 0..seq_len { for j in (i+1)..seq_len { m[i*seq_len+j] = f32::NEG_INFINITY; } }
315            candle_core::Tensor::from_vec(m, &[1, 1, seq_len, seq_len], device)?
316        };
317        let scores = scores.broadcast_add(&mask)?;
318
319        let v = v.contiguous()?;
320        let attn = candle_nn::ops::softmax_last_dim(&scores)?.matmul(&v)?;
321        let out = attn.transpose(1, 2)?.reshape(&[batch, seq_len, self.num_heads * self.head_dim])?;
322        ops_fn::matmul(&Tensor::from_candle(out), &self.o_proj)
323    }
324
325    fn to_device(&mut self, device: &Device) -> Result<()> {
326        self.q_proj = self.q_proj.to_device(device)?;
327        self.k_proj = self.k_proj.to_device(device)?;
328        self.v_proj = self.v_proj.to_device(device)?;
329        self.o_proj = self.o_proj.to_device(device)?;
330        Ok(())
331    }
332}
333
334impl DeepSeekMoEBlock {
335    fn new(config: &DeepSeekMoEConfig, device: &Device, layer_idx: usize) -> Result<Self> {
336        // First few layers use dense MLP, rest use MoE
337        let use_moe = layer_idx >= config.first_k_dense_replace;
338
339        let mut shared_experts = Vec::with_capacity(config.num_shared_experts);
340        for _ in 0..config.num_shared_experts {
341            shared_experts.push(DeepSeekExpert::new(config.hidden_size, config.moe_intermediate_size, device)?);
342        }
343
344        let mut routed_experts = Vec::new();
345        if use_moe {
346            for _ in 0..config.num_experts {
347                routed_experts.push(DeepSeekExpert::new(config.hidden_size, config.moe_intermediate_size, device)?);
348            }
349        }
350
351        let router = if use_moe {
352            ops_fn::zeros(&[config.hidden_size, config.num_experts], DataType::Float32, device)?
353        } else {
354            ops_fn::zeros(&[1, 1], DataType::Float32, device)?
355        };
356
357        Ok(Self {
358            router,
359            shared_experts,
360            routed_experts,
361            num_experts_per_tok: config.num_experts_per_tok,
362        })
363    }
364
365    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
366        let shape = hidden_states.shape();
367        let (batch_size, seq_len, hidden_size) = (shape[0], shape[1], shape[2]);
368
369        // Shared experts (always active)
370        let mut output = ops_fn::zeros(&[batch_size, seq_len, hidden_size], hidden_states.dtype(), hidden_states.device())?;
371        for expert in &self.shared_experts {
372            let expert_out = expert.forward(hidden_states)?;
373            output = ops_fn::add(&output, &expert_out)?;
374        }
375
376        // Routed experts
377        if !self.routed_experts.is_empty() {
378            let num_tokens = batch_size * seq_len;
379            let k = self.num_experts_per_tok;
380
381            let flat_hidden = hidden_states.reshape(&[num_tokens, hidden_size])?;
382            let router_logits = ops_fn::matmul(&flat_hidden, &self.router)?;
383
384            let (topk_weights, topk_indices) = ops_fn::topk(&router_logits, k, -1)?;
385            let routing_weights = ops_fn::softmax(&topk_weights, -1)?;
386
387            let all_indices: Vec<i64> = topk_indices.to_candle()?.flatten_all()?.to_vec1()?;
388            let all_weights: Vec<f32> = routing_weights.to_candle()?.flatten_all()?.to_vec1()?;
389            let flat_hidden_candle = flat_hidden.to_candle()?;
390
391            let mut routed_output_data = vec![0.0f32; num_tokens * hidden_size];
392
393            for tok_idx in 0..num_tokens {
394                let token_hidden = flat_hidden_candle.get(tok_idx)?;
395                let token_tensor = Tensor::from_candle(token_hidden.unsqueeze(0)?);
396
397                let start = tok_idx * k;
398                let indices = &all_indices[start..start + k];
399                let weights = &all_weights[start..start + k];
400
401                let mut token_output = ops_fn::zeros(&[1, hidden_size], hidden_states.dtype(), hidden_states.device())?;
402
403                for (i, &expert_idx) in indices.iter().enumerate() {
404                    if (expert_idx as usize) < self.routed_experts.len() {
405                        let expert = &self.routed_experts[expert_idx as usize];
406                        let expert_output = expert.forward(&token_tensor)?;
407                        let scaled_output = ops_fn::scale(&expert_output, weights[i])?;
408                        token_output = ops_fn::add(&token_output, &scaled_output)?;
409                    }
410                }
411
412                let token_data: Vec<f32> = token_output.to_candle()?.flatten_all()?.to_vec1()?;
413                for (i, &v) in token_data.iter().enumerate() {
414                    routed_output_data[tok_idx * hidden_size + i] = v;
415                }
416            }
417
418            let routed_output = Tensor::from_f32_slice(&routed_output_data, &[num_tokens, hidden_size], hidden_states.device())?;
419            let routed_output = routed_output.reshape(&[batch_size, seq_len, hidden_size])?;
420            output = ops_fn::add(&output, &routed_output)?;
421        }
422
423        Ok(output)
424    }
425
426    fn to_device(&mut self, device: &Device) -> Result<()> {
427        self.router = self.router.to_device(device)?;
428        for expert in &mut self.shared_experts { expert.to_device(device)?; }
429        for expert in &mut self.routed_experts { expert.to_device(device)?; }
430        Ok(())
431    }
432}
433
434impl DeepSeekExpert {
435    fn new(hidden_size: usize, intermediate_size: usize, device: &Device) -> Result<Self> {
436        Ok(Self {
437            gate_proj: ops_fn::zeros(&[hidden_size, intermediate_size], DataType::Float32, device)?,
438            up_proj: ops_fn::zeros(&[hidden_size, intermediate_size], DataType::Float32, device)?,
439            down_proj: ops_fn::zeros(&[intermediate_size, hidden_size], DataType::Float32, device)?,
440        })
441    }
442
443    fn forward(&self, x: &Tensor) -> Result<Tensor> {
444        let gate = ops_fn::matmul(x, &self.gate_proj)?;
445        let up = ops_fn::matmul(x, &self.up_proj)?;
446        let h = ops_fn::mul(&ops_fn::silu(&gate)?, &up)?;
447        ops_fn::matmul(&h, &self.down_proj)
448    }
449
450    fn to_device(&mut self, device: &Device) -> Result<()> {
451        self.gate_proj = self.gate_proj.to_device(device)?;
452        self.up_proj = self.up_proj.to_device(device)?;
453        self.down_proj = self.down_proj.to_device(device)?;
454        Ok(())
455    }
456}
457
458#[cfg(test)]
459mod tests {
460    use super::*;
461    #[test]
462    fn test_deepseek_moe_creation() {
463        let config = DeepSeekMoEConfig { vocab_size: 1000, hidden_size: 64, intermediate_size: 256, moe_intermediate_size: 64, num_hidden_layers: 2, num_attention_heads: 4, num_key_value_heads: 4, num_experts: 4, num_experts_per_tok: 2, num_shared_experts: 1, first_k_dense_replace: 1, ..Default::default() };
464        let model = DeepSeekMoEModelV2::new(config).unwrap();
465        assert_eq!(model.config().vocab_size(), 1000);
466    }
467    #[test]
468    fn test_deepseek_moe_forward() {
469        let config = DeepSeekMoEConfig { vocab_size: 100, hidden_size: 64, intermediate_size: 256, moe_intermediate_size: 64, num_hidden_layers: 2, num_attention_heads: 4, num_key_value_heads: 4, num_experts: 4, num_experts_per_tok: 2, num_shared_experts: 1, first_k_dense_replace: 1, ..Default::default() };
470        let model = DeepSeekMoEModelV2::new(config).unwrap();
471        let inputs = ModelInputs::text(ops_fn::zeros(&[1, 4], DataType::Int64, &Device::CPU).unwrap());
472        match model.forward(&inputs).unwrap() { ModelOutputs::Logits { logits, .. } => assert_eq!(logits.shape(), &[1, 4, 100]), _ => panic!() }
473    }
474}