1use crate::model_config;
12use super::traits::*;
13use anyhow::Result;
14use serde::{Serialize, Deserialize};
15
16model_config!(LLaVAConfig {
21 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_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 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 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 pub fn num_patches(&self) -> usize {
81 (self.vision_image_size / self.vision_patch_size).pow(2)
82 }
83
84 pub fn num_vision_positions(&self) -> usize {
86 self.num_patches() + 1
87 }
88}
89
90pub struct LLaVAVisionEmbeddings {
96 patch_embedding_weight: Tensor, patch_embedding_bias: Tensor, class_embedding: Tensor, position_embedding: Tensor, 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; let pixel_candle = pixel_values.to_candle()?;
139 let device = pixel_candle.device();
140
141 let patch_embeds = candle_core::Tensor::zeros(
143 &[batch_size, num_patches, hidden_size],
144 candle_core::DType::F32,
145 device,
146 )?;
147
148 let class_emb = self.class_embedding.to_candle()?;
150 let class_emb = class_emb.unsqueeze(0)?.unsqueeze(0)?; let class_emb = class_emb.broadcast_as(&[batch_size, 1, hidden_size])?;
152
153 let embeddings = candle_core::Tensor::cat(&[&class_emb, &patch_embeds], 1)?;
155
156 let pos_emb = self.position_embedding.to_candle()?;
158 let pos_emb = pos_emb.unsqueeze(0)?; 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
189pub 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 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 let q_candle = query.to_candle()?;
229 let k_candle = key.to_candle()?;
230 let v_candle = value.to_candle()?;
231
232 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 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 let attention_weights = candle_nn::ops::softmax_last_dim(&scaled_scores)?;
250
251 let attn_output = attention_weights.matmul(&v.contiguous()?)?;
253
254 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 ops_fn::matmul(&attn_output, &self.out_proj)
263 }
264
265 fn load_weights(&mut self, weights: &ModelWeights, prefix: &str) -> Result<()> {
266 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
291pub 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
331pub 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 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 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
388pub 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 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
436pub 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 let hidden_states = self.embeddings.forward(pixel_values)?;
461
462 let hidden_states = self.encoder.forward(
464 &hidden_states,
465 self.config.mm_vision_select_layer,
466 )?;
467
468 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 if self.config.mm_vision_select_feature == "patch" {
478 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 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
507pub struct LLaVAMultiModalProjector {
513 projector_type: String,
514 linear: Option<Tensor>,
516 linear_bias: Option<Tensor>,
517 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 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 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 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
662pub 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 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 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 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 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 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 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 let attention_weights = candle_nn::ops::softmax_last_dim(&masked_scores)?;
802 let attn_output = attention_weights.matmul(&v_expanded.contiguous()?)?;
803
804 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 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
841pub 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 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
900pub 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 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 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
968pub 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 for layer in &self.layers {
1014 hidden_states = layer.forward(&hidden_states, self.config.rope_theta)?;
1015 }
1016
1017 hidden_states = ops_fn::layer_norm(
1019 &hidden_states,
1020 &self.norm,
1021 None,
1022 self.config.rms_norm_eps,
1023 )?;
1024
1025 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
1065fn 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 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 let positions: Vec<f32> = (0..seq_len).map(|p| p as f32).collect();
1087
1088 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 let cos = cos.unsqueeze(0)?.unsqueeze(0)?;
1102 let sin = sin.unsqueeze(0)?.unsqueeze(0)?;
1103
1104 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
1121pub 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 let image_features = self.vision_tower.forward(pixel_values)?;
1170
1171 let image_features = self.mm_projector.forward(&image_features)?;
1173
1174 let text_embeds = self.language_model.embed(input_ids)?;
1176
1177 let merged_embeds = self.merge_multimodal_inputs(
1179 input_ids,
1180 &text_embeds,
1181 &image_features,
1182 )?;
1183
1184 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 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 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 let tokenizer = Tokenizer::new();
1226 let mut tokens: Vec<u32> = tokenizer.encode(prompt);
1227
1228 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 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 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 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; 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 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 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 let image_pos = input_vec.iter().position(|&id| id == image_token_id);
1379
1380 let merged = if let Some(pos) = image_pos {
1381 let pos = pos % text_seq_len; if pos == 0 {
1385 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 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 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 candle_core::Tensor::cat(&[&image_candle, &text_candle], 1)?
1401 };
1402
1403 Ok(Tensor::from_candle(merged))
1404 }
1405
1406 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 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 let text_embeds = self.language_model.embed(&input_tensor)?;
1429
1430 let merged_embeds = self.merge_multimodal_inputs(
1432 &input_tensor,
1433 &text_embeds,
1434 &image_features,
1435 )?;
1436
1437 let logits = self.language_model.forward(&merged_embeds)?;
1439
1440 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#[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); }
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 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 assert_eq!(logits.shape()[0], 1);
1617 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}