1use crate::model_config;
10use super::traits::*;
11use anyhow::Result;
12use serde::{Serialize, Deserialize};
13
14model_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 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
55pub struct QwenModelV2 {
57 config: QwenConfig,
58 device: Device,
59
60 embed_tokens: Tensor,
62 layers: Vec<QwenLayer>,
63 norm: Tensor,
64 lm_head: Tensor,
65}
66
67pub struct QwenLayer {
69 self_attn: QwenAttention,
70 mlp: QwenMLP,
71 input_layernorm: Tensor,
72 post_attention_layernorm: Tensor,
73}
74
75pub 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
87pub 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 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 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 if let Some(lm_head_weights) = weights.get("lm_head.weight") {
154 model.lm_head = ops_fn::transpose(lm_head_weights)?;
155 }
156
157 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 let mut hidden_states = ops_fn::embedding(input_ids, &self.embed_tokens)?;
170
171 for layer in &self.layers {
173 hidden_states = layer.forward(&hidden_states, attention_mask.as_ref(), self.config.rope_theta)?;
174 }
175
176 hidden_states = ops_fn::layer_norm(&hidden_states, &self.norm, None, self.config.rms_norm_eps)?;
178
179 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 let tokenizer = Tokenizer::new();
205 let mut tokens: Vec<u32> = tokenizer.encode(prompt);
206
207 for _ in 0..config.max_new_tokens {
209 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 let outputs = self.forward(&inputs)?;
221
222 let logits = match outputs {
224 ModelOutputs::Logits { logits, .. } => logits,
225 _ => return Err(anyhow::anyhow!("Expected logits output")),
226 };
227
228 let logits_candle = logits.to_candle()?;
230 let shape = logits_candle.dims();
231
232 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 let logits_vec: Vec<f32> = last_logits.to_vec1()?;
248
249 let next_token = if config.do_sample && config.temperature > 0.0 {
250 let scaled: Vec<f32> = logits_vec.iter()
252 .map(|&x| x / config.temperature)
253 .collect();
254
255 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 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 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 if next_token == config.eos_token_id {
291 break;
292 }
293
294 tokens.push(next_token);
296 }
297
298 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 let normed = ops_fn::layer_norm(hidden_states, &self.input_layernorm, None, 1e-6)?;
359
360 let attn_output = self.self_attn.forward(&normed, attention_mask, rope_theta)?;
362
363 let hidden_states = ops_fn::add(hidden_states, &attn_output)?;
365
366 let normed = ops_fn::layer_norm(&hidden_states, &self.post_attention_layernorm, None, 1e-6)?;
368
369 let mlp_output = self.mlp.forward(&normed)?;
371
372 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 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 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 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
426fn 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 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 let positions: Vec<f32> = (0..seq_len).map(|p| p as f32).collect();
444
445 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 let cos = angles_tensor.cos()?;
457 let sin = angles_tensor.sin()?;
458
459 let cos = cos.unsqueeze(0)?.unsqueeze(0)?;
461 let sin = sin.unsqueeze(0)?.unsqueeze(0)?;
462
463 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 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 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 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 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 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 let (q_with_rope, k_with_rope) = apply_rope(&q_reshaped, &k_reshaped, seq_len, self.head_dim, rope_theta)?;
541
542 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 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 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 let attention_weights = candle_nn::ops::softmax_last_dim(&masked_scores)?;
585
586 let v_contiguous = v_expanded.contiguous()?;
588 let attn_output = attention_weights.matmul(&v_contiguous)?;
589
590 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 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 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 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 let gated = ops_fn::mul(&gate_activated, &up_output)?;
640
641 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}