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 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#[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 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 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 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(¤t_state)?;
170 let current_hidden_states = current_hidden_states.broadcast_mul(&selected_experts)?;
171 ys = ys.index_add(&top_x, ¤t_hidden_states, 0)?;
172 }
173
174 ys.reshape((b_size, seq_len, hidden_dim))
175 }
176}
177
178#[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 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 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}