Skip to main content

runtime/models_v2/
llava.rs

1//! LLaVA Model V2 - Clean implementation using solid abstractions
2//!
3//! This implements the LLaVA (Large Language and Vision Assistant) architecture including:
4//! - LLaVA-1.5-7B, LLaVA-1.5-13B, LLaVA-1.6-7B, LLaVA-1.6-13B, LLaVA-1.6-34B
5//!
6//! Architecture components:
7//! - Vision Tower: CLIP-style vision encoder with patch embeddings and bidirectional attention
8//! - Multimodal Projector: Projects vision features to language model dimension (Linear or MLP)
9//! - Language Model: LLaMA-style decoder with RoPE, GQA, and SwiGLU MLP
10
11use crate::model_config;
12use super::traits::*;
13use anyhow::Result;
14use serde::{Serialize, Deserialize};
15
16// ============================================================================
17// Configuration
18// ============================================================================
19
20model_config!(LLaVAConfig {
21    // Language model config (based on Llama/Vicuna)
22    vocab_size: usize = 32000,
23    hidden_size: usize = 4096,
24    intermediate_size: usize = 11008,
25    num_hidden_layers: usize = 32,
26    num_attention_heads: usize = 32,
27    num_key_value_heads: usize = 32,
28    hidden_act: String = "silu".to_string(),
29    max_position_embeddings: usize = 4096,
30    initializer_range: f32 = 0.02,
31    rms_norm_eps: f32 = 1e-5,
32    use_cache: bool = true,
33    pad_token_id: i64 = 0,
34    bos_token_id: i64 = 1,
35    eos_token_id: i64 = 2,
36    tie_word_embeddings: bool = false,
37    rope_theta: f32 = 10000.0,
38    attention_bias: bool = false,
39    attention_dropout: f32 = 0.0,
40
41    // Vision config
42    vision_hidden_size: usize = 1024,
43    vision_intermediate_size: usize = 4096,
44    vision_num_hidden_layers: usize = 24,
45    vision_num_attention_heads: usize = 16,
46    vision_num_channels: usize = 3,
47    vision_patch_size: usize = 14,
48    vision_image_size: usize = 336,
49    vision_layer_norm_eps: f32 = 1e-5,
50
51    // Multimodal config
52    mm_projector_type: String = "mlp2x_gelu".to_string(),
53    mm_hidden_size: usize = 1024,
54    mm_vision_select_layer: i32 = -2,
55    mm_vision_select_feature: String = "patch".to_string(),
56    image_token_len: usize = 576,
57    im_patch_token: i64 = 32000,
58    im_start_token: i64 = 32001,
59    im_end_token: i64 = 32002,
60});
61
62impl LLaVAConfig {
63    /// Create LLaVAConfig from GGUF model configuration
64    pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
65        Self {
66            vocab_size: gguf.vocab_size,
67            hidden_size: gguf.hidden_size,
68            intermediate_size: gguf.intermediate_size,
69            num_hidden_layers: gguf.num_hidden_layers,
70            num_attention_heads: gguf.num_attention_heads,
71            num_key_value_heads: gguf.num_key_value_heads,
72            rms_norm_eps: gguf.rms_norm_eps,
73            rope_theta: gguf.rope_theta,
74            max_position_embeddings: gguf.max_position_embeddings,
75            ..Default::default()
76        }
77    }
78
79    /// Get the number of vision patches
80    pub fn num_patches(&self) -> usize {
81        (self.vision_image_size / self.vision_patch_size).pow(2)
82    }
83
84    /// Get the number of vision positions (patches + CLS token)
85    pub fn num_vision_positions(&self) -> usize {
86        self.num_patches() + 1
87    }
88}
89
90// ============================================================================
91// Vision Tower (CLIP-style)
92// ============================================================================
93
94/// CLIP-style Vision Embeddings
95pub struct LLaVAVisionEmbeddings {
96    patch_embedding_weight: Tensor,  // Conv2d kernel: [hidden, channels, patch, patch]
97    patch_embedding_bias: Tensor,    // Conv2d bias: [hidden]
98    class_embedding: Tensor,         // [hidden]
99    position_embedding: Tensor,      // [num_positions, hidden]
100    config: LLaVAConfig,
101}
102
103impl LLaVAVisionEmbeddings {
104    fn new(config: &LLaVAConfig, device: &Device) -> Result<Self> {
105        let num_positions = config.num_vision_positions();
106        let hidden = config.vision_hidden_size;
107        let patch = config.vision_patch_size;
108        let channels = config.vision_num_channels;
109
110        Ok(Self {
111            patch_embedding_weight: ops_fn::zeros(
112                &[hidden, channels, patch, patch],
113                DataType::Float32,
114                device,
115            )?,
116            patch_embedding_bias: ops_fn::zeros(&[hidden], DataType::Float32, device)?,
117            class_embedding: ops_fn::zeros(&[hidden], DataType::Float32, device)?,
118            position_embedding: ops_fn::zeros(&[num_positions, hidden], DataType::Float32, device)?,
119            config: config.clone(),
120        })
121    }
122
123    fn forward(&self, pixel_values: &Tensor) -> Result<Tensor> {
124        let shape = pixel_values.shape();
125        let batch_size = shape[0];
126        let hidden_size = self.config.vision_hidden_size;
127        let num_patches = self.config.num_patches();
128        let seq_len = num_patches + 1; // patches + CLS token
129
130        // In a real implementation, we would:
131        // 1. Apply Conv2d to extract patch embeddings
132        // 2. Flatten patches to [batch, num_patches, hidden]
133        // 3. Prepend class embedding
134        // 4. Add position embeddings
135
136        // For now, simulate the output shape
137        // Real implementation would use conv2d operation
138        let pixel_candle = pixel_values.to_candle()?;
139        let device = pixel_candle.device();
140
141        // Create patch embeddings by simulating conv2d output
142        let patch_embeds = candle_core::Tensor::zeros(
143            &[batch_size, num_patches, hidden_size],
144            candle_core::DType::F32,
145            device,
146        )?;
147
148        // Create class embedding for each batch
149        let class_emb = self.class_embedding.to_candle()?;
150        let class_emb = class_emb.unsqueeze(0)?.unsqueeze(0)?; // [1, 1, hidden]
151        let class_emb = class_emb.broadcast_as(&[batch_size, 1, hidden_size])?;
152
153        // Concatenate: [CLS, patches] -> [batch, seq_len, hidden]
154        let embeddings = candle_core::Tensor::cat(&[&class_emb, &patch_embeds], 1)?;
155
156        // Add position embeddings
157        let pos_emb = self.position_embedding.to_candle()?;
158        let pos_emb = pos_emb.unsqueeze(0)?; // [1, seq_len, hidden]
159        let embeddings = embeddings.broadcast_add(&pos_emb)?;
160
161        Ok(Tensor::from_candle(embeddings))
162    }
163
164    fn load_weights(&mut self, weights: &ModelWeights, prefix: &str) -> Result<()> {
165        if let Some(w) = weights.get(&format!("{}.patch_embedding.weight", prefix)) {
166            self.patch_embedding_weight = w.clone();
167        }
168        if let Some(w) = weights.get(&format!("{}.patch_embedding.bias", prefix)) {
169            self.patch_embedding_bias = w.clone();
170        }
171        if let Some(w) = weights.get(&format!("{}.class_embedding", prefix)) {
172            self.class_embedding = w.clone();
173        }
174        if let Some(w) = weights.get(&format!("{}.position_embedding.weight", prefix)) {
175            self.position_embedding = w.clone();
176        }
177        Ok(())
178    }
179
180    fn to_device(&mut self, device: &Device) -> Result<()> {
181        self.patch_embedding_weight = self.patch_embedding_weight.to_device(device)?;
182        self.patch_embedding_bias = self.patch_embedding_bias.to_device(device)?;
183        self.class_embedding = self.class_embedding.to_device(device)?;
184        self.position_embedding = self.position_embedding.to_device(device)?;
185        Ok(())
186    }
187}
188
189/// CLIP Vision Attention (bidirectional, no causal mask)
190pub struct LLaVAVisionAttention {
191    q_proj: Tensor,
192    k_proj: Tensor,
193    v_proj: Tensor,
194    out_proj: Tensor,
195    num_heads: usize,
196    head_dim: usize,
197    scale: f32,
198}
199
200impl LLaVAVisionAttention {
201    fn new(config: &LLaVAConfig, device: &Device) -> Result<Self> {
202        let hidden = config.vision_hidden_size;
203        let num_heads = config.vision_num_attention_heads;
204        let head_dim = hidden / num_heads;
205        let scale = 1.0 / (head_dim as f32).sqrt();
206
207        Ok(Self {
208            q_proj: ops_fn::zeros(&[hidden, hidden], DataType::Float32, device)?,
209            k_proj: ops_fn::zeros(&[hidden, hidden], DataType::Float32, device)?,
210            v_proj: ops_fn::zeros(&[hidden, hidden], DataType::Float32, device)?,
211            out_proj: ops_fn::zeros(&[hidden, hidden], DataType::Float32, device)?,
212            num_heads,
213            head_dim,
214            scale,
215        })
216    }
217
218    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
219        let shape = hidden_states.shape();
220        let (batch_size, seq_len, _) = (shape[0], shape[1], shape[2]);
221
222        // Project to Q, K, V
223        let query = ops_fn::matmul(hidden_states, &self.q_proj)?;
224        let key = ops_fn::matmul(hidden_states, &self.k_proj)?;
225        let value = ops_fn::matmul(hidden_states, &self.v_proj)?;
226
227        // Reshape for multi-head attention
228        let q_candle = query.to_candle()?;
229        let k_candle = key.to_candle()?;
230        let v_candle = value.to_candle()?;
231
232        // [batch, seq, hidden] -> [batch, seq, heads, head_dim] -> [batch, heads, seq, head_dim]
233        let q = q_candle
234            .reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?
235            .transpose(1, 2)?;
236        let k = k_candle
237            .reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?
238            .transpose(1, 2)?;
239        let v = v_candle
240            .reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?
241            .transpose(1, 2)?;
242
243        // Attention scores: Q @ K^T
244        let k_t = k.transpose(2, 3)?;
245        let scores = q.contiguous()?.matmul(&k_t.contiguous()?)?;
246        let scaled_scores = (scores * (self.scale as f64))?;
247
248        // Softmax (bidirectional - no causal mask for vision)
249        let attention_weights = candle_nn::ops::softmax_last_dim(&scaled_scores)?;
250
251        // Apply attention to values
252        let attn_output = attention_weights.matmul(&v.contiguous()?)?;
253
254        // Reshape back: [batch, heads, seq, head_dim] -> [batch, seq, hidden]
255        let attn_output = attn_output
256            .transpose(1, 2)?
257            .reshape(&[batch_size, seq_len, self.num_heads * self.head_dim])?;
258
259        let attn_output = Tensor::from_candle(attn_output);
260
261        // Output projection
262        ops_fn::matmul(&attn_output, &self.out_proj)
263    }
264
265    fn load_weights(&mut self, weights: &ModelWeights, prefix: &str) -> Result<()> {
266        // Load and transpose for matmul: [out, in] -> [in, out]
267        if let Some(w) = weights.get(&format!("{}.q_proj.weight", prefix)) {
268            self.q_proj = ops_fn::transpose(w)?;
269        }
270        if let Some(w) = weights.get(&format!("{}.k_proj.weight", prefix)) {
271            self.k_proj = ops_fn::transpose(w)?;
272        }
273        if let Some(w) = weights.get(&format!("{}.v_proj.weight", prefix)) {
274            self.v_proj = ops_fn::transpose(w)?;
275        }
276        if let Some(w) = weights.get(&format!("{}.out_proj.weight", prefix)) {
277            self.out_proj = ops_fn::transpose(w)?;
278        }
279        Ok(())
280    }
281
282    fn to_device(&mut self, device: &Device) -> Result<()> {
283        self.q_proj = self.q_proj.to_device(device)?;
284        self.k_proj = self.k_proj.to_device(device)?;
285        self.v_proj = self.v_proj.to_device(device)?;
286        self.out_proj = self.out_proj.to_device(device)?;
287        Ok(())
288    }
289}
290
291/// CLIP Vision MLP (GELU activation)
292pub struct LLaVAVisionMLP {
293    fc1: Tensor,
294    fc2: Tensor,
295}
296
297impl LLaVAVisionMLP {
298    fn new(config: &LLaVAConfig, device: &Device) -> Result<Self> {
299        let hidden = config.vision_hidden_size;
300        let intermediate = config.vision_intermediate_size;
301
302        Ok(Self {
303            fc1: ops_fn::zeros(&[hidden, intermediate], DataType::Float32, device)?,
304            fc2: ops_fn::zeros(&[intermediate, hidden], DataType::Float32, device)?,
305        })
306    }
307
308    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
309        let hidden = ops_fn::matmul(hidden_states, &self.fc1)?;
310        let hidden = ops_fn::gelu(&hidden)?;
311        ops_fn::matmul(&hidden, &self.fc2)
312    }
313
314    fn load_weights(&mut self, weights: &ModelWeights, prefix: &str) -> Result<()> {
315        if let Some(w) = weights.get(&format!("{}.fc1.weight", prefix)) {
316            self.fc1 = ops_fn::transpose(w)?;
317        }
318        if let Some(w) = weights.get(&format!("{}.fc2.weight", prefix)) {
319            self.fc2 = ops_fn::transpose(w)?;
320        }
321        Ok(())
322    }
323
324    fn to_device(&mut self, device: &Device) -> Result<()> {
325        self.fc1 = self.fc1.to_device(device)?;
326        self.fc2 = self.fc2.to_device(device)?;
327        Ok(())
328    }
329}
330
331/// CLIP Vision Encoder Layer
332pub struct LLaVAVisionLayer {
333    self_attn: LLaVAVisionAttention,
334    layer_norm1: Tensor,
335    mlp: LLaVAVisionMLP,
336    layer_norm2: Tensor,
337    eps: f32,
338}
339
340impl LLaVAVisionLayer {
341    fn new(config: &LLaVAConfig, device: &Device) -> Result<Self> {
342        let hidden = config.vision_hidden_size;
343
344        Ok(Self {
345            self_attn: LLaVAVisionAttention::new(config, device)?,
346            layer_norm1: ops_fn::zeros(&[hidden], DataType::Float32, device)?,
347            mlp: LLaVAVisionMLP::new(config, device)?,
348            layer_norm2: ops_fn::zeros(&[hidden], DataType::Float32, device)?,
349            eps: config.vision_layer_norm_eps,
350        })
351    }
352
353    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
354        // Pre-norm attention with residual
355        let residual = hidden_states.clone();
356        let hidden_states = ops_fn::layer_norm(hidden_states, &self.layer_norm1, None, self.eps)?;
357        let hidden_states = self.self_attn.forward(&hidden_states)?;
358        let hidden_states = ops_fn::add(&residual, &hidden_states)?;
359
360        // Pre-norm MLP with residual
361        let residual = hidden_states.clone();
362        let hidden_states = ops_fn::layer_norm(&hidden_states, &self.layer_norm2, None, self.eps)?;
363        let hidden_states = self.mlp.forward(&hidden_states)?;
364        ops_fn::add(&residual, &hidden_states)
365    }
366
367    fn load_weights(&mut self, weights: &ModelWeights, prefix: &str) -> Result<()> {
368        if let Some(w) = weights.get(&format!("{}.layer_norm1.weight", prefix)) {
369            self.layer_norm1 = w.clone();
370        }
371        if let Some(w) = weights.get(&format!("{}.layer_norm2.weight", prefix)) {
372            self.layer_norm2 = w.clone();
373        }
374        self.self_attn.load_weights(weights, &format!("{}.self_attn", prefix))?;
375        self.mlp.load_weights(weights, &format!("{}.mlp", prefix))?;
376        Ok(())
377    }
378
379    fn to_device(&mut self, device: &Device) -> Result<()> {
380        self.layer_norm1 = self.layer_norm1.to_device(device)?;
381        self.layer_norm2 = self.layer_norm2.to_device(device)?;
382        self.self_attn.to_device(device)?;
383        self.mlp.to_device(device)?;
384        Ok(())
385    }
386}
387
388/// CLIP Vision Encoder
389pub struct LLaVAVisionEncoder {
390    layers: Vec<LLaVAVisionLayer>,
391}
392
393impl LLaVAVisionEncoder {
394    fn new(config: &LLaVAConfig, device: &Device) -> Result<Self> {
395        let mut layers = Vec::with_capacity(config.vision_num_hidden_layers);
396        for _ in 0..config.vision_num_hidden_layers {
397            layers.push(LLaVAVisionLayer::new(config, device)?);
398        }
399        Ok(Self { layers })
400    }
401
402    fn forward(&self, hidden_states: &Tensor, select_layer: i32) -> Result<Tensor> {
403        let num_layers = self.layers.len() as i32;
404        let target_layer = if select_layer < 0 {
405            (num_layers + select_layer) as usize
406        } else {
407            select_layer as usize
408        };
409
410        let mut hidden_states = hidden_states.clone();
411        for (i, layer) in self.layers.iter().enumerate() {
412            hidden_states = layer.forward(&hidden_states)?;
413            // Return early if we reached the target layer
414            if i == target_layer {
415                return Ok(hidden_states);
416            }
417        }
418        Ok(hidden_states)
419    }
420
421    fn load_weights(&mut self, weights: &ModelWeights, prefix: &str) -> Result<()> {
422        for (i, layer) in self.layers.iter_mut().enumerate() {
423            layer.load_weights(weights, &format!("{}.layers.{}", prefix, i))?;
424        }
425        Ok(())
426    }
427
428    fn to_device(&mut self, device: &Device) -> Result<()> {
429        for layer in &mut self.layers {
430            layer.to_device(device)?;
431        }
432        Ok(())
433    }
434}
435
436/// Complete Vision Tower (CLIP ViT)
437pub struct LLaVAVisionTower {
438    embeddings: LLaVAVisionEmbeddings,
439    encoder: LLaVAVisionEncoder,
440    post_layernorm: Tensor,
441    config: LLaVAConfig,
442}
443
444impl LLaVAVisionTower {
445    fn new(config: &LLaVAConfig, device: &Device) -> Result<Self> {
446        Ok(Self {
447            embeddings: LLaVAVisionEmbeddings::new(config, device)?,
448            encoder: LLaVAVisionEncoder::new(config, device)?,
449            post_layernorm: ops_fn::zeros(
450                &[config.vision_hidden_size],
451                DataType::Float32,
452                device,
453            )?,
454            config: config.clone(),
455        })
456    }
457
458    fn forward(&self, pixel_values: &Tensor) -> Result<Tensor> {
459        // Embed patches
460        let hidden_states = self.embeddings.forward(pixel_values)?;
461
462        // Encode through transformer layers (select layer for multimodal)
463        let hidden_states = self.encoder.forward(
464            &hidden_states,
465            self.config.mm_vision_select_layer,
466        )?;
467
468        // Post layer norm
469        let hidden_states = ops_fn::layer_norm(
470            &hidden_states,
471            &self.post_layernorm,
472            None,
473            self.config.vision_layer_norm_eps,
474        )?;
475
476        // Select features based on config
477        if self.config.mm_vision_select_feature == "patch" {
478            // Remove CLS token, keep only patch features
479            let candle_tensor = hidden_states.to_candle()?;
480            let shape = candle_tensor.dims();
481            let patch_features = candle_tensor.narrow(1, 1, shape[1] - 1)?;
482            Ok(Tensor::from_candle(patch_features))
483        } else {
484            // Return all features including CLS
485            Ok(hidden_states)
486        }
487    }
488
489    fn load_weights(&mut self, weights: &ModelWeights) -> Result<()> {
490        let prefix = "vision_tower.vision_model";
491        self.embeddings.load_weights(weights, &format!("{}.embeddings", prefix))?;
492        self.encoder.load_weights(weights, &format!("{}.encoder", prefix))?;
493        if let Some(w) = weights.get(&format!("{}.post_layernorm.weight", prefix)) {
494            self.post_layernorm = w.clone();
495        }
496        Ok(())
497    }
498
499    fn to_device(&mut self, device: &Device) -> Result<()> {
500        self.embeddings.to_device(device)?;
501        self.encoder.to_device(device)?;
502        self.post_layernorm = self.post_layernorm.to_device(device)?;
503        Ok(())
504    }
505}
506
507// ============================================================================
508// Multimodal Projector
509// ============================================================================
510
511/// Projects vision features to language model dimension
512pub struct LLaVAMultiModalProjector {
513    projector_type: String,
514    // For "linear" type
515    linear: Option<Tensor>,
516    linear_bias: Option<Tensor>,
517    // For "mlp2x_gelu" type (2-layer MLP with GELU)
518    mlp_fc1: Option<Tensor>,
519    mlp_fc1_bias: Option<Tensor>,
520    mlp_fc2: Option<Tensor>,
521    mlp_fc2_bias: Option<Tensor>,
522}
523
524impl LLaVAMultiModalProjector {
525    fn new(config: &LLaVAConfig, device: &Device) -> Result<Self> {
526        let vision_hidden = config.vision_hidden_size;
527        let text_hidden = config.hidden_size;
528
529        match config.mm_projector_type.as_str() {
530            "linear" => Ok(Self {
531                projector_type: "linear".to_string(),
532                linear: Some(ops_fn::zeros(
533                    &[vision_hidden, text_hidden],
534                    DataType::Float32,
535                    device,
536                )?),
537                linear_bias: Some(ops_fn::zeros(&[text_hidden], DataType::Float32, device)?),
538                mlp_fc1: None,
539                mlp_fc1_bias: None,
540                mlp_fc2: None,
541                mlp_fc2_bias: None,
542            }),
543            "mlp2x_gelu" => Ok(Self {
544                projector_type: "mlp2x_gelu".to_string(),
545                linear: None,
546                linear_bias: None,
547                mlp_fc1: Some(ops_fn::zeros(
548                    &[vision_hidden, text_hidden],
549                    DataType::Float32,
550                    device,
551                )?),
552                mlp_fc1_bias: Some(ops_fn::zeros(&[text_hidden], DataType::Float32, device)?),
553                mlp_fc2: Some(ops_fn::zeros(
554                    &[text_hidden, text_hidden],
555                    DataType::Float32,
556                    device,
557                )?),
558                mlp_fc2_bias: Some(ops_fn::zeros(&[text_hidden], DataType::Float32, device)?),
559            }),
560            _ => Err(anyhow::anyhow!(
561                "Unsupported projector type: {}",
562                config.mm_projector_type
563            )),
564        }
565    }
566
567    fn forward(&self, vision_features: &Tensor) -> Result<Tensor> {
568        match self.projector_type.as_str() {
569            "linear" => {
570                let linear = self.linear.as_ref().ok_or_else(|| {
571                    anyhow::anyhow!("Linear projector not initialized")
572                })?;
573                let output = ops_fn::matmul(vision_features, linear)?;
574                if let Some(ref bias) = self.linear_bias {
575                    ops_fn::add(&output, bias)
576                } else {
577                    Ok(output)
578                }
579            }
580            "mlp2x_gelu" => {
581                let fc1 = self.mlp_fc1.as_ref().ok_or_else(|| {
582                    anyhow::anyhow!("MLP fc1 not initialized")
583                })?;
584                let fc2 = self.mlp_fc2.as_ref().ok_or_else(|| {
585                    anyhow::anyhow!("MLP fc2 not initialized")
586                })?;
587
588                // First layer with GELU
589                let mut hidden = ops_fn::matmul(vision_features, fc1)?;
590                if let Some(ref bias) = self.mlp_fc1_bias {
591                    hidden = ops_fn::add(&hidden, bias)?;
592                }
593                hidden = ops_fn::gelu(&hidden)?;
594
595                // Second layer
596                let mut output = ops_fn::matmul(&hidden, fc2)?;
597                if let Some(ref bias) = self.mlp_fc2_bias {
598                    output = ops_fn::add(&output, bias)?;
599                }
600                Ok(output)
601            }
602            _ => Err(anyhow::anyhow!(
603                "Unsupported projector type: {}",
604                self.projector_type
605            )),
606        }
607    }
608
609    fn load_weights(&mut self, weights: &ModelWeights) -> Result<()> {
610        match self.projector_type.as_str() {
611            "linear" => {
612                if let Some(w) = weights.get("mm_projector.weight") {
613                    self.linear = Some(ops_fn::transpose(w)?);
614                }
615                if let Some(w) = weights.get("mm_projector.bias") {
616                    self.linear_bias = Some(w.clone());
617                }
618            }
619            "mlp2x_gelu" => {
620                // MLP projector has 2 layers: 0 and 2 (with GELU at index 1)
621                if let Some(w) = weights.get("mm_projector.0.weight") {
622                    self.mlp_fc1 = Some(ops_fn::transpose(w)?);
623                }
624                if let Some(w) = weights.get("mm_projector.0.bias") {
625                    self.mlp_fc1_bias = Some(w.clone());
626                }
627                if let Some(w) = weights.get("mm_projector.2.weight") {
628                    self.mlp_fc2 = Some(ops_fn::transpose(w)?);
629                }
630                if let Some(w) = weights.get("mm_projector.2.bias") {
631                    self.mlp_fc2_bias = Some(w.clone());
632                }
633            }
634            _ => {}
635        }
636        Ok(())
637    }
638
639    fn to_device(&mut self, device: &Device) -> Result<()> {
640        if let Some(ref mut w) = self.linear {
641            *w = w.to_device(device)?;
642        }
643        if let Some(ref mut w) = self.linear_bias {
644            *w = w.to_device(device)?;
645        }
646        if let Some(ref mut w) = self.mlp_fc1 {
647            *w = w.to_device(device)?;
648        }
649        if let Some(ref mut w) = self.mlp_fc1_bias {
650            *w = w.to_device(device)?;
651        }
652        if let Some(ref mut w) = self.mlp_fc2 {
653            *w = w.to_device(device)?;
654        }
655        if let Some(ref mut w) = self.mlp_fc2_bias {
656            *w = w.to_device(device)?;
657        }
658        Ok(())
659    }
660}
661
662// ============================================================================
663// Language Model (LLaMA-style)
664// ============================================================================
665
666/// LLaMA-style Attention with RoPE and GQA
667pub struct LLaVALanguageAttention {
668    q_proj: Tensor,
669    k_proj: Tensor,
670    v_proj: Tensor,
671    o_proj: Tensor,
672    num_heads: usize,
673    num_key_value_heads: usize,
674    head_dim: usize,
675    scale: f32,
676}
677
678impl LLaVALanguageAttention {
679    fn new(config: &LLaVAConfig, device: &Device) -> Result<Self> {
680        let num_heads = config.num_attention_heads;
681        let num_kv_heads = config.num_key_value_heads;
682        let head_dim = config.hidden_size / num_heads;
683        let scale = 1.0 / (head_dim as f32).sqrt();
684
685        Ok(Self {
686            q_proj: ops_fn::zeros(
687                &[config.hidden_size, num_heads * head_dim],
688                DataType::Float32,
689                device,
690            )?,
691            k_proj: ops_fn::zeros(
692                &[config.hidden_size, num_kv_heads * head_dim],
693                DataType::Float32,
694                device,
695            )?,
696            v_proj: ops_fn::zeros(
697                &[config.hidden_size, num_kv_heads * head_dim],
698                DataType::Float32,
699                device,
700            )?,
701            o_proj: ops_fn::zeros(
702                &[num_heads * head_dim, config.hidden_size],
703                DataType::Float32,
704                device,
705            )?,
706            num_heads,
707            num_key_value_heads: num_kv_heads,
708            head_dim,
709            scale,
710        })
711    }
712
713    fn forward(&self, hidden_states: &Tensor, rope_theta: f32) -> Result<Tensor> {
714        let shape = hidden_states.shape();
715        let (batch_size, seq_len, _) = if shape.len() == 3 {
716            (shape[0], shape[1], shape[2])
717        } else if shape.len() == 2 {
718            (1, shape[0], shape[1])
719        } else {
720            return Err(anyhow::anyhow!("Invalid hidden_states shape: {:?}", shape));
721        };
722
723        // Project to Q, K, V
724        let query_states = ops_fn::matmul(hidden_states, &self.q_proj)?;
725        let key_states = ops_fn::matmul(hidden_states, &self.k_proj)?;
726        let value_states = ops_fn::matmul(hidden_states, &self.v_proj)?;
727
728        // Reshape for multi-head attention
729        let q_candle = query_states.to_candle()?;
730        let k_candle = key_states.to_candle()?;
731        let v_candle = value_states.to_candle()?;
732
733        let q_reshaped = q_candle
734            .reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?
735            .transpose(1, 2)?;
736        let k_reshaped = k_candle
737            .reshape(&[batch_size, seq_len, self.num_key_value_heads, self.head_dim])?
738            .transpose(1, 2)?;
739        let v_reshaped = v_candle
740            .reshape(&[batch_size, seq_len, self.num_key_value_heads, self.head_dim])?
741            .transpose(1, 2)?;
742
743        // Apply RoPE to Q and K
744        let (q_with_rope, k_with_rope) = apply_rope(
745            &q_reshaped,
746            &k_reshaped,
747            seq_len,
748            self.head_dim,
749            rope_theta,
750        )?;
751
752        // Handle GQA - repeat K/V heads to match Q heads
753        let num_groups = self.num_heads / self.num_key_value_heads;
754        let (k_expanded, v_expanded) = if num_groups > 1 {
755            let k_rep = k_with_rope
756                .unsqueeze(2)?
757                .broadcast_as(&[
758                    batch_size,
759                    self.num_key_value_heads,
760                    num_groups,
761                    seq_len,
762                    self.head_dim,
763                ])?
764                .reshape(&[batch_size, self.num_heads, seq_len, self.head_dim])?;
765            let v_rep = v_reshaped
766                .unsqueeze(2)?
767                .broadcast_as(&[
768                    batch_size,
769                    self.num_key_value_heads,
770                    num_groups,
771                    seq_len,
772                    self.head_dim,
773                ])?
774                .reshape(&[batch_size, self.num_heads, seq_len, self.head_dim])?;
775            (k_rep, v_rep)
776        } else {
777            (k_with_rope, v_reshaped)
778        };
779
780        // Attention scores with causal mask
781        let k_t = k_expanded.transpose(2, 3)?;
782        let scores = q_with_rope.contiguous()?.matmul(&k_t.contiguous()?)?;
783        let scaled_scores = (scores * (self.scale as f64))?;
784
785        // Apply causal mask
786        let device = scaled_scores.device();
787        let causal_mask = {
788            let mut mask_data = vec![0.0f32; seq_len * seq_len];
789            for i in 0..seq_len {
790                for j in 0..seq_len {
791                    if j > i {
792                        mask_data[i * seq_len + j] = f32::NEG_INFINITY;
793                    }
794                }
795            }
796            candle_core::Tensor::from_vec(mask_data, &[1, 1, seq_len, seq_len], device)?
797        };
798        let masked_scores = scaled_scores.broadcast_add(&causal_mask)?;
799
800        // Softmax and apply to values
801        let attention_weights = candle_nn::ops::softmax_last_dim(&masked_scores)?;
802        let attn_output = attention_weights.matmul(&v_expanded.contiguous()?)?;
803
804        // Reshape back
805        let attn_output = attn_output
806            .transpose(1, 2)?
807            .reshape(&[batch_size, seq_len, self.num_heads * self.head_dim])?;
808
809        let attn_output = Tensor::from_candle(attn_output);
810
811        // Output projection
812        ops_fn::matmul(&attn_output, &self.o_proj)
813    }
814
815    fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
816        let prefix = format!("language_model.model.layers.{}.self_attn", layer_idx);
817        if let Some(w) = weights.get(&format!("{}.q_proj.weight", prefix)) {
818            self.q_proj = ops_fn::transpose(w)?;
819        }
820        if let Some(w) = weights.get(&format!("{}.k_proj.weight", prefix)) {
821            self.k_proj = ops_fn::transpose(w)?;
822        }
823        if let Some(w) = weights.get(&format!("{}.v_proj.weight", prefix)) {
824            self.v_proj = ops_fn::transpose(w)?;
825        }
826        if let Some(w) = weights.get(&format!("{}.o_proj.weight", prefix)) {
827            self.o_proj = ops_fn::transpose(w)?;
828        }
829        Ok(())
830    }
831
832    fn to_device(&mut self, device: &Device) -> Result<()> {
833        self.q_proj = self.q_proj.to_device(device)?;
834        self.k_proj = self.k_proj.to_device(device)?;
835        self.v_proj = self.v_proj.to_device(device)?;
836        self.o_proj = self.o_proj.to_device(device)?;
837        Ok(())
838    }
839}
840
841/// LLaMA-style MLP with SwiGLU
842pub struct LLaVALanguageMLP {
843    gate_proj: Tensor,
844    up_proj: Tensor,
845    down_proj: Tensor,
846}
847
848impl LLaVALanguageMLP {
849    fn new(config: &LLaVAConfig, device: &Device) -> Result<Self> {
850        Ok(Self {
851            gate_proj: ops_fn::zeros(
852                &[config.hidden_size, config.intermediate_size],
853                DataType::Float32,
854                device,
855            )?,
856            up_proj: ops_fn::zeros(
857                &[config.hidden_size, config.intermediate_size],
858                DataType::Float32,
859                device,
860            )?,
861            down_proj: ops_fn::zeros(
862                &[config.intermediate_size, config.hidden_size],
863                DataType::Float32,
864                device,
865            )?,
866        })
867    }
868
869    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
870        // SwiGLU: down(silu(gate(x)) * up(x))
871        let gate_output = ops_fn::matmul(hidden_states, &self.gate_proj)?;
872        let up_output = ops_fn::matmul(hidden_states, &self.up_proj)?;
873        let gate_activated = ops_fn::silu(&gate_output)?;
874        let gated = ops_fn::mul(&gate_activated, &up_output)?;
875        ops_fn::matmul(&gated, &self.down_proj)
876    }
877
878    fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
879        let prefix = format!("language_model.model.layers.{}.mlp", layer_idx);
880        if let Some(w) = weights.get(&format!("{}.gate_proj.weight", prefix)) {
881            self.gate_proj = ops_fn::transpose(w)?;
882        }
883        if let Some(w) = weights.get(&format!("{}.up_proj.weight", prefix)) {
884            self.up_proj = ops_fn::transpose(w)?;
885        }
886        if let Some(w) = weights.get(&format!("{}.down_proj.weight", prefix)) {
887            self.down_proj = ops_fn::transpose(w)?;
888        }
889        Ok(())
890    }
891
892    fn to_device(&mut self, device: &Device) -> Result<()> {
893        self.gate_proj = self.gate_proj.to_device(device)?;
894        self.up_proj = self.up_proj.to_device(device)?;
895        self.down_proj = self.down_proj.to_device(device)?;
896        Ok(())
897    }
898}
899
900/// LLaMA-style Transformer Layer
901pub struct LLaVALanguageLayer {
902    self_attn: LLaVALanguageAttention,
903    mlp: LLaVALanguageMLP,
904    input_layernorm: Tensor,
905    post_attention_layernorm: Tensor,
906    rms_norm_eps: f32,
907}
908
909impl LLaVALanguageLayer {
910    fn new(config: &LLaVAConfig, device: &Device) -> Result<Self> {
911        Ok(Self {
912            self_attn: LLaVALanguageAttention::new(config, device)?,
913            mlp: LLaVALanguageMLP::new(config, device)?,
914            input_layernorm: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
915            post_attention_layernorm: ops_fn::zeros(
916                &[config.hidden_size],
917                DataType::Float32,
918                device,
919            )?,
920            rms_norm_eps: config.rms_norm_eps,
921        })
922    }
923
924    fn forward(&self, hidden_states: &Tensor, rope_theta: f32) -> Result<Tensor> {
925        // Pre-norm attention with residual
926        let normed = ops_fn::layer_norm(
927            hidden_states,
928            &self.input_layernorm,
929            None,
930            self.rms_norm_eps,
931        )?;
932        let attn_output = self.self_attn.forward(&normed, rope_theta)?;
933        let hidden_states = ops_fn::add(hidden_states, &attn_output)?;
934
935        // Pre-norm MLP with residual
936        let normed = ops_fn::layer_norm(
937            &hidden_states,
938            &self.post_attention_layernorm,
939            None,
940            self.rms_norm_eps,
941        )?;
942        let mlp_output = self.mlp.forward(&normed)?;
943        ops_fn::add(&hidden_states, &mlp_output)
944    }
945
946    fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
947        let prefix = format!("language_model.model.layers.{}", layer_idx);
948        if let Some(w) = weights.get(&format!("{}.input_layernorm.weight", prefix)) {
949            self.input_layernorm = w.clone();
950        }
951        if let Some(w) = weights.get(&format!("{}.post_attention_layernorm.weight", prefix)) {
952            self.post_attention_layernorm = w.clone();
953        }
954        self.self_attn.load_weights(weights, layer_idx)?;
955        self.mlp.load_weights(weights, layer_idx)?;
956        Ok(())
957    }
958
959    fn to_device(&mut self, device: &Device) -> Result<()> {
960        self.input_layernorm = self.input_layernorm.to_device(device)?;
961        self.post_attention_layernorm = self.post_attention_layernorm.to_device(device)?;
962        self.self_attn.to_device(device)?;
963        self.mlp.to_device(device)?;
964        Ok(())
965    }
966}
967
968/// Complete Language Model
969pub struct LLaVALanguageModel {
970    embed_tokens: Tensor,
971    layers: Vec<LLaVALanguageLayer>,
972    norm: Tensor,
973    lm_head: Tensor,
974    config: LLaVAConfig,
975}
976
977impl LLaVALanguageModel {
978    fn new(config: &LLaVAConfig, device: &Device) -> Result<Self> {
979        let embed_tokens = ops_fn::zeros(
980            &[config.vocab_size, config.hidden_size],
981            DataType::Float32,
982            device,
983        )?;
984        let norm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
985        let lm_head = if config.tie_word_embeddings {
986            embed_tokens.clone()
987        } else {
988            ops_fn::zeros(
989                &[config.hidden_size, config.vocab_size],
990                DataType::Float32,
991                device,
992            )?
993        };
994
995        let mut layers = Vec::with_capacity(config.num_hidden_layers);
996        for _ in 0..config.num_hidden_layers {
997            layers.push(LLaVALanguageLayer::new(config, device)?);
998        }
999
1000        Ok(Self {
1001            embed_tokens,
1002            layers,
1003            norm,
1004            lm_head,
1005            config: config.clone(),
1006        })
1007    }
1008
1009    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
1010        let mut hidden_states = hidden_states.clone();
1011
1012        // Apply transformer layers
1013        for layer in &self.layers {
1014            hidden_states = layer.forward(&hidden_states, self.config.rope_theta)?;
1015        }
1016
1017        // Final norm
1018        hidden_states = ops_fn::layer_norm(
1019            &hidden_states,
1020            &self.norm,
1021            None,
1022            self.config.rms_norm_eps,
1023        )?;
1024
1025        // LM head
1026        ops_fn::matmul(&hidden_states, &self.lm_head)
1027    }
1028
1029    fn forward_from_ids(&self, input_ids: &Tensor) -> Result<Tensor> {
1030        let hidden_states = ops_fn::embedding(input_ids, &self.embed_tokens)?;
1031        self.forward(&hidden_states)
1032    }
1033
1034    fn embed(&self, input_ids: &Tensor) -> Result<Tensor> {
1035        ops_fn::embedding(input_ids, &self.embed_tokens)
1036    }
1037
1038    fn load_weights(&mut self, weights: &ModelWeights) -> Result<()> {
1039        if let Some(w) = weights.get("language_model.model.embed_tokens.weight") {
1040            self.embed_tokens = w.clone();
1041        }
1042        if let Some(w) = weights.get("language_model.model.norm.weight") {
1043            self.norm = w.clone();
1044        }
1045        if let Some(w) = weights.get("language_model.lm_head.weight") {
1046            self.lm_head = ops_fn::transpose(w)?;
1047        }
1048        for (i, layer) in self.layers.iter_mut().enumerate() {
1049            layer.load_weights(weights, i)?;
1050        }
1051        Ok(())
1052    }
1053
1054    fn to_device(&mut self, device: &Device) -> Result<()> {
1055        self.embed_tokens = self.embed_tokens.to_device(device)?;
1056        self.norm = self.norm.to_device(device)?;
1057        self.lm_head = self.lm_head.to_device(device)?;
1058        for layer in &mut self.layers {
1059            layer.to_device(device)?;
1060        }
1061        Ok(())
1062    }
1063}
1064
1065// ============================================================================
1066// RoPE Implementation
1067// ============================================================================
1068
1069/// Apply Rotary Position Embedding (RoPE) to Q and K tensors
1070fn apply_rope(
1071    q: &candle_core::Tensor,
1072    k: &candle_core::Tensor,
1073    seq_len: usize,
1074    head_dim: usize,
1075    rope_theta: f32,
1076) -> Result<(candle_core::Tensor, candle_core::Tensor)> {
1077    let device = q.device();
1078    let half_dim = head_dim / 2;
1079
1080    // Compute inverse frequencies
1081    let inv_freq: Vec<f32> = (0..half_dim)
1082        .map(|i| 1.0 / rope_theta.powf((2 * i) as f32 / head_dim as f32))
1083        .collect();
1084
1085    // Create position indices
1086    let positions: Vec<f32> = (0..seq_len).map(|p| p as f32).collect();
1087
1088    // Compute angles: pos * inv_freq -> [seq_len, half_dim]
1089    let mut angles = Vec::with_capacity(seq_len * half_dim);
1090    for pos in &positions {
1091        for freq in &inv_freq {
1092            angles.push(pos * freq);
1093        }
1094    }
1095
1096    let angles_tensor = candle_core::Tensor::from_vec(angles, &[seq_len, half_dim], device)?;
1097    let cos = angles_tensor.cos()?;
1098    let sin = angles_tensor.sin()?;
1099
1100    // Reshape for broadcasting: [1, 1, seq_len, half_dim]
1101    let cos = cos.unsqueeze(0)?.unsqueeze(0)?;
1102    let sin = sin.unsqueeze(0)?.unsqueeze(0)?;
1103
1104    // Apply rotation
1105    let q_half1 = q.narrow(3, 0, half_dim)?;
1106    let q_half2 = q.narrow(3, half_dim, half_dim)?;
1107    let k_half1 = k.narrow(3, 0, half_dim)?;
1108    let k_half2 = k.narrow(3, half_dim, half_dim)?;
1109
1110    let q_rot1 = (q_half1.broadcast_mul(&cos)? - q_half2.broadcast_mul(&sin)?)?;
1111    let q_rot2 = (q_half1.broadcast_mul(&sin)? + q_half2.broadcast_mul(&cos)?)?;
1112    let k_rot1 = (k_half1.broadcast_mul(&cos)? - k_half2.broadcast_mul(&sin)?)?;
1113    let k_rot2 = (k_half1.broadcast_mul(&sin)? + k_half2.broadcast_mul(&cos)?)?;
1114
1115    let q_rotated = candle_core::Tensor::cat(&[&q_rot1, &q_rot2], 3)?;
1116    let k_rotated = candle_core::Tensor::cat(&[&k_rot1, &k_rot2], 3)?;
1117
1118    Ok((q_rotated, k_rotated))
1119}
1120
1121// ============================================================================
1122// Main LLaVA Model
1123// ============================================================================
1124
1125/// Complete LLaVA Model
1126pub struct LLaVAModelV2 {
1127    config: LLaVAConfig,
1128    device: Device,
1129    vision_tower: LLaVAVisionTower,
1130    mm_projector: LLaVAMultiModalProjector,
1131    language_model: LLaVALanguageModel,
1132}
1133
1134impl Model for LLaVAModelV2 {
1135    type Config = LLaVAConfig;
1136
1137    fn new(config: LLaVAConfig) -> Result<Self> {
1138        let device = Device::CPU;
1139        let vision_tower = LLaVAVisionTower::new(&config, &device)?;
1140        let mm_projector = LLaVAMultiModalProjector::new(&config, &device)?;
1141        let language_model = LLaVALanguageModel::new(&config, &device)?;
1142
1143        Ok(Self {
1144            config,
1145            device,
1146            vision_tower,
1147            mm_projector,
1148            language_model,
1149        })
1150    }
1151
1152    fn from_weights(config: LLaVAConfig, weights: ModelWeights) -> Result<Self> {
1153        let mut model = Self::new(config)?;
1154        model.vision_tower.load_weights(&weights)?;
1155        model.mm_projector.load_weights(&weights)?;
1156        model.language_model.load_weights(&weights)?;
1157        Ok(model)
1158    }
1159
1160    fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
1161        match inputs {
1162            ModelInputs::Multimodal {
1163                input_ids,
1164                pixel_values,
1165                ..
1166            } => {
1167                if let Some(pixel_values) = pixel_values {
1168                    // Process image through vision tower
1169                    let image_features = self.vision_tower.forward(pixel_values)?;
1170
1171                    // Project to language model dimension
1172                    let image_features = self.mm_projector.forward(&image_features)?;
1173
1174                    // Get text embeddings
1175                    let text_embeds = self.language_model.embed(input_ids)?;
1176
1177                    // Merge multimodal inputs
1178                    let merged_embeds = self.merge_multimodal_inputs(
1179                        input_ids,
1180                        &text_embeds,
1181                        &image_features,
1182                    )?;
1183
1184                    // Forward through language model
1185                    let logits = self.language_model.forward(&merged_embeds)?;
1186
1187                    Ok(ModelOutputs::Logits {
1188                        logits,
1189                        hidden_states: Some(merged_embeds),
1190                    })
1191                } else {
1192                    // No image, just process text
1193                    let logits = self.language_model.forward_from_ids(input_ids)?;
1194                    Ok(ModelOutputs::Logits {
1195                        logits,
1196                        hidden_states: None,
1197                    })
1198                }
1199            }
1200            ModelInputs::Text { input_ids, .. } => {
1201                let logits = self.language_model.forward_from_ids(input_ids)?;
1202                Ok(ModelOutputs::Logits {
1203                    logits,
1204                    hidden_states: None,
1205                })
1206            }
1207            ModelInputs::Image { pixel_values, .. } => {
1208                // Image-only: return image embeddings
1209                let image_features = self.vision_tower.forward(pixel_values)?;
1210                let image_features = self.mm_projector.forward(&image_features)?;
1211                Ok(ModelOutputs::Embeddings {
1212                    embeddings: image_features,
1213                    pooled: None,
1214                })
1215            }
1216            _ => Err(anyhow::anyhow!("LLaVA expects text, image, or multimodal input")),
1217        }
1218    }
1219
1220    fn generate(&self, prompt: &str, config: &GenerationConfig) -> Result<String> {
1221        use crate::tokenizer::Tokenizer;
1222        use rand::Rng;
1223
1224        // Tokenize prompt
1225        let tokenizer = Tokenizer::new();
1226        let mut tokens: Vec<u32> = tokenizer.encode(prompt);
1227
1228        // Generation loop
1229        for _ in 0..config.max_new_tokens {
1230            let tokens_i64: Vec<i64> = tokens.iter().map(|&t| t as i64).collect();
1231            let input_tensor = Tensor::from_i64_slice(&tokens_i64, &[1, tokens.len()], &self.device)?;
1232
1233            let inputs = ModelInputs::Text {
1234                input_ids: input_tensor,
1235                attention_mask: None,
1236                position_ids: None,
1237            };
1238
1239            let outputs = self.forward(&inputs)?;
1240            let logits = match outputs {
1241                ModelOutputs::Logits { logits, .. } => logits,
1242                _ => return Err(anyhow::anyhow!("Expected logits output")),
1243            };
1244
1245            // Get last token logits
1246            let logits_candle = logits.to_candle()?;
1247            let shape = logits_candle.dims();
1248            let last_logits = if shape.len() == 3 {
1249                let seq_len = shape[1];
1250                logits_candle
1251                    .narrow(1, seq_len - 1, 1)?
1252                    .squeeze(1)?
1253                    .squeeze(0)?
1254            } else {
1255                let seq_len = shape[0];
1256                logits_candle.narrow(0, seq_len - 1, 1)?.squeeze(0)?
1257            };
1258
1259            let logits_vec: Vec<f32> = last_logits.to_vec1()?;
1260
1261            let next_token = if config.do_sample && config.temperature > 0.0 {
1262                // Temperature sampling
1263                let scaled: Vec<f32> = logits_vec
1264                    .iter()
1265                    .map(|&x| x / config.temperature)
1266                    .collect();
1267                let max_val = scaled.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
1268                let exp_sum: f32 = scaled.iter().map(|&x| (x - max_val).exp()).sum();
1269                let probs: Vec<f32> = scaled
1270                    .iter()
1271                    .map(|&x| (x - max_val).exp() / exp_sum)
1272                    .collect();
1273
1274                let mut rng = rand::thread_rng();
1275                let random_val: f32 = rng.gen();
1276                let mut cumulative = 0.0;
1277                let mut sampled = 0u32;
1278                for (idx, &prob) in probs.iter().enumerate() {
1279                    cumulative += prob;
1280                    if random_val <= cumulative {
1281                        sampled = idx as u32;
1282                        break;
1283                    }
1284                }
1285                sampled
1286            } else {
1287                // Greedy sampling
1288                let mut max_idx = 0;
1289                let mut max_val = logits_vec[0];
1290                for (idx, &val) in logits_vec.iter().enumerate() {
1291                    if val > max_val {
1292                        max_val = val;
1293                        max_idx = idx;
1294                    }
1295                }
1296                max_idx as u32
1297            };
1298
1299            if next_token == config.eos_token_id {
1300                break;
1301            }
1302
1303            tokens.push(next_token);
1304        }
1305
1306        Ok(tokenizer.decode(&tokens))
1307    }
1308
1309    fn config(&self) -> &Self::Config {
1310        &self.config
1311    }
1312
1313    fn memory_requirements(&self) -> MemoryRequirements {
1314        let language_params = self.config.vocab_size * self.config.hidden_size
1315            + self.config.num_hidden_layers
1316                * (4 * self.config.hidden_size * self.config.hidden_size
1317                    + 3 * self.config.hidden_size * self.config.intermediate_size);
1318
1319        let vision_params = self.config.vision_num_hidden_layers
1320            * (4 * self.config.vision_hidden_size * self.config.vision_hidden_size
1321                + 2 * self.config.vision_hidden_size * self.config.vision_intermediate_size);
1322
1323        let projector_params = self.config.vision_hidden_size * self.config.hidden_size * 2;
1324
1325        let total_params = language_params + vision_params + projector_params;
1326        let param_bytes = total_params * 4; // float32
1327
1328        let kv_cache_bytes = 2
1329            * self.config.num_hidden_layers
1330            * self.config.max_position_embeddings
1331            * self.config.hidden_size
1332            * 4;
1333
1334        MemoryRequirements {
1335            gpu_memory: param_bytes,
1336            cpu_memory: param_bytes / 4,
1337            kv_cache_memory: kv_cache_bytes,
1338            peak_memory: param_bytes + kv_cache_bytes,
1339        }
1340    }
1341
1342    fn to_device(&mut self, device: &Device) -> Result<()> {
1343        self.vision_tower.to_device(device)?;
1344        self.mm_projector.to_device(device)?;
1345        self.language_model.to_device(device)?;
1346        self.device = device.clone();
1347        Ok(())
1348    }
1349}
1350
1351impl LLaVAModelV2 {
1352    /// Merge vision features into text sequence at image token positions
1353    ///
1354    /// This replaces <image> tokens in the text sequence with actual image features.
1355    /// The resulting sequence is: [text_before_image, image_features, text_after_image]
1356    fn merge_multimodal_inputs(
1357        &self,
1358        input_ids: &Tensor,
1359        text_embeds: &Tensor,
1360        image_features: &Tensor,
1361    ) -> Result<Tensor> {
1362        let input_candle = input_ids.to_candle()?;
1363        let text_candle = text_embeds.to_candle()?;
1364        let image_candle = image_features.to_candle()?;
1365
1366        let batch_size = text_candle.dims()[0];
1367        let text_seq_len = text_candle.dims()[1];
1368        let hidden_size = text_candle.dims()[2];
1369        let image_seq_len = image_candle.dims()[1];
1370
1371        // Find image token positions
1372        let input_flat = input_candle.flatten_all()?;
1373        let input_vec: Vec<i64> = input_flat.to_vec1()?;
1374
1375        let image_token_id = self.config.im_patch_token;
1376
1377        // Find first image token position (if any)
1378        let image_pos = input_vec.iter().position(|&id| id == image_token_id);
1379
1380        let merged = if let Some(pos) = image_pos {
1381            // Split text at image position
1382            let pos = pos % text_seq_len; // Handle batch dimension
1383
1384            if pos == 0 {
1385                // Image at start: [image, text]
1386                let text_after = text_candle.narrow(1, 1, text_seq_len - 1)?;
1387                candle_core::Tensor::cat(&[&image_candle, &text_after], 1)?
1388            } else if pos >= text_seq_len - 1 {
1389                // Image at end: [text, image]
1390                let text_before = text_candle.narrow(1, 0, text_seq_len - 1)?;
1391                candle_core::Tensor::cat(&[&text_before, &image_candle], 1)?
1392            } else {
1393                // Image in middle: [text_before, image, text_after]
1394                let text_before = text_candle.narrow(1, 0, pos)?;
1395                let text_after = text_candle.narrow(1, pos + 1, text_seq_len - pos - 1)?;
1396                candle_core::Tensor::cat(&[&text_before, &image_candle, &text_after], 1)?
1397            }
1398        } else {
1399            // No image token found, prepend image features
1400            candle_core::Tensor::cat(&[&image_candle, &text_candle], 1)?
1401        };
1402
1403        Ok(Tensor::from_candle(merged))
1404    }
1405
1406    /// Generate response for multimodal input (text + image)
1407    pub fn generate_multimodal(
1408        &self,
1409        prompt: &str,
1410        image: &Tensor,
1411        config: &GenerationConfig,
1412    ) -> Result<String> {
1413        use crate::tokenizer::Tokenizer;
1414        use rand::Rng;
1415
1416        let tokenizer = Tokenizer::new();
1417        let mut tokens: Vec<u32> = tokenizer.encode(prompt);
1418
1419        // Process image once
1420        let image_features = self.vision_tower.forward(image)?;
1421        let image_features = self.mm_projector.forward(&image_features)?;
1422
1423        for _ in 0..config.max_new_tokens {
1424            let tokens_i64: Vec<i64> = tokens.iter().map(|&t| t as i64).collect();
1425            let input_tensor = Tensor::from_i64_slice(&tokens_i64, &[1, tokens.len()], &self.device)?;
1426
1427            // Get text embeddings
1428            let text_embeds = self.language_model.embed(&input_tensor)?;
1429
1430            // Merge with image features
1431            let merged_embeds = self.merge_multimodal_inputs(
1432                &input_tensor,
1433                &text_embeds,
1434                &image_features,
1435            )?;
1436
1437            // Forward through language model
1438            let logits = self.language_model.forward(&merged_embeds)?;
1439
1440            // Sample next token
1441            let logits_candle = logits.to_candle()?;
1442            let shape = logits_candle.dims();
1443            let last_logits = if shape.len() == 3 {
1444                let seq_len = shape[1];
1445                logits_candle
1446                    .narrow(1, seq_len - 1, 1)?
1447                    .squeeze(1)?
1448                    .squeeze(0)?
1449            } else {
1450                let seq_len = shape[0];
1451                logits_candle.narrow(0, seq_len - 1, 1)?.squeeze(0)?
1452            };
1453
1454            let logits_vec: Vec<f32> = last_logits.to_vec1()?;
1455
1456            let next_token = if config.do_sample && config.temperature > 0.0 {
1457                let scaled: Vec<f32> = logits_vec
1458                    .iter()
1459                    .map(|&x| x / config.temperature)
1460                    .collect();
1461                let max_val = scaled.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
1462                let exp_sum: f32 = scaled.iter().map(|&x| (x - max_val).exp()).sum();
1463                let probs: Vec<f32> = scaled
1464                    .iter()
1465                    .map(|&x| (x - max_val).exp() / exp_sum)
1466                    .collect();
1467
1468                let mut rng = rand::thread_rng();
1469                let random_val: f32 = rng.gen();
1470                let mut cumulative = 0.0;
1471                let mut sampled = 0u32;
1472                for (idx, &prob) in probs.iter().enumerate() {
1473                    cumulative += prob;
1474                    if random_val <= cumulative {
1475                        sampled = idx as u32;
1476                        break;
1477                    }
1478                }
1479                sampled
1480            } else {
1481                let mut max_idx = 0;
1482                let mut max_val = logits_vec[0];
1483                for (idx, &val) in logits_vec.iter().enumerate() {
1484                    if val > max_val {
1485                        max_val = val;
1486                        max_idx = idx;
1487                    }
1488                }
1489                max_idx as u32
1490            };
1491
1492            if next_token == config.eos_token_id {
1493                break;
1494            }
1495
1496            tokens.push(next_token);
1497        }
1498
1499        Ok(tokenizer.decode(&tokens))
1500    }
1501}
1502
1503// ============================================================================
1504// Tests
1505// ============================================================================
1506
1507#[cfg(test)]
1508mod tests {
1509    use super::*;
1510
1511    #[test]
1512    fn test_llava_config_defaults() {
1513        let config = LLaVAConfig::default();
1514        assert_eq!(config.vocab_size, 32000);
1515        assert_eq!(config.hidden_size, 4096);
1516        assert_eq!(config.vision_hidden_size, 1024);
1517        assert_eq!(config.num_patches(), 576); // (336/14)^2
1518    }
1519
1520    #[test]
1521    fn test_llava_model_creation() {
1522        let config = LLaVAConfig {
1523            vocab_size: 1000,
1524            hidden_size: 128,
1525            intermediate_size: 512,
1526            num_hidden_layers: 2,
1527            num_attention_heads: 4,
1528            num_key_value_heads: 4,
1529            vision_hidden_size: 64,
1530            vision_intermediate_size: 256,
1531            vision_num_hidden_layers: 2,
1532            vision_num_attention_heads: 4,
1533            vision_patch_size: 14,
1534            vision_image_size: 224,
1535            ..Default::default()
1536        };
1537
1538        let model = LLaVAModelV2::new(config).unwrap();
1539        assert_eq!(model.config().vocab_size(), 1000);
1540        assert_eq!(model.config().hidden_size(), 128);
1541    }
1542
1543    #[test]
1544    fn test_llava_text_forward() {
1545        let config = LLaVAConfig {
1546            vocab_size: 100,
1547            hidden_size: 64,
1548            intermediate_size: 256,
1549            num_hidden_layers: 1,
1550            num_attention_heads: 4,
1551            num_key_value_heads: 4,
1552            vision_hidden_size: 32,
1553            vision_intermediate_size: 128,
1554            vision_num_hidden_layers: 1,
1555            vision_num_attention_heads: 4,
1556            vision_patch_size: 14,
1557            vision_image_size: 56,
1558            ..Default::default()
1559        };
1560
1561        let model = LLaVAModelV2::new(config).unwrap();
1562        let input_ids = ops_fn::zeros(&[1, 8], DataType::Int64, &Device::CPU).unwrap();
1563        let inputs = ModelInputs::text(input_ids);
1564
1565        let outputs = model.forward(&inputs).unwrap();
1566        match outputs {
1567            ModelOutputs::Logits { logits, .. } => {
1568                assert_eq!(logits.shape()[0], 1);
1569                assert_eq!(logits.shape()[1], 8);
1570                assert_eq!(logits.shape()[2], 100);
1571            }
1572            _ => panic!("Expected logits output"),
1573        }
1574    }
1575
1576    #[test]
1577    fn test_llava_multimodal_forward() {
1578        let config = LLaVAConfig {
1579            vocab_size: 100,
1580            hidden_size: 64,
1581            intermediate_size: 256,
1582            num_hidden_layers: 1,
1583            num_attention_heads: 4,
1584            num_key_value_heads: 4,
1585            vision_hidden_size: 32,
1586            vision_intermediate_size: 128,
1587            vision_num_hidden_layers: 1,
1588            vision_num_attention_heads: 4,
1589            vision_patch_size: 14,
1590            vision_image_size: 56,
1591            ..Default::default()
1592        };
1593
1594        let model = LLaVAModelV2::new(config.clone()).unwrap();
1595
1596        // Create multimodal inputs
1597        let input_ids = ops_fn::zeros(&[1, 8], DataType::Int64, &Device::CPU).unwrap();
1598        let pixel_values = ops_fn::zeros(
1599            &[1, 3, config.vision_image_size, config.vision_image_size],
1600            DataType::Float32,
1601            &Device::CPU,
1602        )
1603        .unwrap();
1604
1605        let inputs = ModelInputs::Multimodal {
1606            input_ids,
1607            pixel_values: Some(pixel_values),
1608            attention_mask: None,
1609            image_mask: None,
1610        };
1611
1612        let outputs = model.forward(&inputs).unwrap();
1613        match outputs {
1614            ModelOutputs::Logits { logits, .. } => {
1615                // Output should have combined sequence length
1616                assert_eq!(logits.shape()[0], 1);
1617                // Sequence length = text_len + image_patches - 1 (one text token replaced)
1618                // But our simple merge prepends image if no token found
1619                assert!(logits.shape()[1] > 0);
1620                assert_eq!(logits.shape()[2], 100);
1621            }
1622            _ => panic!("Expected logits output"),
1623        }
1624    }
1625
1626    #[test]
1627    fn test_llava_generation() {
1628        let config = LLaVAConfig {
1629            vocab_size: 256,
1630            hidden_size: 64,
1631            intermediate_size: 256,
1632            num_hidden_layers: 1,
1633            num_attention_heads: 4,
1634            num_key_value_heads: 4,
1635            vision_hidden_size: 32,
1636            vision_intermediate_size: 128,
1637            vision_num_hidden_layers: 1,
1638            vision_num_attention_heads: 4,
1639            vision_patch_size: 14,
1640            vision_image_size: 56,
1641            ..Default::default()
1642        };
1643
1644        let model = LLaVAModelV2::new(config).unwrap();
1645        let gen_config = GenerationConfig {
1646            max_new_tokens: 5,
1647            ..Default::default()
1648        };
1649
1650        let output = model.generate("Hello", &gen_config).unwrap();
1651        assert!(!output.is_empty());
1652    }
1653}