1use crate::model_config;
9use super::traits::*;
10use anyhow::Result;
11use serde::{Serialize, Deserialize};
12
13model_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 num_experts: Option<usize> = None,
34 num_experts_per_tok: Option<usize> = None,
35 expert_capacity: Option<usize> = None,
36});
37
38impl DeepSeekConfig {
39 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
56pub struct DeepSeekModelV2 {
58 config: DeepSeekConfig,
59 device: Device,
60
61 embed_tokens: Tensor,
63 layers: Vec<DeepSeekLayer>,
64 norm: Tensor,
65 lm_head: Tensor,
66}
67
68pub struct DeepSeekLayer {
70 self_attn: DeepSeekAttention,
71 mlp: DeepSeekMLP,
72 input_layernorm: Tensor,
73 post_attention_layernorm: Tensor,
74}
75
76pub 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
88pub 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 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 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 if let Some(lm_head_weights) = weights.get("lm_head.weight") {
155 model.lm_head = ops_fn::transpose(lm_head_weights)?;
156 }
157
158 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 let mut hidden_states = ops_fn::embedding(input_ids, &self.embed_tokens)?;
171
172 for layer in &self.layers {
174 hidden_states = layer.forward(&hidden_states, attention_mask.as_ref(), self.config.rope_theta)?;
175 }
176
177 hidden_states = ops_fn::layer_norm(&hidden_states, &self.norm, None, self.config.rms_norm_eps)?;
179
180 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 let tokenizer = Tokenizer::new();
206 let mut tokens: Vec<u32> = tokenizer.encode(prompt);
207
208 for _ in 0..config.max_new_tokens {
210 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 let outputs = self.forward(&inputs)?;
222
223 let logits = match outputs {
225 ModelOutputs::Logits { logits, .. } => logits,
226 _ => return Err(anyhow::anyhow!("Expected logits output")),
227 };
228
229 let logits_candle = logits.to_candle()?;
231 let shape = logits_candle.dims();
232
233 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 let logits_vec: Vec<f32> = last_logits.to_vec1()?;
249
250 let next_token = if config.do_sample && config.temperature > 0.0 {
251 let scaled: Vec<f32> = logits_vec.iter()
253 .map(|&x| x / config.temperature)
254 .collect();
255
256 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 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 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 if next_token == config.eos_token_id {
292 break;
293 }
294
295 tokens.push(next_token);
297 }
298
299 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 let normed = ops_fn::layer_norm(hidden_states, &self.input_layernorm, None, 1e-6)?;
360
361 let attn_output = self.self_attn.forward(&normed, attention_mask, rope_theta)?;
363
364 let hidden_states = ops_fn::add(hidden_states, &attn_output)?;
366
367 let normed = ops_fn::layer_norm(&hidden_states, &self.post_attention_layernorm, None, 1e-6)?;
369
370 let mlp_output = self.mlp.forward(&normed)?;
372
373 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 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 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 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
427fn 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 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 let positions: Vec<f32> = (0..seq_len).map(|p| p as f32).collect();
445
446 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 let cos = angles_tensor.cos()?;
458 let sin = angles_tensor.sin()?;
459
460 let cos = cos.unsqueeze(0)?.unsqueeze(0)?;
462 let sin = sin.unsqueeze(0)?.unsqueeze(0)?;
463
464 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 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 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 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 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 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 let (q_with_rope, k_with_rope) = apply_rope(&q_reshaped, &k_reshaped, seq_len, self.head_dim, rope_theta)?;
542
543 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 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 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 let attention_weights = candle_nn::ops::softmax_last_dim(&masked_scores)?;
586
587 let v_contiguous = v_expanded.contiguous()?;
589 let attn_output = attention_weights.matmul(&v_contiguous)?;
590
591 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 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 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 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 let gated = ops_fn::mul(&gate_activated, &up_output)?;
641
642 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}