1use 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 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 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 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}