Skip to main content

runtime/models_v2/
qwen.rs

1//! Qwen Model V2 - Clean implementation using solid abstractions
2//!
3//! This implements the Qwen/Qwen2 architecture which is used by 15+ models including:
4//! - Qwen2, Qwen-7B, Qwen-14B, CodeQwen, Qwen-VL, Qwen-Chat
5//! - Uses unified Tensor type from tensor_core
6//! - Implements Model trait from model_core
7//! - Supports loading via weight_loader_core
8
9use crate::model_config;
10use super::traits::*;
11use anyhow::Result;
12use serde::{Serialize, Deserialize};
13
14/// Qwen model configuration using the model_config macro
15model_config!(QwenConfig {
16    vocab_size: usize = 151936,
17    hidden_size: usize = 4096,
18    intermediate_size: usize = 11008,
19    num_hidden_layers: usize = 32,
20    num_attention_heads: usize = 32,
21    num_key_value_heads: usize = 32,
22    hidden_act: String = "silu".to_string(),
23    max_position_embeddings: usize = 32768,
24    initializer_range: f32 = 0.02,
25    rms_norm_eps: f32 = 1e-6,
26    use_cache: bool = true,
27    pad_token_id: i64 = 151643,
28    bos_token_id: i64 = 151643,
29    eos_token_id: i64 = 151643,
30    tie_word_embeddings: bool = false,
31    rope_theta: f32 = 1000000.0,
32    use_sliding_window: bool = false,
33    sliding_window: usize = 4096,
34    attention_dropout: f32 = 0.0,
35});
36
37impl QwenConfig {
38    /// Create QwenConfig from GGUF model configuration
39    pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
40        Self {
41            vocab_size: gguf.vocab_size,
42            hidden_size: gguf.hidden_size,
43            intermediate_size: gguf.intermediate_size,
44            num_hidden_layers: gguf.num_hidden_layers,
45            num_attention_heads: gguf.num_attention_heads,
46            num_key_value_heads: gguf.num_key_value_heads,
47            rms_norm_eps: gguf.rms_norm_eps,
48            rope_theta: gguf.rope_theta,
49            max_position_embeddings: gguf.max_position_embeddings,
50            ..Default::default()
51        }
52    }
53}
54
55/// Main Qwen model implementation
56pub struct QwenModelV2 {
57    config: QwenConfig,
58    device: Device,
59
60    // Model components using unified Tensor type
61    embed_tokens: Tensor,
62    layers: Vec<QwenLayer>,
63    norm: Tensor,
64    lm_head: Tensor,
65}
66
67/// Qwen transformer layer
68pub struct QwenLayer {
69    self_attn: QwenAttention,
70    mlp: QwenMLP,
71    input_layernorm: Tensor,
72    post_attention_layernorm: Tensor,
73}
74
75/// Qwen attention mechanism
76pub struct QwenAttention {
77    q_proj: Tensor,
78    k_proj: Tensor,
79    v_proj: Tensor,
80    o_proj: Tensor,
81    num_heads: usize,
82    num_key_value_heads: usize,
83    head_dim: usize,
84    scale: f32,
85}
86
87/// Qwen MLP (feed-forward network)
88pub struct QwenMLP {
89    gate_proj: Tensor,
90    up_proj: Tensor,
91    down_proj: Tensor,
92    hidden_act: String,
93}
94
95impl Model for QwenModelV2 {
96    type Config = QwenConfig;
97
98    fn new(config: QwenConfig) -> Result<Self> {
99        let device = Device::CPU;
100
101        let embed_tokens = ops_fn::zeros(
102            &[config.vocab_size, config.hidden_size],
103            DataType::Float32,
104            &device
105        )?;
106
107        let norm = ops_fn::zeros(
108            &[config.hidden_size],
109            DataType::Float32,
110            &device
111        )?;
112
113        let lm_head = if config.tie_word_embeddings {
114            embed_tokens.clone()
115        } else {
116            ops_fn::zeros(
117                &[config.hidden_size, config.vocab_size],
118                DataType::Float32,
119                &device
120            )?
121        };
122
123        // Create transformer layers
124        let mut layers = Vec::with_capacity(config.num_hidden_layers);
125        for _ in 0..config.num_hidden_layers {
126            layers.push(QwenLayer::new(&config, &device)?);
127        }
128
129        Ok(Self {
130            config,
131            device,
132            embed_tokens,
133            layers,
134            norm,
135            lm_head,
136        })
137    }
138
139    fn from_weights(config: QwenConfig, weights: ModelWeights) -> Result<Self> {
140        let mut model = Self::new(config)?;
141
142        // Load weights from the unified weight container
143        // Note: embedding weight is not transposed (used for index lookup)
144        if let Some(embed_weights) = weights.get("model.embed_tokens.weight") {
145            model.embed_tokens = embed_weights.clone();
146        }
147
148        if let Some(norm_weights) = weights.get("model.norm.weight") {
149            model.norm = norm_weights.clone();
150        }
151
152        // lm_head weight needs transpose: [vocab, hidden] -> [hidden, vocab] for matmul
153        if let Some(lm_head_weights) = weights.get("lm_head.weight") {
154            model.lm_head = ops_fn::transpose(lm_head_weights)?;
155        }
156
157        // Load layer weights
158        for (i, layer) in model.layers.iter_mut().enumerate() {
159            layer.load_weights(&weights, i)?;
160        }
161
162        Ok(model)
163    }
164
165    fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
166        match inputs {
167            ModelInputs::Text { input_ids, attention_mask, .. } => {
168                // 1. Token embedding
169                let mut hidden_states = ops_fn::embedding(input_ids, &self.embed_tokens)?;
170
171                // 2. Apply transformer layers (with RoPE)
172                for layer in &self.layers {
173                    hidden_states = layer.forward(&hidden_states, attention_mask.as_ref(), self.config.rope_theta)?;
174                }
175
176                // 3. Final layer norm
177                hidden_states = ops_fn::layer_norm(&hidden_states, &self.norm, None, self.config.rms_norm_eps)?;
178
179                // 4. Language modeling head
180                let logits = ops_fn::matmul(&hidden_states, &self.lm_head)?;
181
182                Ok(ModelOutputs::Logits {
183                    logits,
184                    hidden_states: None,
185                })
186            }
187            ModelInputs::Multimodal { input_ids, .. } => {
188                let text_inputs = ModelInputs::Text {
189                    input_ids: input_ids.clone(),
190                    attention_mask: None,
191                    position_ids: None,
192                };
193                self.forward(&text_inputs)
194            }
195            _ => Err(anyhow::anyhow!("Qwen model only supports text and multimodal inputs")),
196        }
197    }
198
199    fn generate(&self, prompt: &str, config: &GenerationConfig) -> Result<String> {
200        use crate::tokenizer::Tokenizer;
201        use rand::Rng;
202
203        // 1. Tokenize prompt
204        let tokenizer = Tokenizer::new();
205        let mut tokens: Vec<u32> = tokenizer.encode(prompt);
206
207        // 2. Generation loop
208        for _ in 0..config.max_new_tokens {
209            // Create input tensor from current tokens
210            let tokens_i64: Vec<i64> = tokens.iter().map(|&t| t as i64).collect();
211            let input_tensor = Tensor::from_i64_slice(&tokens_i64, &[1, tokens.len()], &self.device)?;
212
213            let inputs = ModelInputs::Text {
214                input_ids: input_tensor,
215                attention_mask: None,
216                position_ids: None,
217            };
218
219            // 3. Forward pass
220            let outputs = self.forward(&inputs)?;
221
222            // 4. Get logits and sample next token
223            let logits = match outputs {
224                ModelOutputs::Logits { logits, .. } => logits,
225                _ => return Err(anyhow::anyhow!("Expected logits output")),
226            };
227
228            // Get last token logits
229            let logits_candle = logits.to_candle()?;
230            let shape = logits_candle.dims();
231
232            // Extract last position logits [batch, seq, vocab] -> [vocab]
233            let last_logits = if shape.len() == 3 {
234                let seq_len = shape[1];
235                logits_candle
236                    .narrow(1, seq_len - 1, 1)?
237                    .squeeze(1)?
238                    .squeeze(0)?
239            } else {
240                let seq_len = shape[0];
241                logits_candle
242                    .narrow(0, seq_len - 1, 1)?
243                    .squeeze(0)?
244            };
245
246            // Convert to probabilities and sample
247            let logits_vec: Vec<f32> = last_logits.to_vec1()?;
248
249            let next_token = if config.do_sample && config.temperature > 0.0 {
250                // Temperature sampling
251                let scaled: Vec<f32> = logits_vec.iter()
252                    .map(|&x| x / config.temperature)
253                    .collect();
254
255                // Softmax
256                let max_val = scaled.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
257                let exp_sum: f32 = scaled.iter().map(|&x| (x - max_val).exp()).sum();
258                let probs: Vec<f32> = scaled.iter()
259                    .map(|&x| (x - max_val).exp() / exp_sum)
260                    .collect();
261
262                // Sample from distribution
263                let mut rng = rand::thread_rng();
264                let random_val: f32 = rng.gen();
265                let mut cumulative = 0.0;
266                let mut sampled = 0u32;
267
268                for (idx, &prob) in probs.iter().enumerate() {
269                    cumulative += prob;
270                    if random_val <= cumulative {
271                        sampled = idx as u32;
272                        break;
273                    }
274                }
275                sampled
276            } else {
277                // Greedy sampling
278                let mut max_idx = 0;
279                let mut max_val = logits_vec[0];
280                for (idx, &val) in logits_vec.iter().enumerate() {
281                    if val > max_val {
282                        max_val = val;
283                        max_idx = idx;
284                    }
285                }
286                max_idx as u32
287            };
288
289            // 5. Check for EOS
290            if next_token == config.eos_token_id {
291                break;
292            }
293
294            // 6. Append token
295            tokens.push(next_token);
296        }
297
298        // 7. Decode and return
299        Ok(tokenizer.decode(&tokens))
300    }
301
302    fn config(&self) -> &Self::Config {
303        &self.config
304    }
305
306    fn memory_requirements(&self) -> MemoryRequirements {
307        let param_size = self.config.vocab_size * self.config.hidden_size +
308                        self.config.num_hidden_layers * (
309                            4 * self.config.hidden_size * self.config.hidden_size +
310                            3 * self.config.hidden_size * self.config.intermediate_size
311                        );
312
313        let param_bytes = param_size * 4;
314        let kv_cache_bytes = 2 * self.config.num_hidden_layers *
315                           self.config.max_position_embeddings *
316                           self.config.hidden_size * 4;
317
318        MemoryRequirements {
319            gpu_memory: param_bytes,
320            cpu_memory: param_bytes / 4,
321            kv_cache_memory: kv_cache_bytes,
322            peak_memory: param_bytes + kv_cache_bytes,
323        }
324    }
325
326    fn to_device(&mut self, device: &Device) -> Result<()> {
327        self.embed_tokens = self.embed_tokens.to_device(device)?;
328        self.norm = self.norm.to_device(device)?;
329        self.lm_head = self.lm_head.to_device(device)?;
330
331        for layer in &mut self.layers {
332            layer.to_device(device)?;
333        }
334
335        self.device = device.clone();
336        Ok(())
337    }
338}
339
340impl QwenLayer {
341    fn new(config: &QwenConfig, device: &Device) -> Result<Self> {
342        let self_attn = QwenAttention::new(config, device)?;
343        let mlp = QwenMLP::new(config, device)?;
344
345        let input_layernorm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
346        let post_attention_layernorm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
347
348        Ok(Self {
349            self_attn,
350            mlp,
351            input_layernorm,
352            post_attention_layernorm,
353        })
354    }
355
356    fn forward(&self, hidden_states: &Tensor, attention_mask: Option<&Tensor>, rope_theta: f32) -> Result<Tensor> {
357        // 1. Pre-attention layer norm
358        let normed = ops_fn::layer_norm(hidden_states, &self.input_layernorm, None, 1e-6)?;
359
360        // 2. Self attention (with RoPE)
361        let attn_output = self.self_attn.forward(&normed, attention_mask, rope_theta)?;
362
363        // 3. Residual connection
364        let hidden_states = ops_fn::add(hidden_states, &attn_output)?;
365
366        // 4. Pre-MLP layer norm
367        let normed = ops_fn::layer_norm(&hidden_states, &self.post_attention_layernorm, None, 1e-6)?;
368
369        // 5. MLP
370        let mlp_output = self.mlp.forward(&normed)?;
371
372        // 6. Residual connection
373        let output = ops_fn::add(&hidden_states, &mlp_output)?;
374
375        Ok(output)
376    }
377
378    fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
379        let prefix = format!("model.layers.{}", layer_idx);
380
381        // Load attention weights (transpose for matmul: [out, in] -> [in, out])
382        if let Some(q_proj) = weights.get(&format!("{}.self_attn.q_proj.weight", prefix)) {
383            self.self_attn.q_proj = ops_fn::transpose(q_proj)?;
384        }
385        if let Some(k_proj) = weights.get(&format!("{}.self_attn.k_proj.weight", prefix)) {
386            self.self_attn.k_proj = ops_fn::transpose(k_proj)?;
387        }
388        if let Some(v_proj) = weights.get(&format!("{}.self_attn.v_proj.weight", prefix)) {
389            self.self_attn.v_proj = ops_fn::transpose(v_proj)?;
390        }
391        if let Some(o_proj) = weights.get(&format!("{}.self_attn.o_proj.weight", prefix)) {
392            self.self_attn.o_proj = ops_fn::transpose(o_proj)?;
393        }
394
395        // Load MLP weights (transpose for matmul: [out, in] -> [in, out])
396        if let Some(gate_proj) = weights.get(&format!("{}.mlp.gate_proj.weight", prefix)) {
397            self.mlp.gate_proj = ops_fn::transpose(gate_proj)?;
398        }
399        if let Some(up_proj) = weights.get(&format!("{}.mlp.up_proj.weight", prefix)) {
400            self.mlp.up_proj = ops_fn::transpose(up_proj)?;
401        }
402        if let Some(down_proj) = weights.get(&format!("{}.mlp.down_proj.weight", prefix)) {
403            self.mlp.down_proj = ops_fn::transpose(down_proj)?;
404        }
405
406        // Load layer norm weights (no transpose needed - 1D tensors)
407        if let Some(input_ln) = weights.get(&format!("{}.input_layernorm.weight", prefix)) {
408            self.input_layernorm = input_ln.clone();
409        }
410        if let Some(post_ln) = weights.get(&format!("{}.post_attention_layernorm.weight", prefix)) {
411            self.post_attention_layernorm = post_ln.clone();
412        }
413
414        Ok(())
415    }
416
417    fn to_device(&mut self, device: &Device) -> Result<()> {
418        self.self_attn.to_device(device)?;
419        self.mlp.to_device(device)?;
420        self.input_layernorm = self.input_layernorm.to_device(device)?;
421        self.post_attention_layernorm = self.post_attention_layernorm.to_device(device)?;
422        Ok(())
423    }
424}
425
426/// Apply Rotary Position Embedding (RoPE) to Q and K tensors
427fn apply_rope(
428    q: &candle_core::Tensor,
429    k: &candle_core::Tensor,
430    seq_len: usize,
431    head_dim: usize,
432    rope_theta: f32,
433) -> Result<(candle_core::Tensor, candle_core::Tensor)> {
434    let device = q.device();
435
436    // Compute inverse frequencies: 1 / (theta^(2i/d)) for i in [0, d/2)
437    let half_dim = head_dim / 2;
438    let inv_freq: Vec<f32> = (0..half_dim)
439        .map(|i| 1.0 / rope_theta.powf((2 * i) as f32 / head_dim as f32))
440        .collect();
441
442    // Create position indices [0, 1, 2, ..., seq_len-1]
443    let positions: Vec<f32> = (0..seq_len).map(|p| p as f32).collect();
444
445    // Compute angles: pos * inv_freq -> [seq_len, half_dim]
446    let mut angles = Vec::with_capacity(seq_len * half_dim);
447    for pos in &positions {
448        for freq in &inv_freq {
449            angles.push(pos * freq);
450        }
451    }
452
453    let angles_tensor = candle_core::Tensor::from_vec(angles, &[seq_len, half_dim], device)?;
454
455    // Compute cos and sin
456    let cos = angles_tensor.cos()?;
457    let sin = angles_tensor.sin()?;
458
459    // Reshape for broadcasting: [1, 1, seq_len, half_dim]
460    let cos = cos.unsqueeze(0)?.unsqueeze(0)?;
461    let sin = sin.unsqueeze(0)?.unsqueeze(0)?;
462
463    // Apply RoPE rotation
464    let q_half1 = q.narrow(3, 0, half_dim)?;
465    let q_half2 = q.narrow(3, half_dim, half_dim)?;
466    let k_half1 = k.narrow(3, 0, half_dim)?;
467    let k_half2 = k.narrow(3, half_dim, half_dim)?;
468
469    // Apply rotation
470    let q_rot1 = (q_half1.broadcast_mul(&cos)? - q_half2.broadcast_mul(&sin)?)?;
471    let q_rot2 = (q_half1.broadcast_mul(&sin)? + q_half2.broadcast_mul(&cos)?)?;
472    let k_rot1 = (k_half1.broadcast_mul(&cos)? - k_half2.broadcast_mul(&sin)?)?;
473    let k_rot2 = (k_half1.broadcast_mul(&sin)? + k_half2.broadcast_mul(&cos)?)?;
474
475    // Concatenate rotated halves
476    let q_rotated = candle_core::Tensor::cat(&[&q_rot1, &q_rot2], 3)?;
477    let k_rotated = candle_core::Tensor::cat(&[&k_rot1, &k_rot2], 3)?;
478
479    Ok((q_rotated, k_rotated))
480}
481
482impl QwenAttention {
483    fn new(config: &QwenConfig, device: &Device) -> Result<Self> {
484        let num_heads = config.num_attention_heads;
485        let num_key_value_heads = config.num_key_value_heads;
486        let head_dim = config.hidden_size / num_heads;
487        let scale = 1.0 / (head_dim as f32).sqrt();
488
489        let q_proj = ops_fn::zeros(&[config.hidden_size, num_heads * head_dim], DataType::Float32, device)?;
490        let k_proj = ops_fn::zeros(&[config.hidden_size, num_key_value_heads * head_dim], DataType::Float32, device)?;
491        let v_proj = ops_fn::zeros(&[config.hidden_size, num_key_value_heads * head_dim], DataType::Float32, device)?;
492        let o_proj = ops_fn::zeros(&[num_heads * head_dim, config.hidden_size], DataType::Float32, device)?;
493
494        Ok(Self {
495            q_proj,
496            k_proj,
497            v_proj,
498            o_proj,
499            num_heads,
500            num_key_value_heads,
501            head_dim,
502            scale,
503        })
504    }
505
506    fn forward(&self, hidden_states: &Tensor, _attention_mask: Option<&Tensor>, rope_theta: f32) -> Result<Tensor> {
507        // Get batch and sequence length from hidden_states shape
508        let shape = hidden_states.shape();
509        let (batch_size, seq_len, _hidden_size) = if shape.len() == 3 {
510            (shape[0], shape[1], shape[2])
511        } else if shape.len() == 2 {
512            (1, shape[0], shape[1])
513        } else {
514            return Err(anyhow::anyhow!("Invalid hidden_states shape: {:?}", shape));
515        };
516
517        // 1. Project to Q, K, V
518        let query_states = ops_fn::matmul(hidden_states, &self.q_proj)?;
519        let key_states = ops_fn::matmul(hidden_states, &self.k_proj)?;
520        let value_states = ops_fn::matmul(hidden_states, &self.v_proj)?;
521
522        // 2. Reshape for multi-head attention
523        let q_candle = query_states.to_candle()?;
524        let k_candle = key_states.to_candle()?;
525        let v_candle = value_states.to_candle()?;
526
527        let q_reshaped = q_candle
528            .reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?
529            .transpose(1, 2)?;
530
531        let k_reshaped = k_candle
532            .reshape(&[batch_size, seq_len, self.num_key_value_heads, self.head_dim])?
533            .transpose(1, 2)?;
534
535        let v_reshaped = v_candle
536            .reshape(&[batch_size, seq_len, self.num_key_value_heads, self.head_dim])?
537            .transpose(1, 2)?;
538
539        // 3. Apply RoPE (Rotary Position Embedding) to Q and K
540        let (q_with_rope, k_with_rope) = apply_rope(&q_reshaped, &k_reshaped, seq_len, self.head_dim, rope_theta)?;
541
542        // 4. Handle GQA (Grouped Query Attention) - repeat K/V heads to match Q heads
543        let num_groups = self.num_heads / self.num_key_value_heads;
544        let (k_expanded, v_expanded) = if num_groups > 1 {
545            let k_rep = k_with_rope
546                .unsqueeze(2)?
547                .broadcast_as(&[batch_size, self.num_key_value_heads, num_groups, seq_len, self.head_dim])?
548                .reshape(&[batch_size, self.num_heads, seq_len, self.head_dim])?;
549            let v_rep = v_reshaped
550                .unsqueeze(2)?
551                .broadcast_as(&[batch_size, self.num_key_value_heads, num_groups, seq_len, self.head_dim])?
552                .reshape(&[batch_size, self.num_heads, seq_len, self.head_dim])?;
553            (k_rep, v_rep)
554        } else {
555            (k_with_rope, v_reshaped)
556        };
557
558        // 5. Scaled dot-product attention
559        let k_t = k_expanded.transpose(2, 3)?;
560
561        let q_contiguous = q_with_rope.contiguous()?;
562        let k_contiguous = k_t.contiguous()?;
563
564        let scores = q_contiguous.matmul(&k_contiguous)?;
565        let scaled_scores = (scores * (self.scale as f64))?;
566
567        // Apply causal mask
568        let device = scaled_scores.device();
569        let causal_mask = {
570            let mut mask_data = vec![0.0f32; seq_len * seq_len];
571            for i in 0..seq_len {
572                for j in 0..seq_len {
573                    if j > i {
574                        mask_data[i * seq_len + j] = f32::NEG_INFINITY;
575                    }
576                }
577            }
578            candle_core::Tensor::from_vec(mask_data, &[1, 1, seq_len, seq_len], device)?
579        };
580
581        let masked_scores = scaled_scores.broadcast_add(&causal_mask)?;
582
583        // Softmax over last dimension
584        let attention_weights = candle_nn::ops::softmax_last_dim(&masked_scores)?;
585
586        // Apply attention to values
587        let v_contiguous = v_expanded.contiguous()?;
588        let attn_output = attention_weights.matmul(&v_contiguous)?;
589
590        // 6. Reshape back
591        let attn_output = attn_output
592            .transpose(1, 2)?
593            .reshape(&[batch_size, seq_len, self.num_heads * self.head_dim])?;
594
595        let attn_output = Tensor::from_candle(attn_output);
596
597        // 7. Output projection
598        let output = ops_fn::matmul(&attn_output, &self.o_proj)?;
599
600        Ok(output)
601    }
602
603    fn to_device(&mut self, device: &Device) -> Result<()> {
604        self.q_proj = self.q_proj.to_device(device)?;
605        self.k_proj = self.k_proj.to_device(device)?;
606        self.v_proj = self.v_proj.to_device(device)?;
607        self.o_proj = self.o_proj.to_device(device)?;
608        Ok(())
609    }
610}
611
612impl QwenMLP {
613    fn new(config: &QwenConfig, device: &Device) -> Result<Self> {
614        let gate_proj = ops_fn::zeros(&[config.hidden_size, config.intermediate_size], DataType::Float32, device)?;
615        let up_proj = ops_fn::zeros(&[config.hidden_size, config.intermediate_size], DataType::Float32, device)?;
616        let down_proj = ops_fn::zeros(&[config.intermediate_size, config.hidden_size], DataType::Float32, device)?;
617
618        Ok(Self {
619            gate_proj,
620            up_proj,
621            down_proj,
622            hidden_act: config.hidden_act.clone(),
623        })
624    }
625
626    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
627        // 1. Gate and up projections
628        let gate_output = ops_fn::matmul(hidden_states, &self.gate_proj)?;
629        let up_output = ops_fn::matmul(hidden_states, &self.up_proj)?;
630
631        // 2. Apply activation (SiLU for Qwen)
632        let gate_activated = match self.hidden_act.as_str() {
633            "silu" | "swish" => ops_fn::silu(&gate_output)?,
634            "gelu" => ops_fn::gelu(&gate_output)?,
635            _ => return Err(anyhow::anyhow!("Unsupported activation: {}", self.hidden_act)),
636        };
637
638        // 3. Element-wise multiplication (gating)
639        let gated = ops_fn::mul(&gate_activated, &up_output)?;
640
641        // 4. Down projection
642        let output = ops_fn::matmul(&gated, &self.down_proj)?;
643
644        Ok(output)
645    }
646
647    fn to_device(&mut self, device: &Device) -> Result<()> {
648        self.gate_proj = self.gate_proj.to_device(device)?;
649        self.up_proj = self.up_proj.to_device(device)?;
650        self.down_proj = self.down_proj.to_device(device)?;
651        Ok(())
652    }
653}
654
655#[cfg(test)]
656mod tests {
657    use super::*;
658
659    #[test]
660    fn test_qwen_model_creation() {
661        let config = QwenConfig {
662            vocab_size: 1000,
663            hidden_size: 128,
664            intermediate_size: 512,
665            num_hidden_layers: 2,
666            num_attention_heads: 8,
667            num_key_value_heads: 8,
668            ..Default::default()
669        };
670
671        let model = QwenModelV2::new(config).unwrap();
672        assert_eq!(model.config().vocab_size(), 1000);
673        assert_eq!(model.config().hidden_size(), 128);
674        assert_eq!(model.config().num_layers(), 2);
675    }
676
677    #[test]
678    fn test_qwen_forward_pass() {
679        let config = QwenConfig {
680            vocab_size: 100,
681            hidden_size: 64,
682            intermediate_size: 256,
683            num_hidden_layers: 1,
684            num_attention_heads: 4,
685            num_key_value_heads: 4,
686            ..Default::default()
687        };
688
689        let model = QwenModelV2::new(config).unwrap();
690        let input_ids = ops_fn::zeros(&[2, 8], DataType::Int64, &Device::CPU).unwrap();
691        let inputs = ModelInputs::text(input_ids);
692
693        let outputs = model.forward(&inputs).unwrap();
694        match outputs {
695            ModelOutputs::Logits { logits, .. } => {
696                assert_eq!(logits.shape(), &[2, 8, 100]);
697            }
698            _ => panic!("Expected logits output"),
699        }
700    }
701
702    #[test]
703    fn test_qwen_generation() {
704        let config = QwenConfig {
705            vocab_size: 256,
706            hidden_size: 64,
707            intermediate_size: 256,
708            num_hidden_layers: 1,
709            num_attention_heads: 4,
710            num_key_value_heads: 4,
711            ..Default::default()
712        };
713        let model = QwenModelV2::new(config).unwrap();
714        let gen_config = GenerationConfig {
715            max_new_tokens: 5,
716            ..Default::default()
717        };
718
719        let output = model.generate("Hello", &gen_config).unwrap();
720        assert!(!output.is_empty());
721    }
722}