1use crate::model_config;
12use super::traits::*;
13use anyhow::Result;
14use serde::{Serialize, Deserialize};
15
16model_config!(Qwen2VLConfig {
18 vocab_size: usize = 151936,
20 hidden_size: usize = 3584,
21 intermediate_size: usize = 18944,
22 num_hidden_layers: usize = 28,
23 num_attention_heads: usize = 28,
24 num_key_value_heads: usize = 4,
25 max_position_embeddings: usize = 32768,
26 rms_norm_eps: f32 = 1e-6,
27 rope_theta: f32 = 1000000.0,
28 tie_word_embeddings: bool = false,
29
30 vision_hidden_size: usize = 1280,
32 vision_intermediate_size: usize = 5120,
33 vision_num_hidden_layers: usize = 32,
34 vision_num_attention_heads: usize = 16,
35 vision_patch_size: usize = 14,
36 vision_image_size: usize = 448,
37 vision_temporal_patch_size: usize = 2,
38
39 projector_hidden_size: usize = 0, pad_token_id: i64 = 151643,
44 bos_token_id: i64 = 151643,
45 eos_token_id: i64 = 151645,
46 image_token_id: i64 = 151655,
47 video_token_id: i64 = 151656,
48});
49
50impl Qwen2VLConfig {
51 pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
52 Self {
53 vocab_size: gguf.vocab_size,
54 hidden_size: gguf.hidden_size,
55 intermediate_size: gguf.intermediate_size,
56 num_hidden_layers: gguf.num_hidden_layers,
57 num_attention_heads: gguf.num_attention_heads,
58 num_key_value_heads: gguf.num_key_value_heads,
59 rms_norm_eps: gguf.rms_norm_eps,
60 rope_theta: gguf.rope_theta,
61 max_position_embeddings: gguf.max_position_embeddings,
62 ..Default::default()
63 }
64 }
65
66 pub fn effective_projector_hidden_size(&self) -> usize {
67 if self.projector_hidden_size > 0 {
68 self.projector_hidden_size
69 } else {
70 self.hidden_size
71 }
72 }
73}
74
75pub struct Qwen2VLModelV2 {
77 config: Qwen2VLConfig,
78 device: Device,
79 vision_encoder: Qwen2VisionEncoder,
80 projector: Qwen2VLProjector,
81 language_model: Qwen2VLLanguageModel,
82}
83
84pub struct Qwen2VisionEncoder {
86 patch_embed: PatchEmbedding3D,
87 blocks: Vec<VisionTransformerBlock>,
88 merger: VisionMerger,
89 config: Qwen2VLConfig,
90}
91
92pub struct PatchEmbedding3D {
94 proj: Tensor,
95 temporal_patch_size: usize,
96 patch_size: usize,
97 hidden_size: usize,
98}
99
100pub struct VisionTransformerBlock {
102 norm1: Tensor,
103 attn: VisionAttention,
104 norm2: Tensor,
105 mlp: VisionMLP,
106}
107
108pub struct VisionAttention {
110 qkv: Tensor,
111 proj: Tensor,
112 num_heads: usize,
113 head_dim: usize,
114}
115
116pub struct VisionMLP {
118 fc1: Tensor,
119 fc2: Tensor,
120}
121
122pub struct VisionMerger {
124 mlp: Vec<Tensor>,
125 hidden_size: usize,
126 target_hidden_size: usize,
127}
128
129pub struct Qwen2VLProjector {
131 linear1: Tensor,
132 linear2: Tensor,
133}
134
135pub struct Qwen2VLLanguageModel {
137 embed_tokens: Tensor,
138 layers: Vec<Qwen2VLDecoderLayer>,
139 norm: Tensor,
140 lm_head: Tensor,
141 config: Qwen2VLConfig,
142}
143
144pub struct Qwen2VLDecoderLayer {
146 self_attn: Qwen2VLAttention,
147 mlp: Qwen2VLMLP,
148 input_layernorm: Tensor,
149 post_attention_layernorm: Tensor,
150}
151
152pub struct Qwen2VLAttention {
154 q_proj: Tensor,
155 k_proj: Tensor,
156 v_proj: Tensor,
157 o_proj: Tensor,
158 num_heads: usize,
159 num_key_value_heads: usize,
160 head_dim: usize,
161 scale: f32,
162}
163
164pub struct Qwen2VLMLP {
166 gate_proj: Tensor,
167 up_proj: Tensor,
168 down_proj: Tensor,
169}
170
171impl Model for Qwen2VLModelV2 {
172 type Config = Qwen2VLConfig;
173
174 fn new(config: Qwen2VLConfig) -> Result<Self> {
175 let device = Device::CPU;
176
177 let vision_encoder = Qwen2VisionEncoder::new(&config, &device)?;
178 let projector = Qwen2VLProjector::new(&config, &device)?;
179 let language_model = Qwen2VLLanguageModel::new(&config, &device)?;
180
181 Ok(Self {
182 config,
183 device,
184 vision_encoder,
185 projector,
186 language_model,
187 })
188 }
189
190 fn from_weights(config: Qwen2VLConfig, weights: ModelWeights) -> Result<Self> {
191 let mut model = Self::new(config)?;
192
193 model.vision_encoder.load_weights(&weights)?;
194 model.projector.load_weights(&weights)?;
195 model.language_model.load_weights(&weights)?;
196
197 Ok(model)
198 }
199
200 fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
201 match inputs {
202 ModelInputs::Multimodal { input_ids, pixel_values, attention_mask, .. } => {
203 let image_embeds = if let Some(pixels) = pixel_values {
205 let vision_features = self.vision_encoder.forward(pixels)?;
206 Some(self.projector.forward(&vision_features)?)
207 } else {
208 None
209 };
210
211 let text_embeds = ops_fn::embedding(input_ids, &self.language_model.embed_tokens)?;
213
214 let hidden_states = if let Some(img_emb) = image_embeds {
216 self.merge_embeddings(&text_embeds, &img_emb, input_ids)?
217 } else {
218 text_embeds
219 };
220
221 let logits = self.language_model.forward(&hidden_states)?;
223
224 Ok(ModelOutputs::Logits {
225 logits,
226 hidden_states: None,
227 })
228 }
229 ModelInputs::Text { input_ids, .. } => {
230 let hidden_states = ops_fn::embedding(input_ids, &self.language_model.embed_tokens)?;
232 let logits = self.language_model.forward(&hidden_states)?;
233
234 Ok(ModelOutputs::Logits {
235 logits,
236 hidden_states: None,
237 })
238 }
239 _ => Err(anyhow::anyhow!("Qwen2-VL requires text or vision-language inputs")),
240 }
241 }
242
243 fn generate(&self, prompt: &str, config: &GenerationConfig) -> Result<String> {
244 use crate::tokenizer::Tokenizer;
245 use rand::Rng;
246
247 let tokenizer = Tokenizer::new();
248 let mut tokens: Vec<u32> = tokenizer.encode(prompt);
249
250 for _ in 0..config.max_new_tokens {
251 let input_ids = Tensor::from_i64_slice(
252 &tokens.iter().map(|&t| t as i64).collect::<Vec<_>>(),
253 &[1, tokens.len()],
254 &self.device
255 )?;
256
257 let inputs = ModelInputs::text(input_ids);
258 let outputs = self.forward(&inputs)?;
259
260 let logits = match outputs {
261 ModelOutputs::Logits { logits, .. } => logits,
262 _ => return Err(anyhow::anyhow!("Expected logits output")),
263 };
264
265 let logits_candle = logits.to_candle()?;
266 let last_logits = logits_candle.squeeze(0)?.narrow(0, logits_candle.dims()[1] - 1, 1)?.squeeze(0)?;
267 let logits_vec: Vec<f32> = last_logits.to_vec1()?;
268
269 let next_token = if config.do_sample && config.temperature > 0.0 {
270 let scaled: Vec<f32> = logits_vec.iter().map(|&x| x / config.temperature).collect();
271 let max_val = scaled.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
272 let exp_sum: f32 = scaled.iter().map(|&x| (x - max_val).exp()).sum();
273 let probs: Vec<f32> = scaled.iter().map(|&x| (x - max_val).exp() / exp_sum).collect();
274
275 let mut rng = rand::thread_rng();
276 let random_val: f32 = rng.gen();
277 let mut cumulative = 0.0;
278 let mut sampled = 0u32;
279
280 for (idx, &prob) in probs.iter().enumerate() {
281 cumulative += prob;
282 if random_val <= cumulative {
283 sampled = idx as u32;
284 break;
285 }
286 }
287 sampled
288 } else {
289 logits_vec.iter()
290 .enumerate()
291 .max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap())
292 .map(|(idx, _)| idx as u32)
293 .unwrap_or(0)
294 };
295
296 if next_token == config.eos_token_id {
297 break;
298 }
299
300 tokens.push(next_token);
301 }
302
303 Ok(tokenizer.decode(&tokens))
304 }
305
306 fn config(&self) -> &Self::Config { &self.config }
307
308 fn memory_requirements(&self) -> MemoryRequirements {
309 let vision_params = self.config.vision_hidden_size * self.config.vision_hidden_size * 4 *
310 self.config.vision_num_hidden_layers;
311 let text_params = self.config.hidden_size * self.config.hidden_size * 4 *
312 self.config.num_hidden_layers;
313 let param_size = (vision_params + text_params + self.config.vocab_size * self.config.hidden_size) * 4;
314
315 MemoryRequirements {
316 gpu_memory: param_size,
317 cpu_memory: param_size / 4,
318 kv_cache_memory: param_size / 8,
319 peak_memory: param_size + param_size / 2,
320 }
321 }
322
323 fn to_device(&mut self, device: &Device) -> Result<()> {
324 self.device = device.clone();
325 self.vision_encoder.to_device(device)?;
326 self.projector.to_device(device)?;
327 self.language_model.to_device(device)?;
328 Ok(())
329 }
330}
331
332impl Qwen2VLModelV2 {
333 fn merge_embeddings(&self, text_embeds: &Tensor, image_embeds: &Tensor, input_ids: &Tensor) -> Result<Tensor> {
334 Ok(text_embeds.clone())
337 }
338}
339
340impl Qwen2VisionEncoder {
341 fn new(config: &Qwen2VLConfig, device: &Device) -> Result<Self> {
342 let patch_embed = PatchEmbedding3D::new(config, device)?;
343
344 let mut blocks = Vec::with_capacity(config.vision_num_hidden_layers);
345 for _ in 0..config.vision_num_hidden_layers {
346 blocks.push(VisionTransformerBlock::new(config, device)?);
347 }
348
349 let merger = VisionMerger::new(config, device)?;
350
351 Ok(Self { patch_embed, blocks, merger, config: config.clone() })
352 }
353
354 fn forward(&self, pixel_values: &Tensor) -> Result<Tensor> {
355 let mut hidden_states = self.patch_embed.forward(pixel_values)?;
357
358 for block in &self.blocks {
360 hidden_states = block.forward(&hidden_states)?;
361 }
362
363 self.merger.forward(&hidden_states)
365 }
366
367 fn load_weights(&mut self, weights: &ModelWeights) -> Result<()> {
368 self.patch_embed.load_weights(weights)?;
369 for (i, block) in self.blocks.iter_mut().enumerate() {
370 block.load_weights(weights, i)?;
371 }
372 self.merger.load_weights(weights)?;
373 Ok(())
374 }
375
376 fn to_device(&mut self, device: &Device) -> Result<()> {
377 self.patch_embed.to_device(device)?;
378 for block in &mut self.blocks {
379 block.to_device(device)?;
380 }
381 self.merger.to_device(device)?;
382 Ok(())
383 }
384}
385
386impl PatchEmbedding3D {
387 fn new(config: &Qwen2VLConfig, device: &Device) -> Result<Self> {
388 let in_channels = 3 * config.vision_temporal_patch_size;
389 let proj = ops_fn::zeros(
390 &[in_channels * config.vision_patch_size * config.vision_patch_size, config.vision_hidden_size],
391 DataType::Float32,
392 device
393 )?;
394
395 Ok(Self {
396 proj,
397 temporal_patch_size: config.vision_temporal_patch_size,
398 patch_size: config.vision_patch_size,
399 hidden_size: config.vision_hidden_size,
400 })
401 }
402
403 fn forward(&self, pixel_values: &Tensor) -> Result<Tensor> {
404 let shape = pixel_values.shape();
407 let batch_size = shape[0];
408
409 let flat = pixel_values.to_candle()?.flatten(1, 4)?;
411 let flat = Tensor::from_candle(flat);
412 ops_fn::matmul(&flat, &self.proj)
413 }
414
415 fn load_weights(&mut self, weights: &ModelWeights) -> Result<()> {
416 if let Some(w) = weights.get("visual.patch_embed.proj.weight") {
417 let w_candle = w.to_candle()?;
419 let shape = w_candle.dims();
420 let flat = w_candle.reshape(&[shape[0], shape[1] * shape[2] * shape[3]])?;
421 self.proj = Tensor::from_candle(flat.t()?);
422 }
423 Ok(())
424 }
425
426 fn to_device(&mut self, device: &Device) -> Result<()> {
427 self.proj = self.proj.to_device(device)?;
428 Ok(())
429 }
430}
431
432impl VisionTransformerBlock {
433 fn new(config: &Qwen2VLConfig, device: &Device) -> Result<Self> {
434 let norm1 = ops_fn::zeros(&[config.vision_hidden_size], DataType::Float32, device)?;
435 let norm2 = ops_fn::zeros(&[config.vision_hidden_size], DataType::Float32, device)?;
436 let attn = VisionAttention::new(config, device)?;
437 let mlp = VisionMLP::new(config, device)?;
438
439 Ok(Self { norm1, attn, norm2, mlp })
440 }
441
442 fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
443 let residual = hidden_states.clone();
444 let normed = ops_fn::layer_norm(hidden_states, &self.norm1, None, 1e-6)?;
445 let attn_out = self.attn.forward(&normed)?;
446 let hidden_states = ops_fn::add(&residual, &attn_out)?;
447
448 let residual = hidden_states.clone();
449 let normed = ops_fn::layer_norm(&hidden_states, &self.norm2, None, 1e-6)?;
450 let mlp_out = self.mlp.forward(&normed)?;
451 ops_fn::add(&residual, &mlp_out)
452 }
453
454 fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
455 let prefix = format!("visual.blocks.{}", layer_idx);
456
457 if let Some(w) = weights.get(&format!("{}.norm1.weight", prefix)) {
458 self.norm1 = w.clone();
459 }
460 if let Some(w) = weights.get(&format!("{}.norm2.weight", prefix)) {
461 self.norm2 = w.clone();
462 }
463
464 self.attn.load_weights(weights, layer_idx)?;
465 self.mlp.load_weights(weights, layer_idx)?;
466
467 Ok(())
468 }
469
470 fn to_device(&mut self, device: &Device) -> Result<()> {
471 self.norm1 = self.norm1.to_device(device)?;
472 self.norm2 = self.norm2.to_device(device)?;
473 self.attn.to_device(device)?;
474 self.mlp.to_device(device)?;
475 Ok(())
476 }
477}
478
479impl VisionAttention {
480 fn new(config: &Qwen2VLConfig, device: &Device) -> Result<Self> {
481 let head_dim = config.vision_hidden_size / config.vision_num_attention_heads;
482 let qkv = ops_fn::zeros(
483 &[config.vision_hidden_size, config.vision_hidden_size * 3],
484 DataType::Float32,
485 device
486 )?;
487 let proj = ops_fn::zeros(
488 &[config.vision_hidden_size, config.vision_hidden_size],
489 DataType::Float32,
490 device
491 )?;
492
493 Ok(Self {
494 qkv,
495 proj,
496 num_heads: config.vision_num_attention_heads,
497 head_dim,
498 })
499 }
500
501 fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
502 let shape = hidden_states.shape();
503 let (batch_size, seq_len, _) = (shape[0], shape[1], shape[2]);
504
505 let qkv = ops_fn::matmul(hidden_states, &self.qkv)?;
506 let qkv_candle = qkv.to_candle()?;
507
508 let q = qkv_candle.narrow(2, 0, self.num_heads * self.head_dim)?;
509 let k = qkv_candle.narrow(2, self.num_heads * self.head_dim, self.num_heads * self.head_dim)?;
510 let v = qkv_candle.narrow(2, self.num_heads * self.head_dim * 2, self.num_heads * self.head_dim)?;
511
512 let q = q.reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?.transpose(1, 2)?;
513 let k = k.reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?.transpose(1, 2)?;
514 let v = v.reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?.transpose(1, 2)?;
515
516 let scale = (self.head_dim as f32).powf(-0.5);
517 let scores = q.contiguous()?.matmul(&k.transpose(2, 3)?.contiguous()?)?;
518 let scores = (scores * (scale as f64))?;
519
520 let attn_weights = candle_nn::ops::softmax_last_dim(&scores)?;
521 let attn_output = attn_weights.matmul(&v.contiguous()?)?;
522
523 let attn_output = attn_output
524 .transpose(1, 2)?
525 .reshape(&[batch_size, seq_len, self.num_heads * self.head_dim])?;
526
527 let attn_output = Tensor::from_candle(attn_output);
528 ops_fn::matmul(&attn_output, &self.proj)
529 }
530
531 fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
532 let prefix = format!("visual.blocks.{}.attn", layer_idx);
533
534 if let Some(w) = weights.get(&format!("{}.qkv.weight", prefix)) {
535 self.qkv = ops_fn::transpose(w)?;
536 }
537 if let Some(w) = weights.get(&format!("{}.proj.weight", prefix)) {
538 self.proj = ops_fn::transpose(w)?;
539 }
540
541 Ok(())
542 }
543
544 fn to_device(&mut self, device: &Device) -> Result<()> {
545 self.qkv = self.qkv.to_device(device)?;
546 self.proj = self.proj.to_device(device)?;
547 Ok(())
548 }
549}
550
551impl VisionMLP {
552 fn new(config: &Qwen2VLConfig, device: &Device) -> Result<Self> {
553 let fc1 = ops_fn::zeros(
554 &[config.vision_hidden_size, config.vision_intermediate_size],
555 DataType::Float32,
556 device
557 )?;
558 let fc2 = ops_fn::zeros(
559 &[config.vision_intermediate_size, config.vision_hidden_size],
560 DataType::Float32,
561 device
562 )?;
563
564 Ok(Self { fc1, fc2 })
565 }
566
567 fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
568 let hidden = ops_fn::matmul(hidden_states, &self.fc1)?;
569 let hidden = ops_fn::gelu(&hidden)?;
570 ops_fn::matmul(&hidden, &self.fc2)
571 }
572
573 fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
574 let prefix = format!("visual.blocks.{}.mlp", layer_idx);
575
576 if let Some(w) = weights.get(&format!("{}.fc1.weight", prefix)) {
577 self.fc1 = ops_fn::transpose(w)?;
578 }
579 if let Some(w) = weights.get(&format!("{}.fc2.weight", prefix)) {
580 self.fc2 = ops_fn::transpose(w)?;
581 }
582
583 Ok(())
584 }
585
586 fn to_device(&mut self, device: &Device) -> Result<()> {
587 self.fc1 = self.fc1.to_device(device)?;
588 self.fc2 = self.fc2.to_device(device)?;
589 Ok(())
590 }
591}
592
593impl VisionMerger {
594 fn new(config: &Qwen2VLConfig, device: &Device) -> Result<Self> {
595 let hidden_size = config.vision_hidden_size * 4; let target_hidden_size = config.hidden_size;
597
598 let mlp = vec![
599 ops_fn::zeros(&[hidden_size, target_hidden_size], DataType::Float32, device)?,
600 ops_fn::zeros(&[target_hidden_size, target_hidden_size], DataType::Float32, device)?,
601 ];
602
603 Ok(Self { mlp, hidden_size, target_hidden_size })
604 }
605
606 fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
607 let mut hidden = hidden_states.clone();
609 for (i, w) in self.mlp.iter().enumerate() {
610 hidden = ops_fn::matmul(&hidden, w)?;
611 if i < self.mlp.len() - 1 {
612 hidden = ops_fn::gelu(&hidden)?;
613 }
614 }
615 Ok(hidden)
616 }
617
618 fn load_weights(&mut self, weights: &ModelWeights) -> Result<()> {
619 for (i, w) in self.mlp.iter_mut().enumerate() {
620 if let Some(weight) = weights.get(&format!("visual.merger.mlp.{}.weight", i * 2)) {
621 *w = ops_fn::transpose(weight)?;
622 }
623 }
624 Ok(())
625 }
626
627 fn to_device(&mut self, device: &Device) -> Result<()> {
628 for w in &mut self.mlp {
629 *w = w.to_device(device)?;
630 }
631 Ok(())
632 }
633}
634
635impl Qwen2VLProjector {
636 fn new(config: &Qwen2VLConfig, device: &Device) -> Result<Self> {
637 let linear1 = ops_fn::zeros(
638 &[config.vision_hidden_size, config.effective_projector_hidden_size()],
639 DataType::Float32,
640 device
641 )?;
642 let linear2 = ops_fn::zeros(
643 &[config.effective_projector_hidden_size(), config.hidden_size],
644 DataType::Float32,
645 device
646 )?;
647
648 Ok(Self { linear1, linear2 })
649 }
650
651 fn forward(&self, vision_features: &Tensor) -> Result<Tensor> {
652 let hidden = ops_fn::matmul(vision_features, &self.linear1)?;
653 let hidden = ops_fn::gelu(&hidden)?;
654 ops_fn::matmul(&hidden, &self.linear2)
655 }
656
657 fn load_weights(&mut self, weights: &ModelWeights) -> Result<()> {
658 if let Some(w) = weights.get("visual.projector.0.weight") {
659 self.linear1 = ops_fn::transpose(w)?;
660 }
661 if let Some(w) = weights.get("visual.projector.2.weight") {
662 self.linear2 = ops_fn::transpose(w)?;
663 }
664 Ok(())
665 }
666
667 fn to_device(&mut self, device: &Device) -> Result<()> {
668 self.linear1 = self.linear1.to_device(device)?;
669 self.linear2 = self.linear2.to_device(device)?;
670 Ok(())
671 }
672}
673
674impl Qwen2VLLanguageModel {
675 fn new(config: &Qwen2VLConfig, device: &Device) -> Result<Self> {
676 let embed_tokens = ops_fn::zeros(&[config.vocab_size, config.hidden_size], DataType::Float32, device)?;
677 let norm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
678
679 let lm_head = if config.tie_word_embeddings {
680 embed_tokens.clone()
681 } else {
682 ops_fn::zeros(&[config.hidden_size, config.vocab_size], DataType::Float32, device)?
683 };
684
685 let mut layers = Vec::with_capacity(config.num_hidden_layers);
686 for _ in 0..config.num_hidden_layers {
687 layers.push(Qwen2VLDecoderLayer::new(config, device)?);
688 }
689
690 Ok(Self {
691 embed_tokens,
692 layers,
693 norm,
694 lm_head,
695 config: config.clone(),
696 })
697 }
698
699 fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
700 let mut hidden = hidden_states.clone();
701
702 for layer in &self.layers {
703 hidden = layer.forward(&hidden)?;
704 }
705
706 hidden = ops_fn::rms_norm(&hidden, &self.norm, self.config.rms_norm_eps)?;
707
708 if self.config.tie_word_embeddings {
709 let embed_t = ops_fn::transpose(&self.embed_tokens)?;
710 ops_fn::matmul(&hidden, &embed_t)
711 } else {
712 ops_fn::matmul(&hidden, &self.lm_head)
713 }
714 }
715
716 fn load_weights(&mut self, weights: &ModelWeights) -> Result<()> {
717 if let Some(w) = weights.get("model.embed_tokens.weight") {
718 self.embed_tokens = w.clone();
719 }
720 if let Some(w) = weights.get("model.norm.weight") {
721 self.norm = w.clone();
722 }
723 if !self.config.tie_word_embeddings {
724 if let Some(w) = weights.get("lm_head.weight") {
725 self.lm_head = ops_fn::transpose(w)?;
726 }
727 }
728
729 for (i, layer) in self.layers.iter_mut().enumerate() {
730 layer.load_weights(weights, i)?;
731 }
732
733 Ok(())
734 }
735
736 fn to_device(&mut self, device: &Device) -> Result<()> {
737 self.embed_tokens = self.embed_tokens.to_device(device)?;
738 self.norm = self.norm.to_device(device)?;
739 if !self.config.tie_word_embeddings {
740 self.lm_head = self.lm_head.to_device(device)?;
741 }
742 for layer in &mut self.layers {
743 layer.to_device(device)?;
744 }
745 Ok(())
746 }
747}
748
749impl Qwen2VLDecoderLayer {
750 fn new(config: &Qwen2VLConfig, device: &Device) -> Result<Self> {
751 let self_attn = Qwen2VLAttention::new(config, device)?;
752 let mlp = Qwen2VLMLP::new(config, device)?;
753 let input_layernorm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
754 let post_attention_layernorm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
755
756 Ok(Self {
757 self_attn,
758 mlp,
759 input_layernorm,
760 post_attention_layernorm,
761 })
762 }
763
764 fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
765 let residual = hidden_states.clone();
766 let hidden = ops_fn::rms_norm(hidden_states, &self.input_layernorm, 1e-6)?;
767 let hidden = self.self_attn.forward(&hidden)?;
768 let hidden = ops_fn::add(&residual, &hidden)?;
769
770 let residual = hidden.clone();
771 let hidden = ops_fn::rms_norm(&hidden, &self.post_attention_layernorm, 1e-6)?;
772 let hidden = self.mlp.forward(&hidden)?;
773 ops_fn::add(&residual, &hidden)
774 }
775
776 fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
777 let prefix = format!("model.layers.{}", layer_idx);
778
779 if let Some(w) = weights.get(&format!("{}.input_layernorm.weight", prefix)) {
780 self.input_layernorm = w.clone();
781 }
782 if let Some(w) = weights.get(&format!("{}.post_attention_layernorm.weight", prefix)) {
783 self.post_attention_layernorm = w.clone();
784 }
785
786 self.self_attn.load_weights(weights, layer_idx)?;
787 self.mlp.load_weights(weights, layer_idx)?;
788
789 Ok(())
790 }
791
792 fn to_device(&mut self, device: &Device) -> Result<()> {
793 self.input_layernorm = self.input_layernorm.to_device(device)?;
794 self.post_attention_layernorm = self.post_attention_layernorm.to_device(device)?;
795 self.self_attn.to_device(device)?;
796 self.mlp.to_device(device)?;
797 Ok(())
798 }
799}
800
801impl Qwen2VLAttention {
802 fn new(config: &Qwen2VLConfig, device: &Device) -> Result<Self> {
803 let head_dim = config.hidden_size / config.num_attention_heads;
804
805 let q_proj = ops_fn::zeros(&[config.hidden_size, config.num_attention_heads * head_dim], DataType::Float32, device)?;
806 let k_proj = ops_fn::zeros(&[config.hidden_size, config.num_key_value_heads * head_dim], DataType::Float32, device)?;
807 let v_proj = ops_fn::zeros(&[config.hidden_size, config.num_key_value_heads * head_dim], DataType::Float32, device)?;
808 let o_proj = ops_fn::zeros(&[config.num_attention_heads * head_dim, config.hidden_size], DataType::Float32, device)?;
809
810 Ok(Self {
811 q_proj,
812 k_proj,
813 v_proj,
814 o_proj,
815 num_heads: config.num_attention_heads,
816 num_key_value_heads: config.num_key_value_heads,
817 head_dim,
818 scale: (head_dim as f32).powf(-0.5),
819 })
820 }
821
822 fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
823 let shape = hidden_states.shape();
824 let (batch_size, seq_len, _) = (shape[0], shape[1], shape[2]);
825
826 let q = ops_fn::matmul(hidden_states, &self.q_proj)?;
827 let k = ops_fn::matmul(hidden_states, &self.k_proj)?;
828 let v = ops_fn::matmul(hidden_states, &self.v_proj)?;
829
830 let q = q.to_candle()?.reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?.transpose(1, 2)?;
831 let k = k.to_candle()?.reshape(&[batch_size, seq_len, self.num_key_value_heads, self.head_dim])?.transpose(1, 2)?;
832 let v = v.to_candle()?.reshape(&[batch_size, seq_len, self.num_key_value_heads, self.head_dim])?.transpose(1, 2)?;
833
834 let num_groups = self.num_heads / self.num_key_value_heads;
836 let (k, v) = if num_groups > 1 {
837 let k = k.unsqueeze(2)?
838 .broadcast_as(&[batch_size, self.num_key_value_heads, num_groups, seq_len, self.head_dim])?
839 .reshape(&[batch_size, self.num_heads, seq_len, self.head_dim])?;
840 let v = v.unsqueeze(2)?
841 .broadcast_as(&[batch_size, self.num_key_value_heads, num_groups, seq_len, self.head_dim])?
842 .reshape(&[batch_size, self.num_heads, seq_len, self.head_dim])?;
843 (k, v)
844 } else {
845 (k, v)
846 };
847
848 let scores = q.contiguous()?.matmul(&k.transpose(2, 3)?.contiguous()?)?;
849 let scores = (scores * (self.scale as f64))?;
850
851 let device = scores.device();
853 let mask = {
854 let mut mask_data = vec![0.0f32; seq_len * seq_len];
855 for i in 0..seq_len {
856 for j in (i + 1)..seq_len {
857 mask_data[i * seq_len + j] = f32::NEG_INFINITY;
858 }
859 }
860 candle_core::Tensor::from_vec(mask_data, &[1, 1, seq_len, seq_len], device)?
861 };
862
863 let scores = scores.broadcast_add(&mask)?;
864 let attn_weights = candle_nn::ops::softmax_last_dim(&scores)?;
865 let attn_output = attn_weights.matmul(&v.contiguous()?)?;
866
867 let attn_output = attn_output
868 .transpose(1, 2)?
869 .reshape(&[batch_size, seq_len, self.num_heads * self.head_dim])?;
870
871 let attn_output = Tensor::from_candle(attn_output);
872 ops_fn::matmul(&attn_output, &self.o_proj)
873 }
874
875 fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
876 let prefix = format!("model.layers.{}.self_attn", layer_idx);
877
878 if let Some(w) = weights.get(&format!("{}.q_proj.weight", prefix)) {
879 self.q_proj = ops_fn::transpose(w)?;
880 }
881 if let Some(w) = weights.get(&format!("{}.k_proj.weight", prefix)) {
882 self.k_proj = ops_fn::transpose(w)?;
883 }
884 if let Some(w) = weights.get(&format!("{}.v_proj.weight", prefix)) {
885 self.v_proj = ops_fn::transpose(w)?;
886 }
887 if let Some(w) = weights.get(&format!("{}.o_proj.weight", prefix)) {
888 self.o_proj = ops_fn::transpose(w)?;
889 }
890
891 Ok(())
892 }
893
894 fn to_device(&mut self, device: &Device) -> Result<()> {
895 self.q_proj = self.q_proj.to_device(device)?;
896 self.k_proj = self.k_proj.to_device(device)?;
897 self.v_proj = self.v_proj.to_device(device)?;
898 self.o_proj = self.o_proj.to_device(device)?;
899 Ok(())
900 }
901}
902
903impl Qwen2VLMLP {
904 fn new(config: &Qwen2VLConfig, device: &Device) -> Result<Self> {
905 let gate_proj = ops_fn::zeros(&[config.hidden_size, config.intermediate_size], DataType::Float32, device)?;
906 let up_proj = ops_fn::zeros(&[config.hidden_size, config.intermediate_size], DataType::Float32, device)?;
907 let down_proj = ops_fn::zeros(&[config.intermediate_size, config.hidden_size], DataType::Float32, device)?;
908
909 Ok(Self { gate_proj, up_proj, down_proj })
910 }
911
912 fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
913 let gate = ops_fn::matmul(hidden_states, &self.gate_proj)?;
914 let gate = ops_fn::silu(&gate)?;
915 let up = ops_fn::matmul(hidden_states, &self.up_proj)?;
916 let hidden = ops_fn::mul(&gate, &up)?;
917 ops_fn::matmul(&hidden, &self.down_proj)
918 }
919
920 fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
921 let prefix = format!("model.layers.{}.mlp", layer_idx);
922
923 if let Some(w) = weights.get(&format!("{}.gate_proj.weight", prefix)) {
924 self.gate_proj = ops_fn::transpose(w)?;
925 }
926 if let Some(w) = weights.get(&format!("{}.up_proj.weight", prefix)) {
927 self.up_proj = ops_fn::transpose(w)?;
928 }
929 if let Some(w) = weights.get(&format!("{}.down_proj.weight", prefix)) {
930 self.down_proj = ops_fn::transpose(w)?;
931 }
932
933 Ok(())
934 }
935
936 fn to_device(&mut self, device: &Device) -> Result<()> {
937 self.gate_proj = self.gate_proj.to_device(device)?;
938 self.up_proj = self.up_proj.to_device(device)?;
939 self.down_proj = self.down_proj.to_device(device)?;
940 Ok(())
941 }
942}
943
944#[cfg(test)]
945mod tests {
946 use super::*;
947
948 #[test]
949 fn test_qwen2vl_config() {
950 let config = Qwen2VLConfig::default();
951 assert_eq!(config.vocab_size, 151936);
952 assert_eq!(config.hidden_size, 3584);
953 assert_eq!(config.vision_hidden_size, 1280);
954 }
955
956 #[test]
957 fn test_qwen2vl_model_creation() {
958 let config = Qwen2VLConfig {
959 vocab_size: 1000,
960 hidden_size: 64,
961 intermediate_size: 256,
962 num_hidden_layers: 2,
963 num_attention_heads: 4,
964 num_key_value_heads: 2,
965 vision_hidden_size: 32,
966 vision_intermediate_size: 128,
967 vision_num_hidden_layers: 2,
968 vision_num_attention_heads: 2,
969 ..Default::default()
970 };
971
972 let model = Qwen2VLModelV2::new(config).unwrap();
973 assert_eq!(model.config().vocab_size(), 1000);
974 }
975
976 #[test]
977 fn test_qwen2vl_text_forward() {
978 let config = Qwen2VLConfig {
979 vocab_size: 100,
980 hidden_size: 32,
981 intermediate_size: 128,
982 num_hidden_layers: 1,
983 num_attention_heads: 2,
984 num_key_value_heads: 1,
985 vision_hidden_size: 16,
986 vision_intermediate_size: 64,
987 vision_num_hidden_layers: 1,
988 vision_num_attention_heads: 2,
989 ..Default::default()
990 };
991
992 let model = Qwen2VLModelV2::new(config).unwrap();
993 let input_ids = ops_fn::zeros(&[1, 4], DataType::Int64, &Device::CPU).unwrap();
994 let inputs = ModelInputs::text(input_ids);
995
996 let outputs = model.forward(&inputs).unwrap();
997 match outputs {
998 ModelOutputs::Logits { logits, .. } => {
999 assert_eq!(logits.shape(), &[1, 4, 100]);
1000 }
1001 _ => panic!("Expected logits output"),
1002 }
1003 }
1004}