Skip to main content

candle_transformers/models/
qwen3_moe.rs

1use crate::{
2    fused_moe::{FusedMoe, MoeCfg},
3    models::{
4        qwen3::{Config as Qwen3Config, Qwen3Attention, Qwen3MLP, Qwen3RotaryEmbedding},
5        with_tracing::{linear_no_bias, Linear, RmsNorm},
6    },
7};
8use candle::{DType, Device, Module, Result, Tensor, D};
9use candle_nn::{Activation, VarBuilder};
10use std::sync::Arc;
11
12#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
13pub struct Config {
14    pub vocab_size: usize,
15    pub hidden_size: usize,
16    pub intermediate_size: usize,
17    pub num_hidden_layers: usize,
18    pub num_attention_heads: usize,
19    pub head_dim: usize,
20    pub attention_bias: bool,
21    pub num_key_value_heads: usize,
22    pub max_position_embeddings: usize,
23    pub sliding_window: Option<usize>,
24    pub max_window_layers: usize,
25    pub tie_word_embeddings: bool,
26    pub rope_theta: f64,
27    pub rms_norm_eps: f64,
28    pub use_sliding_window: bool,
29    pub hidden_act: Activation,
30    // MoE specific configuration
31    pub decoder_sparse_step: usize,
32    pub moe_intermediate_size: usize,
33    pub num_experts_per_tok: usize,
34    pub num_experts: usize,
35    pub norm_topk_prob: bool,
36}
37
38impl From<&Config> for Qwen3Config {
39    fn from(val: &Config) -> Self {
40        Qwen3Config {
41            vocab_size: val.vocab_size,
42            hidden_size: val.hidden_size,
43            intermediate_size: val.intermediate_size,
44            num_hidden_layers: val.num_hidden_layers,
45            num_attention_heads: val.num_attention_heads,
46            head_dim: val.head_dim,
47            attention_bias: val.attention_bias,
48            num_key_value_heads: val.num_key_value_heads,
49            max_position_embeddings: val.max_position_embeddings,
50            sliding_window: val.sliding_window,
51            max_window_layers: val.max_window_layers,
52            tie_word_embeddings: val.tie_word_embeddings,
53            rope_theta: val.rope_theta,
54            rms_norm_eps: val.rms_norm_eps,
55            use_sliding_window: val.use_sliding_window,
56            hidden_act: val.hidden_act,
57        }
58    }
59}
60
61#[derive(Debug, Clone)]
62struct Qwen3MLPExpert {
63    gate_proj: Linear,
64    up_proj: Linear,
65    down_proj: Linear,
66    act_fn: Activation,
67}
68
69impl Qwen3MLPExpert {
70    fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
71        Ok(Self {
72            gate_proj: linear_no_bias(
73                cfg.hidden_size,
74                cfg.moe_intermediate_size,
75                vb.pp("gate_proj"),
76            )?,
77            up_proj: linear_no_bias(cfg.hidden_size, cfg.moe_intermediate_size, vb.pp("up_proj"))?,
78            down_proj: linear_no_bias(
79                cfg.moe_intermediate_size,
80                cfg.hidden_size,
81                vb.pp("down_proj"),
82            )?,
83            act_fn: cfg.hidden_act,
84        })
85    }
86}
87
88impl Module for Qwen3MLPExpert {
89    fn forward(&self, x: &Tensor) -> Result<Tensor> {
90        let lhs = x.apply(&self.gate_proj)?.apply(&self.act_fn)?;
91        let rhs = x.apply(&self.up_proj)?;
92        (lhs * rhs)?.apply(&self.down_proj)
93    }
94}
95
96// Qwen3 Sparse MoE Block implementation
97#[derive(Debug, Clone)]
98struct Qwen3SparseMoeBlock {
99    gate: Linear,
100    experts: Vec<Qwen3MLPExpert>,
101    norm_topk_prob: bool,
102    num_experts_per_tok: usize,
103}
104
105impl Qwen3SparseMoeBlock {
106    fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
107        let gate = linear_no_bias(cfg.hidden_size, cfg.num_experts, vb.pp("gate"))?;
108        let mut experts = Vec::with_capacity(cfg.num_experts);
109        let vb_e = vb.pp("experts");
110        for idx in 0..cfg.num_experts {
111            let expert = Qwen3MLPExpert::new(cfg, vb_e.pp(idx))?;
112            experts.push(expert)
113        }
114        Ok(Self {
115            gate,
116            experts,
117            norm_topk_prob: cfg.norm_topk_prob,
118            num_experts_per_tok: cfg.num_experts_per_tok,
119        })
120    }
121}
122
123impl Module for Qwen3SparseMoeBlock {
124    fn forward(&self, xs: &Tensor) -> Result<Tensor> {
125        let (b_size, seq_len, hidden_dim) = xs.dims3()?;
126        let xs = xs.reshape(((), hidden_dim))?;
127        let router_logits = xs.apply(&self.gate)?;
128        let routing_weights = candle_nn::ops::softmax_last_dim(&router_logits)?;
129
130        // Extract topk experts per token
131        let experts_per_tok = routing_weights
132            .arg_sort_last_dim(false)?
133            .narrow(D::Minus1, 0, self.num_experts_per_tok)?
134            .contiguous()?;
135        let routing_weights = routing_weights.gather(&experts_per_tok, D::Minus1)?;
136
137        // Extract needed data
138        let routing_weights = routing_weights.to_dtype(DType::F32)?.to_vec2::<f32>()?;
139        let experts_per_tok = experts_per_tok.to_vec2::<u32>()?;
140        let mut top_x = vec![vec![]; self.experts.len()];
141        let mut selected_experts = vec![vec![]; self.experts.len()];
142        for (row_idx, (rw, expert_idxs)) in routing_weights
143            .iter()
144            .zip(experts_per_tok.iter())
145            .enumerate()
146        {
147            let sum_rw = rw.iter().sum::<f32>();
148            for (&rw, &expert_idx) in rw.iter().zip(expert_idxs.iter()) {
149                top_x[expert_idx as usize].push(row_idx as u32);
150                let rw = if self.norm_topk_prob { rw / sum_rw } else { rw };
151                selected_experts[expert_idx as usize].push(rw)
152            }
153        }
154
155        // Process through experts
156        let mut ys = xs.zeros_like()?;
157        for (expert_idx, expert_layer) in self.experts.iter().enumerate() {
158            let top_x = &top_x[expert_idx];
159            if top_x.is_empty() {
160                continue;
161            }
162            let top_x = Tensor::new(top_x.as_slice(), xs.device())?;
163            let selected_experts =
164                Tensor::new(selected_experts[expert_idx].as_slice(), xs.device())?
165                    .reshape(((), 1))?
166                    .to_dtype(xs.dtype())?;
167
168            let current_state = xs.index_select(&top_x, 0)?.reshape(((), hidden_dim))?;
169            let current_hidden_states = expert_layer.forward(&current_state)?;
170            let current_hidden_states = current_hidden_states.broadcast_mul(&selected_experts)?;
171            ys = ys.index_add(&top_x, &current_hidden_states, 0)?;
172        }
173
174        ys.reshape((b_size, seq_len, hidden_dim))
175    }
176}
177
178// MLP or MoE decision enum
179#[derive(Debug, Clone)]
180enum Qwen3FeedForward {
181    Mlp(Qwen3MLP),
182    NaiveMoE(Qwen3SparseMoeBlock),
183    FusedMoE(FusedMoe),
184}
185
186impl Qwen3FeedForward {
187    fn forward(&self, xs: &Tensor, is_prefill: bool) -> Result<Tensor> {
188        match self {
189            Self::Mlp(m) => m.forward(xs),
190            Self::NaiveMoE(m) => m.forward(xs),
191            Self::FusedMoE(m) => m.forward(xs, is_prefill),
192        }
193    }
194}
195
196#[derive(Debug, Clone)]
197struct DecoderLayer {
198    self_attn: Qwen3Attention,
199    feed_forward: Qwen3FeedForward,
200    ln1: RmsNorm,
201    ln2: RmsNorm,
202}
203
204impl DecoderLayer {
205    fn new(
206        layer_idx: usize,
207        cfg: &Config,
208        rotary: Arc<Qwen3RotaryEmbedding>,
209        vb: VarBuilder,
210    ) -> Result<Self> {
211        let self_attn = Qwen3Attention::new(&cfg.into(), rotary, vb.pp("self_attn"))?;
212
213        let moe_cfg = MoeCfg {
214            hidden_size: cfg.hidden_size,
215            num_experts: cfg.num_experts,
216            num_experts_per_tok: cfg.num_experts_per_tok,
217            moe_intermediate_size: cfg.moe_intermediate_size,
218            norm_topk_prob: cfg.norm_topk_prob,
219            act: cfg.hidden_act,
220            decoder_sparse_step: None,
221        };
222        // Decide whether to use MoE or regular MLP based on layer_idx and decoder_sparse_step
223        let feed_forward =
224            if cfg.num_experts > 0 && (layer_idx + 1).is_multiple_of(cfg.decoder_sparse_step) {
225                if cfg!(feature = "cuda") {
226                    // Use fused MoE kernel on CUDA
227                    Qwen3FeedForward::FusedMoE(FusedMoe::new(&moe_cfg, vb.pp("mlp"), vb.dtype())?)
228                } else {
229                    Qwen3FeedForward::NaiveMoE(Qwen3SparseMoeBlock::new(cfg, vb.pp("mlp"))?)
230                }
231            } else {
232                Qwen3FeedForward::Mlp(Qwen3MLP::new(&cfg.into(), vb.pp("mlp"))?)
233            };
234
235        let ln1 = RmsNorm::new(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("input_layernorm"))?;
236        let ln2 = RmsNorm::new(
237            cfg.hidden_size,
238            cfg.rms_norm_eps,
239            vb.pp("post_attention_layernorm"),
240        )?;
241
242        Ok(Self {
243            self_attn,
244            feed_forward,
245            ln1,
246            ln2,
247        })
248    }
249
250    fn forward(&mut self, x: &Tensor, mask: Option<&Tensor>, offset: usize) -> Result<Tensor> {
251        let h = self.ln1.forward(x)?;
252        let h = self.self_attn.forward(&h, mask, offset)?;
253        let x = (x + h)?;
254        let h2 = self.ln2.forward(&x)?;
255        let h2 = self.feed_forward.forward(&h2, mask.is_some())?;
256        x + h2
257    }
258
259    fn clear_kv_cache(&mut self) {
260        self.self_attn.clear_kv_cache();
261    }
262}
263
264#[derive(Debug, Clone)]
265pub struct Model {
266    embed_tokens: candle_nn::Embedding,
267    layers: Vec<DecoderLayer>,
268    norm: RmsNorm,
269    device: Device,
270    dtype: DType,
271}
272
273impl Model {
274    pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
275        let embed_tokens =
276            candle_nn::embedding(cfg.vocab_size, cfg.hidden_size, vb.pp("model.embed_tokens"))?;
277        let rotary = Arc::new(Qwen3RotaryEmbedding::new(
278            vb.dtype(),
279            &cfg.into(),
280            vb.device(),
281        )?);
282        let mut layers = Vec::with_capacity(cfg.num_hidden_layers);
283        let vb_l = vb.pp("model.layers");
284        for i in 0..cfg.num_hidden_layers {
285            layers.push(DecoderLayer::new(i, cfg, rotary.clone(), vb_l.pp(i))?);
286        }
287        Ok(Self {
288            embed_tokens,
289            layers,
290            norm: RmsNorm::new(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("model.norm"))?,
291            device: vb.device().clone(),
292            dtype: vb.dtype(),
293        })
294    }
295
296    fn clear_kv_cache(&mut self) {
297        for l in &mut self.layers {
298            l.clear_kv_cache();
299        }
300    }
301
302    fn causal_mask(
303        &self,
304        b: usize,
305        tgt: usize,
306        offset: usize,
307        sw: Option<usize>,
308    ) -> Result<Tensor> {
309        let minf = f32::NEG_INFINITY;
310        let mask: Vec<_> = (0..tgt)
311            .flat_map(|i| {
312                (0..(tgt + offset)).map(move |j| {
313                    let past_ok = j <= i + offset;
314                    let sw_ok = match sw {
315                        Some(w) => (i + offset) as i64 - j as i64 <= w as i64,
316                        None => true,
317                    };
318                    if past_ok && sw_ok {
319                        0.
320                    } else {
321                        minf
322                    }
323                })
324            })
325            .collect();
326        Tensor::from_slice(&mask, (b, 1, tgt, tgt + offset), &self.device)?.to_dtype(self.dtype)
327    }
328
329    pub fn forward(&mut self, input: &Tensor, offset: usize) -> Result<Tensor> {
330        let (b, l) = input.dims2()?;
331        let mut h = self.embed_tokens.forward(input)?;
332
333        let causal = if l == 1 {
334            None
335        } else {
336            Some(self.causal_mask(b, l, offset, None)?)
337        };
338
339        for layer in &mut self.layers {
340            h = layer.forward(&h, causal.as_ref(), offset)?;
341        }
342        self.norm.forward(&h)
343    }
344}
345
346#[derive(Debug, Clone)]
347pub struct ModelForCausalLM {
348    base: Model,
349    lm_head: Linear,
350}
351
352impl ModelForCausalLM {
353    pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
354        let base = Model::new(cfg, vb.clone())?;
355        let lm_head = if cfg.tie_word_embeddings {
356            Linear::from_weights(base.embed_tokens.embeddings().clone(), None)
357        } else {
358            linear_no_bias(cfg.hidden_size, cfg.vocab_size, vb.pp("lm_head"))?
359        };
360        Ok(Self { base, lm_head })
361    }
362
363    pub fn forward(&mut self, input: &Tensor, offset: usize) -> Result<Tensor> {
364        let (_, l) = input.dims2()?;
365        self.base
366            .forward(input, offset)?
367            .narrow(1, l - 1, 1)?
368            .apply(&self.lm_head)
369    }
370
371    pub fn clear_kv_cache(&mut self) {
372        self.base.clear_kv_cache();
373    }
374}