Skip to main content

runtime/models_v2/
phi3_vision.rs

1//! Phi-3-Vision Model V2 - Vision-Language Model
2//!
3//! This implements the Phi-3-Vision architecture which features:
4//! - CLIP ViT vision encoder
5//! - MLP projector to align vision/text embeddings
6//! - Phi-3 decoder for language modeling
7//!
8//! Supports: Phi-3-Vision-128k-Instruct
9
10use crate::model_config;
11use super::traits::*;
12use anyhow::Result;
13use serde::{Serialize, Deserialize};
14
15model_config!(Phi3VisionConfig {
16    // Language model config
17    vocab_size: usize = 32064,
18    hidden_size: usize = 3072,
19    intermediate_size: usize = 8192,
20    num_hidden_layers: usize = 32,
21    num_attention_heads: usize = 32,
22    num_key_value_heads: usize = 32,
23    max_position_embeddings: usize = 131072,
24    rms_norm_eps: f32 = 1e-5,
25    rope_theta: f32 = 10000.0,
26
27    // Vision encoder config (CLIP ViT)
28    vision_hidden_size: usize = 1024,
29    vision_intermediate_size: usize = 4096,
30    vision_num_hidden_layers: usize = 24,
31    vision_num_attention_heads: usize = 16,
32    vision_patch_size: usize = 14,
33    vision_image_size: usize = 336,
34
35    pad_token_id: i64 = 32000,
36    bos_token_id: i64 = 1,
37    eos_token_id: i64 = 32000,
38    image_token_id: i64 = 32044,
39});
40
41impl Phi3VisionConfig {
42    pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
43        Self {
44            vocab_size: gguf.vocab_size,
45            hidden_size: gguf.hidden_size,
46            intermediate_size: gguf.intermediate_size,
47            num_hidden_layers: gguf.num_hidden_layers,
48            num_attention_heads: gguf.num_attention_heads,
49            num_key_value_heads: gguf.num_key_value_heads,
50            rms_norm_eps: gguf.rms_norm_eps,
51            rope_theta: gguf.rope_theta,
52            ..Default::default()
53        }
54    }
55}
56
57pub struct Phi3VisionModelV2 {
58    config: Phi3VisionConfig,
59    device: Device,
60    vision_encoder: Phi3VisionEncoder,
61    projector: Phi3VisionProjector,
62    embed_tokens: Tensor,
63    layers: Vec<Phi3VisionDecoderLayer>,
64    norm: Tensor,
65    lm_head: Tensor,
66}
67
68pub struct Phi3VisionEncoder {
69    patch_embed: Tensor,
70    cls_token: Tensor,
71    pos_embed: Tensor,
72    blocks: Vec<VitBlock>,
73    norm: Tensor,
74    config: Phi3VisionConfig,
75}
76
77pub struct VitBlock {
78    norm1: Tensor,
79    attn_qkv: Tensor,
80    attn_proj: Tensor,
81    norm2: Tensor,
82    mlp_fc1: Tensor,
83    mlp_fc2: Tensor,
84    num_heads: usize,
85    head_dim: usize,
86}
87
88pub struct Phi3VisionProjector {
89    linear1: Tensor,
90    linear2: Tensor,
91}
92
93pub struct Phi3VisionDecoderLayer {
94    self_attn_qkv: Tensor,
95    self_attn_o: Tensor,
96    mlp_gate_up: Tensor,
97    mlp_down: Tensor,
98    input_layernorm: Tensor,
99    post_attention_layernorm: Tensor,
100    num_heads: usize,
101    num_kv_heads: usize,
102    head_dim: usize,
103}
104
105impl Model for Phi3VisionModelV2 {
106    type Config = Phi3VisionConfig;
107
108    fn new(config: Phi3VisionConfig) -> Result<Self> {
109        let device = Device::CPU;
110
111        let vision_encoder = Phi3VisionEncoder::new(&config, &device)?;
112        let projector = Phi3VisionProjector::new(&config, &device)?;
113        let embed_tokens = ops_fn::zeros(&[config.vocab_size, config.hidden_size], DataType::Float32, &device)?;
114        let norm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, &device)?;
115        let lm_head = ops_fn::zeros(&[config.hidden_size, config.vocab_size], DataType::Float32, &device)?;
116
117        let mut layers = Vec::with_capacity(config.num_hidden_layers);
118        for _ in 0..config.num_hidden_layers {
119            layers.push(Phi3VisionDecoderLayer::new(&config, &device)?);
120        }
121
122        Ok(Self { config, device, vision_encoder, projector, embed_tokens, layers, norm, lm_head })
123    }
124
125    fn from_weights(config: Phi3VisionConfig, weights: ModelWeights) -> Result<Self> {
126        let mut model = Self::new(config)?;
127        model.load_weights(&weights)?;
128        Ok(model)
129    }
130
131    fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
132        match inputs {
133            ModelInputs::Text { input_ids, .. } => {
134                let mut hidden = ops_fn::embedding(input_ids, &self.embed_tokens)?;
135
136                for layer in &self.layers {
137                    hidden = layer.forward(&hidden)?;
138                }
139
140                hidden = ops_fn::rms_norm(&hidden, &self.norm, self.config.rms_norm_eps)?;
141                let logits = ops_fn::matmul(&hidden, &self.lm_head)?;
142
143                Ok(ModelOutputs::Logits { logits, hidden_states: None })
144            }
145            _ => Err(anyhow::anyhow!("Phi-3-Vision requires text input")),
146        }
147    }
148
149    fn generate(&self, prompt: &str, config: &GenerationConfig) -> Result<String> {
150        use crate::tokenizer::Tokenizer;
151        let tokenizer = Tokenizer::new();
152        let mut tokens: Vec<u32> = tokenizer.encode(prompt);
153
154        for _ in 0..config.max_new_tokens {
155            let input_ids = Tensor::from_i64_slice(
156                &tokens.iter().map(|&t| t as i64).collect::<Vec<_>>(),
157                &[1, tokens.len()],
158                &self.device
159            )?;
160            let inputs = ModelInputs::text(input_ids);
161            let outputs = self.forward(&inputs)?;
162
163            let logits = match outputs {
164                ModelOutputs::Logits { logits, .. } => logits,
165                _ => return Err(anyhow::anyhow!("Expected logits")),
166            };
167
168            let logits_vec: Vec<f32> = logits.to_candle()?.flatten_all()?.to_vec1()?;
169            let seq_len = tokens.len();
170            let start = (seq_len - 1) * self.config.vocab_size;
171            let last_logits = &logits_vec[start..start + self.config.vocab_size];
172
173            let next_token = last_logits.iter()
174                .enumerate()
175                .max_by(|(_, a), (_, b)| a.partial_cmp(b).unwrap())
176                .map(|(idx, _)| idx as u32)
177                .unwrap_or(0);
178
179            if next_token == config.eos_token_id {
180                break;
181            }
182            tokens.push(next_token);
183        }
184
185        Ok(tokenizer.decode(&tokens))
186    }
187
188    fn config(&self) -> &Self::Config { &self.config }
189
190    fn memory_requirements(&self) -> MemoryRequirements {
191        let param_size = (self.config.vocab_size * self.config.hidden_size +
192            self.config.num_hidden_layers * 4 * self.config.hidden_size * self.config.hidden_size) * 4;
193        MemoryRequirements {
194            gpu_memory: param_size,
195            cpu_memory: param_size / 4,
196            kv_cache_memory: param_size / 8,
197            peak_memory: param_size + param_size / 2,
198        }
199    }
200
201    fn to_device(&mut self, device: &Device) -> Result<()> {
202        self.device = device.clone();
203        self.embed_tokens = self.embed_tokens.to_device(device)?;
204        self.norm = self.norm.to_device(device)?;
205        self.lm_head = self.lm_head.to_device(device)?;
206        Ok(())
207    }
208}
209
210impl Phi3VisionModelV2 {
211    fn load_weights(&mut self, weights: &ModelWeights) -> Result<()> {
212        if let Some(w) = weights.get("model.embed_tokens.weight") {
213            self.embed_tokens = w.clone();
214        }
215        if let Some(w) = weights.get("model.norm.weight") {
216            self.norm = w.clone();
217        }
218        if let Some(w) = weights.get("lm_head.weight") {
219            self.lm_head = ops_fn::transpose(w)?;
220        }
221        Ok(())
222    }
223}
224
225impl Phi3VisionEncoder {
226    fn new(config: &Phi3VisionConfig, device: &Device) -> Result<Self> {
227        let num_patches = (config.vision_image_size / config.vision_patch_size).pow(2);
228        let patch_dim = 3 * config.vision_patch_size * config.vision_patch_size;
229
230        let patch_embed = ops_fn::zeros(&[patch_dim, config.vision_hidden_size], DataType::Float32, device)?;
231        let cls_token = ops_fn::zeros(&[1, 1, config.vision_hidden_size], DataType::Float32, device)?;
232        let pos_embed = ops_fn::zeros(&[1, num_patches + 1, config.vision_hidden_size], DataType::Float32, device)?;
233        let norm = ops_fn::zeros(&[config.vision_hidden_size], DataType::Float32, device)?;
234
235        let mut blocks = Vec::with_capacity(config.vision_num_hidden_layers);
236        for _ in 0..config.vision_num_hidden_layers {
237            blocks.push(VitBlock::new(config, device)?);
238        }
239
240        Ok(Self { patch_embed, cls_token, pos_embed, blocks, norm, config: config.clone() })
241    }
242}
243
244impl VitBlock {
245    fn new(config: &Phi3VisionConfig, device: &Device) -> Result<Self> {
246        let head_dim = config.vision_hidden_size / config.vision_num_attention_heads;
247
248        Ok(Self {
249            norm1: ops_fn::zeros(&[config.vision_hidden_size], DataType::Float32, device)?,
250            attn_qkv: ops_fn::zeros(&[config.vision_hidden_size, config.vision_hidden_size * 3], DataType::Float32, device)?,
251            attn_proj: ops_fn::zeros(&[config.vision_hidden_size, config.vision_hidden_size], DataType::Float32, device)?,
252            norm2: ops_fn::zeros(&[config.vision_hidden_size], DataType::Float32, device)?,
253            mlp_fc1: ops_fn::zeros(&[config.vision_hidden_size, config.vision_intermediate_size], DataType::Float32, device)?,
254            mlp_fc2: ops_fn::zeros(&[config.vision_intermediate_size, config.vision_hidden_size], DataType::Float32, device)?,
255            num_heads: config.vision_num_attention_heads,
256            head_dim,
257        })
258    }
259}
260
261impl Phi3VisionProjector {
262    fn new(config: &Phi3VisionConfig, device: &Device) -> Result<Self> {
263        Ok(Self {
264            linear1: ops_fn::zeros(&[config.vision_hidden_size, config.hidden_size], DataType::Float32, device)?,
265            linear2: ops_fn::zeros(&[config.hidden_size, config.hidden_size], DataType::Float32, device)?,
266        })
267    }
268}
269
270impl Phi3VisionDecoderLayer {
271    fn new(config: &Phi3VisionConfig, device: &Device) -> Result<Self> {
272        let head_dim = config.hidden_size / config.num_attention_heads;
273        let qkv_dim = config.hidden_size + 2 * (config.num_key_value_heads * head_dim);
274
275        Ok(Self {
276            self_attn_qkv: ops_fn::zeros(&[config.hidden_size, qkv_dim], DataType::Float32, device)?,
277            self_attn_o: ops_fn::zeros(&[config.hidden_size, config.hidden_size], DataType::Float32, device)?,
278            mlp_gate_up: ops_fn::zeros(&[config.hidden_size, config.intermediate_size * 2], DataType::Float32, device)?,
279            mlp_down: ops_fn::zeros(&[config.intermediate_size, config.hidden_size], DataType::Float32, device)?,
280            input_layernorm: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
281            post_attention_layernorm: ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?,
282            num_heads: config.num_attention_heads,
283            num_kv_heads: config.num_key_value_heads,
284            head_dim,
285        })
286    }
287
288    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
289        let shape = hidden_states.shape();
290        let (batch_size, seq_len, _) = (shape[0], shape[1], shape[2]);
291
292        // Attention
293        let residual = hidden_states.clone();
294        let hidden = ops_fn::rms_norm(hidden_states, &self.input_layernorm, 1e-5)?;
295
296        let qkv = ops_fn::matmul(&hidden, &self.self_attn_qkv)?;
297        let qkv_candle = qkv.to_candle()?;
298
299        let q_dim = self.num_heads * self.head_dim;
300        let kv_dim = self.num_kv_heads * self.head_dim;
301
302        let q = qkv_candle.narrow(2, 0, q_dim)?
303            .reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?
304            .transpose(1, 2)?;
305        let k = qkv_candle.narrow(2, q_dim, kv_dim)?
306            .reshape(&[batch_size, seq_len, self.num_kv_heads, self.head_dim])?
307            .transpose(1, 2)?;
308        let v = qkv_candle.narrow(2, q_dim + kv_dim, kv_dim)?
309            .reshape(&[batch_size, seq_len, self.num_kv_heads, self.head_dim])?
310            .transpose(1, 2)?;
311
312        // GQA expansion
313        let num_groups = self.num_heads / self.num_kv_heads;
314        let (k, v) = if num_groups > 1 {
315            let k = k.unsqueeze(2)?
316                .broadcast_as(&[batch_size, self.num_kv_heads, num_groups, seq_len, self.head_dim])?
317                .reshape(&[batch_size, self.num_heads, seq_len, self.head_dim])?;
318            let v = v.unsqueeze(2)?
319                .broadcast_as(&[batch_size, self.num_kv_heads, num_groups, seq_len, self.head_dim])?
320                .reshape(&[batch_size, self.num_heads, seq_len, self.head_dim])?;
321            (k, v)
322        } else {
323            (k, v)
324        };
325
326        let scale = (self.head_dim as f32).powf(-0.5);
327        let scores = q.contiguous()?.matmul(&k.transpose(2, 3)?.contiguous()?)?;
328        let scores = (scores * (scale as f64))?;
329
330        // Causal mask
331        let device = scores.device();
332        let mask = {
333            let mut mask_data = vec![0.0f32; seq_len * seq_len];
334            for i in 0..seq_len {
335                for j in (i + 1)..seq_len {
336                    mask_data[i * seq_len + j] = f32::NEG_INFINITY;
337                }
338            }
339            candle_core::Tensor::from_vec(mask_data, &[1, 1, seq_len, seq_len], device)?
340        };
341
342        let scores = scores.broadcast_add(&mask)?;
343        let attn_weights = candle_nn::ops::softmax_last_dim(&scores)?;
344        let attn_output = attn_weights.matmul(&v.contiguous()?)?;
345
346        let attn_output = attn_output
347            .transpose(1, 2)?
348            .reshape(&[batch_size, seq_len, self.num_heads * self.head_dim])?;
349
350        let attn_output = Tensor::from_candle(attn_output);
351        let attn_output = ops_fn::matmul(&attn_output, &self.self_attn_o)?;
352        let hidden = ops_fn::add(&residual, &attn_output)?;
353
354        // MLP
355        let residual = hidden.clone();
356        let hidden = ops_fn::rms_norm(&hidden, &self.post_attention_layernorm, 1e-5)?;
357
358        let gate_up = ops_fn::matmul(&hidden, &self.mlp_gate_up)?;
359        let gate_up_candle = gate_up.to_candle()?;
360        let gate = gate_up_candle.narrow(2, 0, self.mlp_down.shape()[0])?;
361        let up = gate_up_candle.narrow(2, self.mlp_down.shape()[0], self.mlp_down.shape()[0])?;
362
363        let gate = candle_nn::ops::silu(&gate)?;
364        let hidden = gate.mul(&up)?;
365        let hidden = Tensor::from_candle(hidden);
366        let hidden = ops_fn::matmul(&hidden, &self.mlp_down)?;
367
368        ops_fn::add(&residual, &hidden)
369    }
370}
371
372#[cfg(test)]
373mod tests {
374    use super::*;
375
376    #[test]
377    fn test_phi3_vision_config() {
378        let config = Phi3VisionConfig::default();
379        assert_eq!(config.vocab_size, 32064);
380        assert_eq!(config.hidden_size, 3072);
381    }
382
383    #[test]
384    fn test_phi3_vision_model_creation() {
385        let config = Phi3VisionConfig {
386            vocab_size: 100,
387            hidden_size: 32,
388            intermediate_size: 128,
389            num_hidden_layers: 1,
390            num_attention_heads: 2,
391            num_key_value_heads: 2,
392            vision_hidden_size: 16,
393            vision_intermediate_size: 64,
394            vision_num_hidden_layers: 1,
395            vision_num_attention_heads: 2,
396            ..Default::default()
397        };
398
399        let model = Phi3VisionModelV2::new(config).unwrap();
400        assert_eq!(model.config().vocab_size(), 100);
401    }
402
403    #[test]
404    fn test_phi3_vision_forward() {
405        let config = Phi3VisionConfig {
406            vocab_size: 100,
407            hidden_size: 32,
408            intermediate_size: 128,
409            num_hidden_layers: 1,
410            num_attention_heads: 2,
411            num_key_value_heads: 2,
412            vision_hidden_size: 16,
413            vision_intermediate_size: 64,
414            vision_num_hidden_layers: 1,
415            vision_num_attention_heads: 2,
416            ..Default::default()
417        };
418
419        let model = Phi3VisionModelV2::new(config).unwrap();
420        let input_ids = ops_fn::zeros(&[1, 4], DataType::Int64, &Device::CPU).unwrap();
421        let inputs = ModelInputs::text(input_ids);
422
423        let outputs = model.forward(&inputs).unwrap();
424        match outputs {
425            ModelOutputs::Logits { logits, .. } => {
426                assert_eq!(logits.shape(), &[1, 4, 100]);
427            }
428            _ => panic!("Expected logits output"),
429        }
430    }
431}