1use crate::model_config;
11use super::traits::*;
12use anyhow::Result;
13use serde::{Serialize, Deserialize};
14
15model_config!(Phi3VisionConfig {
16 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_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 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 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 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 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}