Skip to main content

runtime/models_v2/
qwen2_vl.rs

1//! Qwen2-VL Model V2 - Vision-Language Model
2//!
3//! This implements the Qwen2-VL architecture which features:
4//! - ViT-based vision encoder
5//! - MLP projector to align vision/text embeddings
6//! - Qwen2 decoder for language modeling
7//! - Dynamic resolution support for images
8//!
9//! Supports: Qwen2-VL-2B, Qwen2-VL-7B, Qwen2-VL-72B
10
11use crate::model_config;
12use super::traits::*;
13use anyhow::Result;
14use serde::{Serialize, Deserialize};
15
16/// Qwen2-VL model configuration
17model_config!(Qwen2VLConfig {
18    // Language model config
19    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 encoder config
31    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 config
40    projector_hidden_size: usize = 0,  // 0 = auto: hidden_size
41
42    // Special tokens
43    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
75/// Main Qwen2-VL model
76pub struct Qwen2VLModelV2 {
77    config: Qwen2VLConfig,
78    device: Device,
79    vision_encoder: Qwen2VisionEncoder,
80    projector: Qwen2VLProjector,
81    language_model: Qwen2VLLanguageModel,
82}
83
84/// Qwen2 Vision Encoder (ViT-based)
85pub struct Qwen2VisionEncoder {
86    patch_embed: PatchEmbedding3D,
87    blocks: Vec<VisionTransformerBlock>,
88    merger: VisionMerger,
89    config: Qwen2VLConfig,
90}
91
92/// 3D patch embedding for spatial + temporal
93pub struct PatchEmbedding3D {
94    proj: Tensor,
95    temporal_patch_size: usize,
96    patch_size: usize,
97    hidden_size: usize,
98}
99
100/// Vision transformer block
101pub struct VisionTransformerBlock {
102    norm1: Tensor,
103    attn: VisionAttention,
104    norm2: Tensor,
105    mlp: VisionMLP,
106}
107
108/// Vision attention
109pub struct VisionAttention {
110    qkv: Tensor,
111    proj: Tensor,
112    num_heads: usize,
113    head_dim: usize,
114}
115
116/// Vision MLP
117pub struct VisionMLP {
118    fc1: Tensor,
119    fc2: Tensor,
120}
121
122/// Vision merger to reduce tokens
123pub struct VisionMerger {
124    mlp: Vec<Tensor>,
125    hidden_size: usize,
126    target_hidden_size: usize,
127}
128
129/// Projector to align vision and text
130pub struct Qwen2VLProjector {
131    linear1: Tensor,
132    linear2: Tensor,
133}
134
135/// Language model portion (Qwen2)
136pub struct Qwen2VLLanguageModel {
137    embed_tokens: Tensor,
138    layers: Vec<Qwen2VLDecoderLayer>,
139    norm: Tensor,
140    lm_head: Tensor,
141    config: Qwen2VLConfig,
142}
143
144/// Decoder layer
145pub struct Qwen2VLDecoderLayer {
146    self_attn: Qwen2VLAttention,
147    mlp: Qwen2VLMLP,
148    input_layernorm: Tensor,
149    post_attention_layernorm: Tensor,
150}
151
152/// Attention with RoPE
153pub 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
164/// MLP (SwiGLU)
165pub 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                // Encode images if provided
204                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                // Get text embeddings
212                let text_embeds = ops_fn::embedding(input_ids, &self.language_model.embed_tokens)?;
213
214                // Merge image and text embeddings (replace image tokens)
215                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                // Forward through language model
222                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                // Text-only forward
231                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        // Find image token positions and replace with image embeddings
335        // For now, return text embeddings (placeholder implementation)
336        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        // Patch embedding
356        let mut hidden_states = self.patch_embed.forward(pixel_values)?;
357
358        // Vision transformer blocks
359        for block in &self.blocks {
360            hidden_states = block.forward(&hidden_states)?;
361        }
362
363        // Merge vision tokens
364        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        // pixel_values: [batch, channels, frames, height, width]
405        // Convert to patches and project
406        let shape = pixel_values.shape();
407        let batch_size = shape[0];
408
409        // Flatten patches and project
410        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            // Reshape conv weight to linear
418            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; // Merge 2x2 patches
596        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        // Merge adjacent patches and project
608        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        // GQA expansion
835        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        // Causal mask
852        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}