Skip to main content

candle_transformers/models/
deepseek2.rs

1#![allow(clippy::cast_possible_truncation, clippy::cast_precision_loss)]
2
3use std::{f32::consts::PI, sync::Arc};
4
5use candle::{
6    shape::Dim, CpuStorage, CustomOp1, DType, Device, Error, IndexOp, Layout, Result, Shape,
7    Tensor, WithDType, D,
8};
9use candle_nn::{embedding, rms_norm, Activation, Embedding, Linear, Module, RmsNorm, VarBuilder};
10use rayon::iter::{IntoParallelRefIterator, ParallelIterator};
11use serde::Deserialize;
12
13struct NonZero {}
14
15impl NonZero {
16    // Sequential version
17    fn nonzero<T: WithDType>(&self, vs: &[T], layout: &Layout) -> Vec<u32> {
18        let n = layout.dims().len();
19        let mut result = Vec::new();
20        let mut indices = vec![0u32; n];
21        for (i, v) in vs.iter().enumerate() {
22            if !v.is_zero() {
23                let mut idx = i;
24                for (dim_index, dim) in layout.dims().iter().enumerate().rev() {
25                    let d = idx % dim;
26                    indices[dim_index] = u32::try_from(d).unwrap();
27                    idx /= dim;
28                }
29                result.extend_from_slice(&indices);
30            }
31        }
32        result
33    }
34}
35
36impl CustomOp1 for NonZero {
37    fn name(&self) -> &'static str {
38        "nonzero"
39    }
40
41    fn cpu_fwd(&self, storage: &CpuStorage, layout: &Layout) -> Result<(CpuStorage, Shape)> {
42        if !layout.is_contiguous() {
43            return Err(Error::RequiresContiguous { op: "nonzero" });
44        }
45        let result = match storage {
46            candle::CpuStorage::U8(vs) => self.nonzero(vs, layout),
47            candle::CpuStorage::U32(vs) => self.nonzero(vs, layout),
48            candle::CpuStorage::I16(vs) => self.nonzero(vs, layout),
49            candle::CpuStorage::I32(vs) => self.nonzero(vs, layout),
50            candle::CpuStorage::I64(vs) => self.nonzero(vs, layout),
51            candle::CpuStorage::BF16(vs) => self.nonzero(vs, layout),
52            candle::CpuStorage::F16(vs) => self.nonzero(vs, layout),
53            candle::CpuStorage::F32(vs) => self.nonzero(vs, layout),
54            candle::CpuStorage::F64(vs) => self.nonzero(vs, layout),
55            candle::CpuStorage::F8E4M3(vs) => self.nonzero(vs, layout),
56            // Dummy types don't support nonzero operation
57            candle::CpuStorage::F6E2M3(_) => {
58                return Err(
59                    candle::Error::UnsupportedDTypeForOp(candle::DType::F6E2M3, "nonzero").bt(),
60                )
61            }
62            candle::CpuStorage::F6E3M2(_) => {
63                return Err(
64                    candle::Error::UnsupportedDTypeForOp(candle::DType::F6E3M2, "nonzero").bt(),
65                )
66            }
67            candle::CpuStorage::F4(_) => {
68                return Err(candle::Error::UnsupportedDTypeForOp(candle::DType::F4, "nonzero").bt())
69            }
70            candle::CpuStorage::F8E8M0(_) => {
71                return Err(
72                    candle::Error::UnsupportedDTypeForOp(candle::DType::F8E8M0, "nonzero").bt(),
73                )
74            }
75        };
76        let index_len = layout.dims().len();
77        let result_len = result.len() / index_len;
78        let result = CpuStorage::U32(result);
79        let shape = Shape::from_dims(&[result_len, index_len]);
80        Ok((result, shape))
81    }
82}
83
84pub trait NonZeroOp {
85    fn nonzero(&self) -> Result<Tensor>;
86}
87
88impl NonZeroOp for Tensor {
89    fn nonzero(&self) -> Result<Tensor> {
90        if !self.is_contiguous() {
91            return Err(candle::Error::RequiresContiguous { op: "nonzero" });
92        }
93        let original_device = self.device();
94        self.to_device(&candle::Device::Cpu)?
95            .apply_op1_no_bwd(&NonZero {})?
96            .to_device(original_device)
97    }
98}
99
100pub struct TopKOutput {
101    pub values: Tensor,
102    pub indices: Tensor,
103}
104
105pub trait TopKLastDimOp {
106    /// Topk in the last dim. `values` retains a gradient but `indices` has none w.r.t self.
107    /// This expects a contiguous tensor.
108    /// Note: this implements torch.topk with sorted=True.
109    fn topk(&self, topk: usize) -> Result<TopKOutput>;
110
111    /// Topk in the last dim. `values` retains a gradient but `indices` has none w.r.t self.
112    /// This expects a contiguous tensor.
113    /// Note: this implements torch.topk with sorted=False.
114    fn topk_unsorted(&self, topk: usize) -> Result<TopKOutput>;
115}
116
117impl TopKLastDimOp for Tensor {
118    fn topk(&self, topk: usize) -> Result<TopKOutput> {
119        // Sorted descending
120        let sorted_indices = self.arg_sort_last_dim(false)?;
121        let topk_indices = sorted_indices.narrow(D::Minus1, 0, topk)?.contiguous()?;
122        Ok(TopKOutput {
123            values: self.gather(&topk_indices, D::Minus1)?,
124            indices: topk_indices,
125        })
126    }
127
128    fn topk_unsorted(&self, topk: usize) -> Result<TopKOutput> {
129        // Sorted descending
130        let sorted_indices_all = self.arg_sort_last_dim(false)?;
131        let topk_indices_sorted = sorted_indices_all
132            .narrow(D::Minus1, 0, topk)?
133            .contiguous()?;
134        let topk_values_sorted = self.gather(&topk_indices_sorted, D::Minus1)?;
135
136        // Reorder the indices ascending
137        let reorder_indices = topk_indices_sorted.arg_sort_last_dim(true)?;
138        let topk_indices_unsorted = topk_indices_sorted.gather(&reorder_indices, D::Minus1)?;
139        let topk_values_unsorted = topk_values_sorted.gather(&reorder_indices, D::Minus1)?;
140        Ok(TopKOutput {
141            values: topk_values_unsorted,
142            indices: topk_indices_unsorted,
143        })
144    }
145}
146
147pub trait SplitOp {
148    fn split<D: Dim>(&self, splits: &[usize], dim: D) -> Result<Vec<Tensor>>;
149}
150
151impl SplitOp for Tensor {
152    fn split<D: Dim>(&self, splits: &[usize], dim: D) -> Result<Vec<Tensor>> {
153        let dim = dim.to_index(self.shape(), "split")?;
154        let mut split_res = Vec::new();
155        let mut index = 0;
156        for split in splits {
157            split_res.push(self.narrow(dim, index, *split)?);
158            index += *split;
159        }
160        Ok(split_res)
161    }
162}
163
164pub trait BincountOp {
165    fn bincount(&self, minlength: u32) -> Result<Vec<u32>>;
166}
167
168fn bincount(values: &[u32], minlength: u32) -> Vec<u32> {
169    // Find the maximum value in `values` (or zero if empty)
170    let max_val = values.par_iter().max().copied().unwrap_or(0);
171
172    // The final size of the bin counts must be at least `minlength`
173    // and large enough to include the largest value in `values`.
174    let result_len = (max_val + 1).max(minlength);
175
176    // Each thread creates a local histogram (`fold`),
177    // and then they are merged together (`reduce`).
178    values
179        .par_iter()
180        .fold(
181            // Create a local histogram
182            || vec![0u32; result_len as usize],
183            // Update the local histogram
184            |mut local_counts, &val| {
185                local_counts[val as usize] += 1;
186                local_counts
187            },
188        )
189        // Merge histograms from all threads
190        .reduce(
191            // Identity (empty histogram)
192            || vec![0u32; result_len as usize],
193            // Combine two histograms
194            |mut global_counts, local_counts| {
195                for (g, l) in global_counts.iter_mut().zip(local_counts) {
196                    *g += l;
197                }
198                global_counts
199            },
200        )
201}
202
203impl BincountOp for Tensor {
204    fn bincount(&self, minlength: u32) -> Result<Vec<u32>> {
205        let values = self.to_vec1::<u32>()?;
206
207        Ok(bincount(&values, minlength))
208    }
209}
210
211fn masked_fill(on_false: &Tensor, mask: &Tensor, on_true: f32) -> Result<Tensor> {
212    let shape = mask.shape();
213    let on_true = Tensor::new(on_true, on_false.device())?.broadcast_as(shape.dims())?;
214    let m = mask.where_cond(&on_true, on_false)?;
215    Ok(m)
216}
217
218#[doc(hidden)]
219#[macro_export]
220macro_rules! serde_default_fn {
221    ($t:ty, $name:ident, $v:expr) => {
222        fn $name() -> $t {
223            $v
224        }
225    };
226}
227
228serde_default_fn!(f64, routed_scaling_factor, 1.0);
229serde_default_fn!(TopkMethod, topk_method, TopkMethod::Greedy);
230serde_default_fn!(usize, moe_layer_freq, 1);
231serde_default_fn!(usize, first_k_dense_replace, 0);
232serde_default_fn!(bool, norm_topk_prob, false);
233serde_default_fn!(ScoringFunc, scoring_func, ScoringFunc::Softmax);
234serde_default_fn!(Activation, hidden_act, Activation::Silu);
235serde_default_fn!(bool, tie_word_embeddings, false);
236
237#[derive(Deserialize, Clone, Debug)]
238enum TopkMethod {
239    #[serde(rename = "greedy")]
240    Greedy,
241    #[serde(rename = "group_limited_greedy")]
242    GroupLimitedGreedy,
243}
244
245#[derive(Deserialize, Clone, Debug)]
246enum ScoringFunc {
247    #[serde(rename = "softmax")]
248    Softmax,
249}
250
251#[derive(Deserialize, Clone, Debug)]
252pub struct DeepSeekV2Config {
253    pub(crate) vocab_size: usize,
254    pub(crate) hidden_size: usize,
255    pub(crate) intermediate_size: usize,
256    pub(crate) moe_intermediate_size: usize,
257    pub(crate) num_hidden_layers: usize,
258    pub(crate) num_attention_heads: usize,
259    pub(crate) n_shared_experts: Option<usize>,
260    pub(crate) n_routed_experts: Option<usize>,
261    #[serde(default = "routed_scaling_factor")]
262    pub(crate) routed_scaling_factor: f64,
263    #[serde(default = "topk_method")]
264    topk_method: TopkMethod,
265    pub(crate) num_experts_per_tok: Option<usize>,
266    #[serde(default = "moe_layer_freq")]
267    pub(crate) moe_layer_freq: usize,
268    #[serde(default = "first_k_dense_replace")]
269    pub(crate) first_k_dense_replace: usize,
270    // k dense layers
271    #[serde(default = "norm_topk_prob")]
272    pub(crate) norm_topk_prob: bool,
273    #[serde(default = "scoring_func")]
274    scoring_func: ScoringFunc,
275    #[serde(default = "hidden_act")]
276    pub(crate) hidden_act: Activation,
277    pub(crate) max_position_embeddings: usize,
278    pub(crate) rms_norm_eps: f64,
279    #[serde(default = "tie_word_embeddings")]
280    pub(crate) tie_word_embeddings: bool,
281    pub(crate) rope_theta: f32,
282    pub(crate) rope_scaling: Option<DeepSeekV2RopeScaling>,
283    pub(crate) attention_bias: bool,
284    pub(crate) q_lora_rank: Option<usize>,
285    pub(crate) qk_rope_head_dim: usize,
286    pub(crate) kv_lora_rank: usize,
287    pub(crate) v_head_dim: usize,
288    pub(crate) qk_nope_head_dim: usize,
289    pub(crate) n_group: usize,
290    pub(crate) topk_group: usize,
291}
292
293#[derive(Debug, Clone, Deserialize)]
294#[serde(rename_all = "lowercase")]
295pub enum ScaledRopeType {
296    #[serde(alias = "su")]
297    #[serde(alias = "longrope")]
298    Su,
299    #[serde(alias = "yarn")]
300    Yarn,
301    #[serde(alias = "dynamic")]
302    Dynamic,
303    #[serde(alias = "linear")]
304    Linear,
305}
306
307#[derive(Debug, Clone)]
308pub struct DeepSeekV2RotaryEmbedding {
309    sin: Tensor,
310    cos: Tensor,
311}
312
313#[derive(Debug, Clone, Deserialize)]
314#[serde(untagged)]
315pub enum DeepSeekV2RopeScaling {
316    Yarn {
317        original_max_position_embeddings: usize,
318        beta_fast: f32,
319        beta_slow: f32,
320        mscale: f32,
321        mscale_all_dim: f32,
322        factor: f32,
323        #[serde(rename = "type")]
324        scaling_type: ScaledRopeType,
325    },
326    LinearOrDynamic {
327        #[serde(rename = "type")]
328        scaling_type: ScaledRopeType,
329        factor: f64,
330    },
331}
332
333pub struct DeepSeekV2RopeConfig {
334    pub rope_scaling: Option<DeepSeekV2RopeScaling>,
335    pub max_position_embeddings: usize,
336    pub rope_theta: f32,
337    pub qk_rope_head_dim: usize,
338}
339
340impl DeepSeekV2RotaryEmbedding {
341    fn new_unscaled(cfg: &DeepSeekV2RopeConfig, dtype: DType, dev: &Device) -> Result<Self> {
342        let max_seq_len = cfg.max_position_embeddings;
343        let dim = cfg.qk_rope_head_dim;
344
345        let inv_freq: Vec<_> = (0..dim)
346            .step_by(2)
347            .map(|i| 1f32 / cfg.rope_theta.powf(i as f32 / dim as f32))
348            .collect();
349        let inv_freq_len = inv_freq.len();
350        let inv_freq = Tensor::from_vec(inv_freq, (1, inv_freq_len), dev)?;
351        let t = Tensor::arange(0u32, max_seq_len as u32, dev)?
352            .to_dtype(DType::F32)?
353            .reshape((max_seq_len, 1))?;
354        let freqs = t.matmul(&inv_freq)?;
355
356        let sin = freqs.sin()?.to_dtype(dtype)?;
357        let cos = freqs.cos()?.to_dtype(dtype)?;
358
359        Ok(Self { sin, cos })
360    }
361
362    fn yarn_find_correction_dim(
363        num_rot: f32,
364        dim: usize,
365        base: f32,
366        max_position_embeddings: usize,
367    ) -> f32 {
368        (dim as f32 * (max_position_embeddings as f32 / (num_rot * 2. * PI)).ln())
369            / (2. * base.ln())
370    }
371
372    fn yarn_find_correction_range(
373        low_rot: f32,
374        high_rot: f32,
375        dim: usize,
376        base: f32,
377        max_position_embeddings: usize,
378    ) -> (f32, f32) {
379        let low =
380            Self::yarn_find_correction_dim(low_rot, dim, base, max_position_embeddings).floor();
381        let high =
382            Self::yarn_find_correction_dim(high_rot, dim, base, max_position_embeddings).ceil();
383        (low.max(0.), high.min(dim as f32 - 1.))
384    }
385
386    fn yarn_linear_ramp_mask(min: f32, mut max: f32, dim: usize, dev: &Device) -> Result<Tensor> {
387        if min == max {
388            // https://huggingface.co/deepseek-ai/DeepSeek-V2-Lite/blob/604d5664dddd88a0433dbae533b7fe9472482de0/modeling_deepseek.py#L255
389            max += 0.001;
390        }
391        let linear_func =
392            ((Tensor::arange(0f32, dim as f32, dev)? - min as f64)? / (max as f64 - min as f64))?;
393        linear_func.clamp(0., 1.)
394    }
395
396    pub(crate) fn yarn_get_mscale(scale: f32, mscale: f32) -> f32 {
397        if scale <= 1. {
398            return 1.;
399        }
400        0.1 * mscale * scale.ln() + 1.
401    }
402
403    #[allow(clippy::too_many_arguments)]
404    fn new_yarn(
405        cfg: &DeepSeekV2RopeConfig,
406        dtype: DType,
407        dev: &Device,
408        original_max_position_embeddings: usize,
409        beta_fast: f32,
410        beta_slow: f32,
411        factor: f32,
412        mscale: f32,
413        mscale_all_dim: f32,
414    ) -> Result<Self> {
415        let freq_extra: Vec<_> = (0..cfg.qk_rope_head_dim)
416            .step_by(2)
417            .map(|i| 1f32 / cfg.rope_theta.powf(i as f32 / cfg.qk_rope_head_dim as f32))
418            .collect();
419        let freq_extra_len = freq_extra.len();
420        let freq_extra = Tensor::from_vec(freq_extra, freq_extra_len, dev)?;
421        let freq_inter: Vec<_> = (0..cfg.qk_rope_head_dim)
422            .step_by(2)
423            .map(|i| 1f32 / (factor * cfg.rope_theta.powf(i as f32 / cfg.qk_rope_head_dim as f32)))
424            .collect();
425        let freq_inter_len = freq_inter.len();
426        let freq_inter = Tensor::from_vec(freq_inter, (1, freq_inter_len), dev)?;
427
428        let (low, high) = Self::yarn_find_correction_range(
429            beta_fast,
430            beta_slow,
431            cfg.qk_rope_head_dim,
432            cfg.rope_theta,
433            original_max_position_embeddings,
434        );
435        let inv_freq_mask =
436            (1. - Self::yarn_linear_ramp_mask(low, high, cfg.qk_rope_head_dim / 2, dev)?)?;
437        let inv_freq = freq_inter
438            .broadcast_mul(&(1. - &inv_freq_mask)?)?
439            .broadcast_add(&freq_extra.broadcast_mul(&inv_freq_mask)?)?;
440
441        let t = Tensor::arange(0u32, cfg.max_position_embeddings as u32, dev)?
442            .to_dtype(DType::F32)?
443            .reshape((cfg.max_position_embeddings, 1))?;
444        let freqs = t.matmul(&inv_freq)?;
445
446        let mscale =
447            Self::yarn_get_mscale(factor, mscale) / Self::yarn_get_mscale(factor, mscale_all_dim);
448        let sin = (freqs.sin()? * mscale as f64)?.to_dtype(dtype)?;
449        let cos = (freqs.cos()? * mscale as f64)?.to_dtype(dtype)?;
450
451        Ok(Self { sin, cos })
452    }
453
454    pub fn new(cfg: &DeepSeekV2RopeConfig, dtype: DType, dev: &Device) -> Result<Self> {
455        match &cfg.rope_scaling {
456            Some(DeepSeekV2RopeScaling::LinearOrDynamic {
457                scaling_type: _,
458                factor: _,
459            }) => candle::bail!("linear and dynamic rope are not implemented yet!"),
460            Some(DeepSeekV2RopeScaling::Yarn {
461                original_max_position_embeddings,
462                beta_fast,
463                beta_slow,
464                factor,
465                mscale,
466                mscale_all_dim,
467                scaling_type: _,
468            }) => Self::new_yarn(
469                cfg,
470                dtype,
471                dev,
472                *original_max_position_embeddings,
473                *beta_fast,
474                *beta_slow,
475                *factor,
476                *mscale,
477                *mscale_all_dim,
478            ),
479            None => Self::new_unscaled(cfg, dtype, dev),
480        }
481    }
482
483    pub fn forward(
484        &self,
485        q: &Tensor,
486        k: &Tensor,
487        seqlen_offset: usize,
488    ) -> Result<(Tensor, Tensor)> {
489        let (_b_sz, _h, seq_len, _n_embd) = q.dims4()?;
490
491        let sin = self.sin.narrow(0, seqlen_offset, seq_len)?;
492        let cos = self.cos.narrow(0, seqlen_offset, seq_len)?;
493
494        let q_embed = candle_nn::rotary_emb::rope_i(&q.contiguous()?, &cos, &sin)?;
495        let k_embed = candle_nn::rotary_emb::rope_i(&k.contiguous()?, &cos, &sin)?;
496
497        Ok((q_embed, k_embed))
498    }
499}
500
501impl DeepSeekV2Config {
502    pub(crate) fn q_head_dim(&self) -> usize {
503        self.qk_rope_head_dim + self.qk_nope_head_dim
504    }
505
506    fn softmax_scale(&self) -> f32 {
507        let mut softmax_scale = 1.0 / (self.q_head_dim() as f32).sqrt();
508        if let Some(DeepSeekV2RopeScaling::Yarn {
509            mscale_all_dim,
510            factor,
511            ..
512        }) = self.rope_scaling
513        {
514            let mscale = DeepSeekV2RotaryEmbedding::yarn_get_mscale(factor, mscale_all_dim);
515            softmax_scale = softmax_scale * mscale * mscale;
516        }
517        softmax_scale
518    }
519}
520
521enum QProj {
522    Plain(Linear),
523    Lora { a: Linear, norm: RmsNorm, b: Linear },
524}
525
526impl QProj {
527    fn forward(&self, xs: &Tensor) -> Result<Tensor> {
528        match self {
529            Self::Lora { a, norm, b } => b.forward(&norm.forward(&a.forward(xs)?)?),
530            Self::Plain(lin) => lin.forward(xs),
531        }
532    }
533}
534
535struct Attention {
536    q: QProj,
537    kv_a_proj_with_mqa: Linear,
538    kv_a_layernorm: RmsNorm,
539    kv_b_proj: Linear,
540    o_proj: Linear,
541    rotary_emb: Arc<DeepSeekV2RotaryEmbedding>,
542    cfg: DeepSeekV2Config,
543    q_head_dim: usize,
544    softmax_scale: f64,
545    kv_cache: Option<(Tensor, Tensor)>,
546}
547
548impl Attention {
549    fn new(
550        rotary_emb: Arc<DeepSeekV2RotaryEmbedding>,
551        cfg: &DeepSeekV2Config,
552        vb: VarBuilder,
553    ) -> Result<Self> {
554        let q_head_dim = cfg.q_head_dim();
555        let q = match cfg.q_lora_rank {
556            Some(lora_rank) => {
557                let a = candle_nn::linear_b(
558                    cfg.hidden_size,
559                    lora_rank,
560                    cfg.attention_bias,
561                    vb.pp("q_a_proj"),
562                )?;
563                let norm = rms_norm(lora_rank, cfg.rms_norm_eps, vb.pp("q_a_layernorm"))?;
564                let b = candle_nn::linear_no_bias(
565                    lora_rank,
566                    cfg.num_attention_heads * q_head_dim,
567                    vb.pp("q_b_proj"),
568                )?;
569                QProj::Lora { a, norm, b }
570            }
571            None => QProj::Plain(candle_nn::linear_no_bias(
572                cfg.hidden_size,
573                cfg.num_attention_heads * q_head_dim,
574                vb.pp("q_proj"),
575            )?),
576        };
577
578        let kv_a_proj_with_mqa = candle_nn::linear_b(
579            cfg.hidden_size,
580            cfg.kv_lora_rank + cfg.qk_rope_head_dim,
581            cfg.attention_bias,
582            vb.pp("kv_a_proj_with_mqa"),
583        )?;
584        let kv_a_layernorm = rms_norm(cfg.kv_lora_rank, cfg.rms_norm_eps, vb.pp("kv_a_layernorm"))?;
585        let kv_b_proj = candle_nn::linear_no_bias(
586            cfg.kv_lora_rank,
587            cfg.num_attention_heads * (q_head_dim - cfg.qk_rope_head_dim + cfg.v_head_dim),
588            vb.pp("kv_b_proj"),
589        )?;
590
591        let o_proj = candle_nn::linear_b(
592            cfg.num_attention_heads * cfg.v_head_dim,
593            cfg.hidden_size,
594            cfg.attention_bias,
595            vb.pp("o_proj"),
596        )?;
597
598        Ok(Self {
599            q,
600            kv_a_proj_with_mqa,
601            kv_a_layernorm,
602            kv_b_proj,
603            o_proj,
604            rotary_emb,
605            cfg: cfg.clone(),
606            q_head_dim,
607            softmax_scale: cfg.softmax_scale() as f64,
608            kv_cache: None,
609        })
610    }
611
612    fn forward(
613        &mut self,
614        xs: &Tensor,
615        attention_mask: Option<&Tensor>,
616        seqlen_offset: usize,
617    ) -> Result<Tensor> {
618        let (bs, seq_len, _) = xs.dims3()?;
619
620        let q = {
621            let q = self.q.forward(xs)?;
622            q.reshape((bs, seq_len, self.cfg.num_attention_heads, self.q_head_dim))?
623                .transpose(1, 2)?
624        };
625        let q_split = q.split(
626            &[self.cfg.qk_nope_head_dim, self.cfg.qk_rope_head_dim],
627            D::Minus1,
628        )?;
629        let q_nope = q_split[0].clone();
630        let q_pe = q_split[1].clone();
631
632        let compressed_kv = self.kv_a_proj_with_mqa.forward(xs)?;
633        let ckv_split = compressed_kv.split(
634            &[self.cfg.kv_lora_rank, self.cfg.qk_rope_head_dim],
635            D::Minus1,
636        )?;
637        let compressed_kv = ckv_split[0].clone();
638        let k_pe = {
639            let k_pe = ckv_split[1].clone();
640            k_pe.reshape((bs, seq_len, 1, self.cfg.qk_rope_head_dim))?
641                .transpose(1, 2)?
642        };
643        let kv = {
644            let kv = self
645                .kv_b_proj
646                .forward(&self.kv_a_layernorm.forward(&compressed_kv)?)?;
647            kv.reshape((
648                bs,
649                seq_len,
650                self.cfg.num_attention_heads,
651                self.cfg.qk_nope_head_dim + self.cfg.v_head_dim,
652            ))?
653            .transpose(1, 2)?
654        };
655
656        let kv_split = kv.split(&[self.cfg.qk_nope_head_dim, self.cfg.v_head_dim], D::Minus1)?;
657        let k_nope = kv_split[0].clone();
658        let v = kv_split[1].clone();
659
660        let (q_pe, k_pe) = self.rotary_emb.forward(&q_pe, &k_pe, seqlen_offset)?;
661
662        let q = Tensor::cat(&[q_nope, q_pe], D::Minus1)?;
663        let k = Tensor::cat(&[k_nope, k_pe.repeat((1, q.dim(1)?, 1, 1))?], D::Minus1)?;
664
665        let (k, v) = match &self.kv_cache {
666            None => (k, v),
667            Some((prev_k, prev_v)) => {
668                let key_states = Tensor::cat(&[prev_k, &k], 2)?;
669                let value_states = Tensor::cat(&[prev_v, &v], 2)?;
670                (key_states, value_states)
671            }
672        };
673        self.kv_cache = Some((k.clone(), v.clone()));
674
675        let attn_out = {
676            let att = (q.contiguous()?.matmul(&k.t()?.contiguous()?)? * self.softmax_scale)?;
677            let att = match attention_mask {
678                Some(mask) => att.broadcast_add(mask)?,
679                None => att,
680            };
681
682            let att = candle_nn::ops::softmax_last_dim(&att)?;
683            // Convert to contiguous as matmul doesn't support strided vs for now.
684            att.matmul(&v.contiguous()?)?
685        };
686
687        let attn_out = if attention_mask.is_some() {
688            attn_out.transpose(1, 2)?.reshape((bs, seq_len, ()))?
689        } else {
690            attn_out.reshape((bs, seq_len, ()))?
691        };
692
693        self.o_proj.forward(&attn_out)
694    }
695
696    fn clear_kv_cache(&mut self) {
697        self.kv_cache = None
698    }
699}
700
701struct Mlp {
702    gate: Linear,
703    up: Linear,
704    down: Linear,
705    act: Activation,
706}
707
708impl Mlp {
709    fn new(
710        cfg: &DeepSeekV2Config,
711        vb: VarBuilder,
712        hidden_size: Option<usize>,
713        intermediate_size: Option<usize>,
714    ) -> Result<Self> {
715        let hidden_size = hidden_size.unwrap_or(cfg.hidden_size);
716        let intermediate_size = intermediate_size.unwrap_or(cfg.intermediate_size);
717
718        Ok(Self {
719            gate: candle_nn::linear_no_bias(hidden_size, intermediate_size, vb.pp("gate_proj"))?,
720            up: candle_nn::linear_no_bias(hidden_size, intermediate_size, vb.pp("up_proj"))?,
721            down: candle_nn::linear_no_bias(intermediate_size, hidden_size, vb.pp("down_proj"))?,
722            act: cfg.hidden_act,
723        })
724    }
725
726    fn forward(&self, xs: &Tensor) -> Result<Tensor> {
727        let lhs = self.gate.forward(xs)?.apply(&self.act)?;
728        let rhs = self.up.forward(xs)?;
729        self.down.forward(&(&lhs * &rhs)?)
730    }
731}
732
733struct MoeGate {
734    weight: Tensor,
735    cfg: DeepSeekV2Config,
736    top_k: usize,
737    n_routed_experts: usize,
738}
739
740impl MoeGate {
741    fn new(cfg: &DeepSeekV2Config, vb: VarBuilder, n_routed_experts: usize) -> Result<Self> {
742        let weight = vb.get((n_routed_experts, cfg.hidden_size), "weight")?;
743        Ok(Self {
744            weight,
745            cfg: cfg.clone(),
746            top_k: cfg.num_experts_per_tok.unwrap(),
747            n_routed_experts,
748        })
749    }
750
751    /// (topk_idx, topk_weight)
752    fn forward(&self, xs: &Tensor) -> Result<(Tensor, Tensor)> {
753        let (bs, seq_len, h) = xs.dims3()?;
754        // Compute gating score
755        let xs = xs.reshape(((), h))?;
756        let logits = xs
757            .to_dtype(DType::F32)?
758            .broadcast_matmul(&self.weight.t()?.to_dtype(DType::F32)?)?;
759        let scores = match self.cfg.scoring_func {
760            ScoringFunc::Softmax => candle_nn::ops::softmax_last_dim(&logits)?,
761        };
762
763        // Select top-k experts
764        let (mut topk_weight, topk_idx) = match self.cfg.topk_method {
765            TopkMethod::Greedy => {
766                let TopKOutput { values, indices } = scores.topk_unsorted(self.top_k)?;
767                (values, indices)
768            }
769            TopkMethod::GroupLimitedGreedy => {
770                // (n, n_group)
771                let group_scores = scores
772                    .reshape((bs * seq_len, self.cfg.n_group, ()))?
773                    .max(D::Minus1)?;
774                // (n, topk_group)
775                let group_idx = scores.topk_unsorted(self.cfg.topk_group)?.indices;
776                // (n, n_group)
777                let group_mask = group_scores.zeros_like()?.scatter_add(
778                    &group_idx,
779                    &group_idx.ones_like()?.to_dtype(group_scores.dtype())?,
780                    1,
781                )?;
782                // (n, e)
783                let score_mask = group_mask
784                    .unsqueeze(D::Minus1)?
785                    .expand((
786                        bs * seq_len,
787                        self.cfg.n_group,
788                        self.n_routed_experts / self.cfg.n_group,
789                    ))?
790                    .reshape((bs, seq_len, ()))?;
791                // (n, e)
792                // Invert the mask
793                let tmp_scores = masked_fill(&score_mask, &(1. - &score_mask.ne(0.)?)?, 0.)?;
794                let TopKOutput { values, indices } = tmp_scores.topk_unsorted(self.top_k)?;
795                (values, indices)
796            }
797        };
798
799        if self.top_k > 1 && self.cfg.norm_topk_prob {
800            let denominator = (topk_weight.sum_keepdim(D::Minus1)? + 1e-20)?;
801            topk_weight = (topk_weight / denominator)?;
802        } else {
803            topk_weight = (topk_weight * self.cfg.routed_scaling_factor)?;
804        }
805        Ok((topk_idx, topk_weight))
806    }
807}
808
809struct Moe {
810    experts: Vec<Mlp>,
811    shared_experts: Option<Mlp>,
812    gate: MoeGate,
813}
814
815impl Moe {
816    fn new(
817        cfg: &DeepSeekV2Config,
818        vb: VarBuilder,
819
820        n_shared_experts: Option<usize>,
821        n_routed_experts: usize,
822    ) -> Result<Self> {
823        let mut experts = Vec::with_capacity(n_routed_experts);
824        for i in 0..n_routed_experts {
825            let vb_e = vb.pp("experts").pp(i);
826            experts.push(Mlp::new(cfg, vb_e, None, Some(cfg.moe_intermediate_size))?);
827        }
828        let shared_experts = if let Some(n_shared_experts) = n_shared_experts {
829            let intermediate_size = cfg.moe_intermediate_size * n_shared_experts;
830            Some(Mlp::new(
831                cfg,
832                vb.pp("shared_experts"),
833                None,
834                Some(intermediate_size),
835            )?)
836        } else {
837            None
838        };
839        let gate = MoeGate::new(cfg, vb.pp("gate"), n_routed_experts)?;
840        Ok(Self {
841            experts,
842            shared_experts,
843            gate,
844        })
845    }
846
847    fn moe_infer(&self, xs: &Tensor, topk_ids: &Tensor, topk_weight: &Tensor) -> Result<Tensor> {
848        let mut y = xs.zeros_like()?;
849        let counts = topk_ids
850            .flatten_all()?
851            .bincount(self.experts.len() as u32)?;
852        for (i, expert) in self.experts.iter().enumerate() {
853            if counts[i] == 0 {
854                continue;
855            }
856            let idx_top = topk_ids.eq(i as f64)?.nonzero()?.t()?;
857            let idx = &idx_top.i(0)?.contiguous()?;
858            let top = &idx_top.i(1)?.contiguous()?;
859
860            y = y.index_add(
861                idx,
862                &expert.forward(&xs.index_select(idx, 0)?)?.broadcast_mul(
863                    &topk_weight
864                        .index_select(idx, 0)?
865                        .gather(&top.unsqueeze(1)?, 1)?
866                        .squeeze(1)?
867                        .unsqueeze(D::Minus1)?
868                        .to_dtype(xs.dtype())?,
869                )?,
870                0,
871            )?;
872        }
873
874        Ok(y)
875    }
876
877    fn forward(&self, xs: &Tensor) -> Result<Tensor> {
878        let identity = xs.clone();
879        let orig_shape = xs.shape();
880        let (topk_idx, topk_weight) = self.gate.forward(xs)?;
881        let xs = xs.reshape(((), xs.dim(D::Minus1)?))?;
882
883        let mut y = self
884            .moe_infer(&xs, &topk_idx, &topk_weight)?
885            .reshape(orig_shape)?;
886        if let Some(ref shared_experts) = self.shared_experts {
887            y = (y + shared_experts.forward(&identity)?)?;
888        }
889        Ok(y)
890    }
891}
892
893enum MoeOrMlp {
894    Moe(Box<Moe>),
895    Mlp(Box<Mlp>),
896}
897
898impl MoeOrMlp {
899    fn forward(&self, xs: &Tensor) -> Result<Tensor> {
900        match self {
901            Self::Mlp(mlp) => mlp.forward(xs),
902            Self::Moe(moe) => moe.forward(xs),
903        }
904    }
905}
906
907struct DecoderLayer {
908    input_layernorm: RmsNorm,
909    post_attention_layernorm: RmsNorm,
910    attn: Attention,
911    moe_or_mlp: MoeOrMlp,
912}
913
914impl DecoderLayer {
915    fn new(
916        rotary_emb: Arc<DeepSeekV2RotaryEmbedding>,
917        cfg: &DeepSeekV2Config,
918        vb: VarBuilder,
919        layer_idx: usize,
920    ) -> Result<Self> {
921        let attn = Attention::new(rotary_emb, cfg, vb.pp("self_attn"))?;
922        let input_layernorm =
923            rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("input_layernorm"))?;
924        let post_attention_layernorm = rms_norm(
925            cfg.hidden_size,
926            cfg.rms_norm_eps,
927            vb.pp("post_attention_layernorm"),
928        )?;
929        let moe_or_mlp = if let Some(n_routed_experts) = cfg.n_routed_experts {
930            if layer_idx >= cfg.first_k_dense_replace
931                && layer_idx.is_multiple_of(cfg.moe_layer_freq)
932            {
933                MoeOrMlp::Moe(
934                    Moe::new(cfg, vb.pp("mlp"), cfg.n_shared_experts, n_routed_experts)?.into(),
935                )
936            } else {
937                MoeOrMlp::Mlp(Mlp::new(cfg, vb.pp("mlp"), None, None)?.into())
938            }
939        } else {
940            MoeOrMlp::Mlp(Mlp::new(cfg, vb.pp("mlp"), None, None)?.into())
941        };
942
943        Ok(Self {
944            input_layernorm,
945            post_attention_layernorm,
946            attn,
947            moe_or_mlp,
948        })
949    }
950
951    fn forward(
952        &mut self,
953        xs: &Tensor,
954        attention_mask: Option<&Tensor>,
955        seqlen_offset: usize,
956    ) -> Result<Tensor> {
957        let residual = xs;
958        let xs = self.input_layernorm.forward(xs)?;
959        let xs = self.attn.forward(&xs, attention_mask, seqlen_offset)?;
960        let xs = (xs + residual)?;
961        let residual = &xs;
962        let xs = self
963            .moe_or_mlp
964            .forward(&xs.apply(&self.post_attention_layernorm)?)?;
965        residual + xs
966    }
967
968    fn clear_kv_cache(&mut self) {
969        self.attn.clear_kv_cache();
970    }
971}
972
973pub struct DeepSeekV2 {
974    lm_head: Linear,
975    embed_tokens: Embedding,
976    norm: RmsNorm,
977    layers: Vec<DecoderLayer>,
978    dtype: DType,
979    device: Device,
980}
981
982impl DeepSeekV2 {
983    pub fn new(cfg: &DeepSeekV2Config, vb: VarBuilder) -> Result<Self> {
984        let vb_m = vb.pp("model");
985
986        let embed_tokens = embedding(cfg.vocab_size, cfg.hidden_size, vb_m.pp("embed_tokens"))?;
987        let lm_head = if !cfg.tie_word_embeddings {
988            candle_nn::linear_no_bias(cfg.hidden_size, cfg.vocab_size, vb.pp("lm_head"))?
989        } else {
990            candle_nn::Linear::new(embed_tokens.embeddings().clone(), None)
991        };
992        let norm = rms_norm(cfg.hidden_size, cfg.rms_norm_eps, vb_m.pp("norm"))?;
993
994        let rope_cfg = DeepSeekV2RopeConfig {
995            rope_scaling: cfg.rope_scaling.clone(),
996            max_position_embeddings: cfg.max_position_embeddings,
997            rope_theta: cfg.rope_theta,
998            qk_rope_head_dim: cfg.qk_rope_head_dim,
999        };
1000        let rotary_emb = Arc::new(DeepSeekV2RotaryEmbedding::new(
1001            &rope_cfg,
1002            vb.dtype(),
1003            vb.device(),
1004        )?);
1005
1006        let mut layers = Vec::with_capacity(cfg.num_hidden_layers);
1007        let vb_l = vb_m.pp("layers");
1008        for layer_idx in 0..cfg.num_hidden_layers {
1009            let layer = DecoderLayer::new(rotary_emb.clone(), cfg, vb_l.pp(layer_idx), layer_idx)?;
1010            layers.push(layer)
1011        }
1012
1013        Ok(Self {
1014            lm_head,
1015            embed_tokens,
1016            norm,
1017            layers,
1018            dtype: vb.dtype(),
1019            device: vb.device().clone(),
1020        })
1021    }
1022
1023    fn prepare_decoder_attention_mask(
1024        &self,
1025        b_size: usize,
1026        tgt_len: usize,
1027        seqlen_offset: usize,
1028    ) -> Result<Tensor> {
1029        let mask: Vec<_> = (0..tgt_len)
1030            .flat_map(|i| (0..tgt_len).map(move |j| if i < j { f32::NEG_INFINITY } else { 0. }))
1031            .collect();
1032        let mask = Tensor::from_slice(&mask, (tgt_len, tgt_len), &self.device)?;
1033        let mask = if seqlen_offset > 0 {
1034            let mask0 = Tensor::zeros((tgt_len, seqlen_offset), DType::F32, &self.device)?;
1035            Tensor::cat(&[&mask0, &mask], D::Minus1)?
1036        } else {
1037            mask
1038        };
1039        mask.expand((b_size, 1, tgt_len, tgt_len + seqlen_offset))?
1040            .to_dtype(self.dtype)
1041    }
1042
1043    pub fn forward(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
1044        let (bs, seq_len) = input_ids.dims2()?;
1045        let mut xs = self.embed_tokens.forward(input_ids)?;
1046        let attention_mask = if seq_len == 1 {
1047            None
1048        } else {
1049            let mask = self.prepare_decoder_attention_mask(bs, seq_len, seqlen_offset)?;
1050            Some(mask)
1051        };
1052        for layer in &mut self.layers {
1053            xs = layer.forward(
1054                &xs,
1055                attention_mask
1056                    .as_ref()
1057                    .map(|m| m.to_device(xs.device()).unwrap())
1058                    .as_ref(),
1059                seqlen_offset,
1060            )?;
1061        }
1062        let xs = xs.apply(&self.norm)?;
1063        let xs = xs.i((.., seq_len - 1, ..))?.contiguous()?;
1064        let logits = self.lm_head.forward(&xs)?;
1065        logits.to_dtype(DType::F32)
1066    }
1067
1068    pub fn clear_kv_cache(&mut self) {
1069        for layer in self.layers.iter_mut() {
1070            layer.clear_kv_cache();
1071        }
1072    }
1073}