Skip to main content

runtime/models_v2/
phi.rs

1//! Phi Model V2 - Clean implementation using solid abstractions
2//!
3//! This implements the Phi architecture (Phi-1, Phi-2, Phi-3) which includes:
4//! - Microsoft Phi-1/1.5/2/3 models
5//! - Uses unified Tensor type from tensor_core
6//! - Implements Model trait from model_core
7
8use crate::model_config;
9use super::traits::*;
10use anyhow::Result;
11use serde::{Serialize, Deserialize};
12
13/// Phi model configuration using the model_config macro
14model_config!(PhiConfig {
15    vocab_size: usize = 51200,
16    hidden_size: usize = 2560,
17    intermediate_size: usize = 10240,
18    num_hidden_layers: usize = 32,
19    num_attention_heads: usize = 32,
20    num_key_value_heads: usize = 32,
21    hidden_act: String = "gelu_new".to_string(),
22    max_position_embeddings: usize = 2048,
23    initializer_range: f32 = 0.02,
24    rms_norm_eps: f32 = 1e-5,
25    use_cache: bool = true,
26    pad_token_id: i64 = 0,
27    bos_token_id: i64 = 1,
28    eos_token_id: i64 = 2,
29    tie_word_embeddings: bool = false,
30    rope_theta: f32 = 10000.0,
31    attention_dropout: f32 = 0.0,
32    partial_rotary_factor: f32 = 0.5,
33});
34
35impl PhiConfig {
36    /// Create PhiConfig from GGUF model configuration
37    pub fn from_gguf_config(gguf: &crate::weight_loader_core::GGUFModelConfig) -> Self {
38        Self {
39            vocab_size: gguf.vocab_size,
40            hidden_size: gguf.hidden_size,
41            intermediate_size: gguf.intermediate_size,
42            num_hidden_layers: gguf.num_hidden_layers,
43            num_attention_heads: gguf.num_attention_heads,
44            num_key_value_heads: gguf.num_key_value_heads,
45            rms_norm_eps: gguf.rms_norm_eps,
46            rope_theta: gguf.rope_theta,
47            max_position_embeddings: gguf.max_position_embeddings,
48            ..Default::default()
49        }
50    }
51}
52
53/// Main Phi model implementation
54pub struct PhiModelV2 {
55    config: PhiConfig,
56    device: Device,
57
58    // Model components using unified Tensor type
59    embed_tokens: Tensor,
60    layers: Vec<PhiLayer>,
61    norm: Tensor,
62    lm_head: Tensor,
63}
64
65/// Phi transformer layer
66pub struct PhiLayer {
67    self_attn: PhiAttention,
68    mlp: PhiMLP,
69    input_layernorm: Tensor,
70}
71
72/// Phi attention mechanism
73pub struct PhiAttention {
74    q_proj: Tensor,
75    k_proj: Tensor,
76    v_proj: Tensor,
77    dense: Tensor, // Phi uses 'dense' instead of 'o_proj'
78    num_heads: usize,
79    num_key_value_heads: usize,
80    head_dim: usize,
81    scale: f32,
82    partial_rotary_factor: f32,
83}
84
85/// Phi MLP (feed-forward network) - standard 2-layer MLP, not SwiGLU
86pub struct PhiMLP {
87    fc1: Tensor,
88    fc2: Tensor,
89    hidden_act: String,
90}
91
92impl Model for PhiModelV2 {
93    type Config = PhiConfig;
94
95    fn new(config: PhiConfig) -> Result<Self> {
96        let device = Device::CPU;
97
98        let embed_tokens = ops_fn::zeros(
99            &[config.vocab_size, config.hidden_size],
100            DataType::Float32,
101            &device
102        )?;
103
104        let norm = ops_fn::zeros(
105            &[config.hidden_size],
106            DataType::Float32,
107            &device
108        )?;
109
110        let lm_head = if config.tie_word_embeddings {
111            embed_tokens.clone()
112        } else {
113            ops_fn::zeros(
114                &[config.hidden_size, config.vocab_size],
115                DataType::Float32,
116                &device
117            )?
118        };
119
120        // Create transformer layers
121        let mut layers = Vec::with_capacity(config.num_hidden_layers);
122        for _ in 0..config.num_hidden_layers {
123            layers.push(PhiLayer::new(&config, &device)?);
124        }
125
126        Ok(Self {
127            config,
128            device,
129            embed_tokens,
130            layers,
131            norm,
132            lm_head,
133        })
134    }
135
136    fn from_weights(config: PhiConfig, weights: ModelWeights) -> Result<Self> {
137        let mut model = Self::new(config)?;
138
139        // Load weights from the unified weight container
140        if let Some(embed_weights) = weights.get("model.embed_tokens.weight") {
141            model.embed_tokens = embed_weights.clone();
142        }
143
144        // Phi uses final_layernorm instead of norm
145        if let Some(norm_weights) = weights.get("model.final_layernorm.weight") {
146            model.norm = norm_weights.clone();
147        } else if let Some(norm_weights) = weights.get("model.norm.weight") {
148            model.norm = norm_weights.clone();
149        }
150
151        // lm_head weight needs transpose for matmul
152        if let Some(lm_head_weights) = weights.get("lm_head.weight") {
153            model.lm_head = ops_fn::transpose(lm_head_weights)?;
154        }
155
156        // Load layer weights
157        for (i, layer) in model.layers.iter_mut().enumerate() {
158            layer.load_weights(&weights, i)?;
159        }
160
161        Ok(model)
162    }
163
164    fn forward(&self, inputs: &ModelInputs) -> Result<ModelOutputs> {
165        match inputs {
166            ModelInputs::Text { input_ids, attention_mask, .. } => {
167                // 1. Token embedding
168                let mut hidden_states = ops_fn::embedding(input_ids, &self.embed_tokens)?;
169
170                // 2. Apply transformer layers (with partial RoPE)
171                for layer in &self.layers {
172                    hidden_states = layer.forward(&hidden_states, attention_mask.as_ref(), self.config.rope_theta, self.config.partial_rotary_factor)?;
173                }
174
175                // 3. Final layer norm
176                hidden_states = ops_fn::layer_norm(&hidden_states, &self.norm, None, self.config.rms_norm_eps)?;
177
178                // 4. Language modeling head
179                let logits = ops_fn::matmul(&hidden_states, &self.lm_head)?;
180
181                Ok(ModelOutputs::Logits {
182                    logits,
183                    hidden_states: None,
184                })
185            }
186            ModelInputs::Multimodal { input_ids, .. } => {
187                let text_inputs = ModelInputs::Text {
188                    input_ids: input_ids.clone(),
189                    attention_mask: None,
190                    position_ids: None,
191                };
192                self.forward(&text_inputs)
193            }
194            _ => Err(anyhow::anyhow!("Phi model only supports text and multimodal inputs")),
195        }
196    }
197
198    fn generate(&self, prompt: &str, config: &GenerationConfig) -> Result<String> {
199        use crate::tokenizer::Tokenizer;
200        use rand::Rng;
201
202        let tokenizer = Tokenizer::new();
203        let mut tokens: Vec<u32> = tokenizer.encode(prompt);
204
205        for _ in 0..config.max_new_tokens {
206            let tokens_i64: Vec<i64> = tokens.iter().map(|&t| t as i64).collect();
207            let input_tensor = Tensor::from_i64_slice(&tokens_i64, &[1, tokens.len()], &self.device)?;
208
209            let inputs = ModelInputs::Text {
210                input_ids: input_tensor,
211                attention_mask: None,
212                position_ids: None,
213            };
214
215            let outputs = self.forward(&inputs)?;
216
217            let logits = match outputs {
218                ModelOutputs::Logits { logits, .. } => logits,
219                _ => return Err(anyhow::anyhow!("Expected logits output")),
220            };
221
222            let logits_candle = logits.to_candle()?;
223            let shape = logits_candle.dims();
224
225            let last_logits = if shape.len() == 3 {
226                let seq_len = shape[1];
227                logits_candle
228                    .narrow(1, seq_len - 1, 1)?
229                    .squeeze(1)?
230                    .squeeze(0)?
231            } else {
232                let seq_len = shape[0];
233                logits_candle
234                    .narrow(0, seq_len - 1, 1)?
235                    .squeeze(0)?
236            };
237
238            let logits_vec: Vec<f32> = last_logits.to_vec1()?;
239
240            let next_token = if config.do_sample && config.temperature > 0.0 {
241                let scaled: Vec<f32> = logits_vec.iter()
242                    .map(|&x| x / config.temperature)
243                    .collect();
244
245                let max_val = scaled.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
246                let exp_sum: f32 = scaled.iter().map(|&x| (x - max_val).exp()).sum();
247                let probs: Vec<f32> = scaled.iter()
248                    .map(|&x| (x - max_val).exp() / exp_sum)
249                    .collect();
250
251                let mut rng = rand::thread_rng();
252                let random_val: f32 = rng.gen();
253                let mut cumulative = 0.0;
254                let mut sampled = 0u32;
255
256                for (idx, &prob) in probs.iter().enumerate() {
257                    cumulative += prob;
258                    if random_val <= cumulative {
259                        sampled = idx as u32;
260                        break;
261                    }
262                }
263                sampled
264            } else {
265                let mut max_idx = 0;
266                let mut max_val = logits_vec[0];
267                for (idx, &val) in logits_vec.iter().enumerate() {
268                    if val > max_val {
269                        max_val = val;
270                        max_idx = idx;
271                    }
272                }
273                max_idx as u32
274            };
275
276            if next_token == config.eos_token_id {
277                break;
278            }
279
280            tokens.push(next_token);
281        }
282
283        Ok(tokenizer.decode(&tokens))
284    }
285
286    fn config(&self) -> &Self::Config {
287        &self.config
288    }
289
290    fn memory_requirements(&self) -> MemoryRequirements {
291        let param_size = self.config.vocab_size * self.config.hidden_size +
292                        self.config.num_hidden_layers * (
293                            4 * self.config.hidden_size * self.config.hidden_size +
294                            2 * self.config.hidden_size * self.config.intermediate_size
295                        );
296
297        let param_bytes = param_size * 4;
298        let kv_cache_bytes = 2 * self.config.num_hidden_layers *
299                           self.config.max_position_embeddings *
300                           self.config.hidden_size * 4;
301
302        MemoryRequirements {
303            gpu_memory: param_bytes,
304            cpu_memory: param_bytes / 4,
305            kv_cache_memory: kv_cache_bytes,
306            peak_memory: param_bytes + kv_cache_bytes,
307        }
308    }
309
310    fn to_device(&mut self, device: &Device) -> Result<()> {
311        self.embed_tokens = self.embed_tokens.to_device(device)?;
312        self.norm = self.norm.to_device(device)?;
313        self.lm_head = self.lm_head.to_device(device)?;
314
315        for layer in &mut self.layers {
316            layer.to_device(device)?;
317        }
318
319        self.device = device.clone();
320        Ok(())
321    }
322}
323
324impl PhiLayer {
325    fn new(config: &PhiConfig, device: &Device) -> Result<Self> {
326        let input_layernorm = ops_fn::zeros(&[config.hidden_size], DataType::Float32, device)?;
327
328        Ok(Self {
329            self_attn: PhiAttention::new(config, device)?,
330            mlp: PhiMLP::new(config, device)?,
331            input_layernorm,
332        })
333    }
334
335    fn forward(&self, hidden_states: &Tensor, attention_mask: Option<&Tensor>, rope_theta: f32, partial_rotary_factor: f32) -> Result<Tensor> {
336        // Phi uses parallel attention and MLP (like Falcon)
337        let residual = hidden_states.clone();
338
339        // Layer norm
340        let normed = ops_fn::layer_norm(hidden_states, &self.input_layernorm, None, 1e-5)?;
341
342        // Attention and MLP in parallel
343        let attn_output = self.self_attn.forward(&normed, attention_mask, rope_theta, partial_rotary_factor)?;
344        let mlp_output = self.mlp.forward(&normed)?;
345
346        // Add both outputs to residual
347        let combined = ops_fn::add(&attn_output, &mlp_output)?;
348        let output = ops_fn::add(&residual, &combined)?;
349
350        Ok(output)
351    }
352
353    fn load_weights(&mut self, weights: &ModelWeights, layer_idx: usize) -> Result<()> {
354        let prefix = format!("model.layers.{}", layer_idx);
355
356        // Load attention weights (transpose for matmul)
357        if let Some(q_proj) = weights.get(&format!("{}.self_attn.q_proj.weight", prefix)) {
358            self.self_attn.q_proj = ops_fn::transpose(q_proj)?;
359        }
360        if let Some(k_proj) = weights.get(&format!("{}.self_attn.k_proj.weight", prefix)) {
361            self.self_attn.k_proj = ops_fn::transpose(k_proj)?;
362        }
363        if let Some(v_proj) = weights.get(&format!("{}.self_attn.v_proj.weight", prefix)) {
364            self.self_attn.v_proj = ops_fn::transpose(v_proj)?;
365        }
366        // Phi uses 'dense' for output projection
367        if let Some(dense) = weights.get(&format!("{}.self_attn.dense.weight", prefix)) {
368            self.self_attn.dense = ops_fn::transpose(dense)?;
369        } else if let Some(o_proj) = weights.get(&format!("{}.self_attn.o_proj.weight", prefix)) {
370            self.self_attn.dense = ops_fn::transpose(o_proj)?;
371        }
372
373        // Load MLP weights (transpose for matmul)
374        if let Some(fc1) = weights.get(&format!("{}.mlp.fc1.weight", prefix)) {
375            self.mlp.fc1 = ops_fn::transpose(fc1)?;
376        }
377        if let Some(fc2) = weights.get(&format!("{}.mlp.fc2.weight", prefix)) {
378            self.mlp.fc2 = ops_fn::transpose(fc2)?;
379        }
380
381        // Load layer norm weight
382        if let Some(input_ln) = weights.get(&format!("{}.input_layernorm.weight", prefix)) {
383            self.input_layernorm = input_ln.clone();
384        }
385
386        Ok(())
387    }
388
389    fn to_device(&mut self, device: &Device) -> Result<()> {
390        self.self_attn.to_device(device)?;
391        self.mlp.to_device(device)?;
392        self.input_layernorm = self.input_layernorm.to_device(device)?;
393        Ok(())
394    }
395}
396
397/// Apply partial Rotary Position Embedding (RoPE) to Q and K tensors
398/// Only applies RoPE to first (partial_rotary_factor * head_dim) dimensions
399fn apply_partial_rope(
400    q: &candle_core::Tensor,
401    k: &candle_core::Tensor,
402    seq_len: usize,
403    head_dim: usize,
404    rope_theta: f32,
405    partial_rotary_factor: f32,
406) -> Result<(candle_core::Tensor, candle_core::Tensor)> {
407    let device = q.device();
408
409    // Calculate how much of head_dim to apply RoPE to
410    let rotary_dim = ((head_dim as f32 * partial_rotary_factor) as usize / 2) * 2; // Make even
411    let half_rotary_dim = rotary_dim / 2;
412
413    if rotary_dim == 0 {
414        // No rotation needed
415        return Ok((q.clone(), k.clone()));
416    }
417
418    // Compute inverse frequencies for the rotary dimensions
419    let inv_freq: Vec<f32> = (0..half_rotary_dim)
420        .map(|i| 1.0 / rope_theta.powf((2 * i) as f32 / rotary_dim as f32))
421        .collect();
422
423    let positions: Vec<f32> = (0..seq_len).map(|p| p as f32).collect();
424
425    let mut angles = Vec::with_capacity(seq_len * half_rotary_dim);
426    for pos in &positions {
427        for freq in &inv_freq {
428            angles.push(pos * freq);
429        }
430    }
431
432    let angles_tensor = candle_core::Tensor::from_vec(angles, &[seq_len, half_rotary_dim], device)?;
433
434    let cos = angles_tensor.cos()?;
435    let sin = angles_tensor.sin()?;
436
437    let cos = cos.unsqueeze(0)?.unsqueeze(0)?;
438    let sin = sin.unsqueeze(0)?.unsqueeze(0)?;
439
440    // Split into rotary and non-rotary parts
441    let q_rot_part = q.narrow(3, 0, rotary_dim)?;
442    let q_pass_part = q.narrow(3, rotary_dim, head_dim - rotary_dim)?;
443    let k_rot_part = k.narrow(3, 0, rotary_dim)?;
444    let k_pass_part = k.narrow(3, rotary_dim, head_dim - rotary_dim)?;
445
446    // Apply rotation to rotary part
447    let q_half1 = q_rot_part.narrow(3, 0, half_rotary_dim)?;
448    let q_half2 = q_rot_part.narrow(3, half_rotary_dim, half_rotary_dim)?;
449    let k_half1 = k_rot_part.narrow(3, 0, half_rotary_dim)?;
450    let k_half2 = k_rot_part.narrow(3, half_rotary_dim, half_rotary_dim)?;
451
452    let q_rot1 = (q_half1.broadcast_mul(&cos)? - q_half2.broadcast_mul(&sin)?)?;
453    let q_rot2 = (q_half1.broadcast_mul(&sin)? + q_half2.broadcast_mul(&cos)?)?;
454    let k_rot1 = (k_half1.broadcast_mul(&cos)? - k_half2.broadcast_mul(&sin)?)?;
455    let k_rot2 = (k_half1.broadcast_mul(&sin)? + k_half2.broadcast_mul(&cos)?)?;
456
457    // Concatenate rotated and non-rotated parts
458    let q_rotated_part = candle_core::Tensor::cat(&[&q_rot1, &q_rot2], 3)?;
459    let k_rotated_part = candle_core::Tensor::cat(&[&k_rot1, &k_rot2], 3)?;
460
461    let q_rotated = candle_core::Tensor::cat(&[&q_rotated_part, &q_pass_part], 3)?;
462    let k_rotated = candle_core::Tensor::cat(&[&k_rotated_part, &k_pass_part], 3)?;
463
464    Ok((q_rotated, k_rotated))
465}
466
467impl PhiAttention {
468    fn new(config: &PhiConfig, device: &Device) -> Result<Self> {
469        let num_heads = config.num_attention_heads;
470        let num_key_value_heads = config.num_key_value_heads;
471        let head_dim = config.hidden_size / num_heads;
472        let scale = 1.0 / (head_dim as f32).sqrt();
473
474        let q_proj = ops_fn::zeros(&[config.hidden_size, num_heads * head_dim], DataType::Float32, device)?;
475        let k_proj = ops_fn::zeros(&[config.hidden_size, num_key_value_heads * head_dim], DataType::Float32, device)?;
476        let v_proj = ops_fn::zeros(&[config.hidden_size, num_key_value_heads * head_dim], DataType::Float32, device)?;
477        let dense = ops_fn::zeros(&[num_heads * head_dim, config.hidden_size], DataType::Float32, device)?;
478
479        Ok(Self {
480            q_proj,
481            k_proj,
482            v_proj,
483            dense,
484            num_heads,
485            num_key_value_heads,
486            head_dim,
487            scale,
488            partial_rotary_factor: config.partial_rotary_factor,
489        })
490    }
491
492    fn forward(&self, hidden_states: &Tensor, _attention_mask: Option<&Tensor>, rope_theta: f32, partial_rotary_factor: f32) -> Result<Tensor> {
493        let shape = hidden_states.shape();
494        let (batch_size, seq_len, _hidden_size) = if shape.len() == 3 {
495            (shape[0], shape[1], shape[2])
496        } else if shape.len() == 2 {
497            (1, shape[0], shape[1])
498        } else {
499            return Err(anyhow::anyhow!("Invalid hidden_states shape: {:?}", shape));
500        };
501
502        let query_states = ops_fn::matmul(hidden_states, &self.q_proj)?;
503        let key_states = ops_fn::matmul(hidden_states, &self.k_proj)?;
504        let value_states = ops_fn::matmul(hidden_states, &self.v_proj)?;
505
506        let q_candle = query_states.to_candle()?;
507        let k_candle = key_states.to_candle()?;
508        let v_candle = value_states.to_candle()?;
509
510        let q_reshaped = q_candle
511            .reshape(&[batch_size, seq_len, self.num_heads, self.head_dim])?
512            .transpose(1, 2)?;
513
514        let k_reshaped = k_candle
515            .reshape(&[batch_size, seq_len, self.num_key_value_heads, self.head_dim])?
516            .transpose(1, 2)?;
517
518        let v_reshaped = v_candle
519            .reshape(&[batch_size, seq_len, self.num_key_value_heads, self.head_dim])?
520            .transpose(1, 2)?;
521
522        // Apply partial RoPE
523        let (q_with_rope, k_with_rope) = apply_partial_rope(&q_reshaped, &k_reshaped, seq_len, self.head_dim, rope_theta, partial_rotary_factor)?;
524
525        // Handle GQA
526        let num_groups = self.num_heads / self.num_key_value_heads;
527        let (k_expanded, v_expanded) = if num_groups > 1 {
528            let k_rep = k_with_rope
529                .unsqueeze(2)?
530                .broadcast_as(&[batch_size, self.num_key_value_heads, num_groups, seq_len, self.head_dim])?
531                .reshape(&[batch_size, self.num_heads, seq_len, self.head_dim])?;
532            let v_rep = v_reshaped
533                .unsqueeze(2)?
534                .broadcast_as(&[batch_size, self.num_key_value_heads, num_groups, seq_len, self.head_dim])?
535                .reshape(&[batch_size, self.num_heads, seq_len, self.head_dim])?;
536            (k_rep, v_rep)
537        } else {
538            (k_with_rope, v_reshaped)
539        };
540
541        let k_t = k_expanded.transpose(2, 3)?;
542
543        let q_contiguous = q_with_rope.contiguous()?;
544        let k_contiguous = k_t.contiguous()?;
545
546        let scores = q_contiguous.matmul(&k_contiguous)?;
547        let scaled_scores = (scores * (self.scale as f64))?;
548
549        // Causal mask
550        let device = scaled_scores.device();
551        let causal_mask = {
552            let mut mask_data = vec![0.0f32; seq_len * seq_len];
553            for i in 0..seq_len {
554                for j in 0..seq_len {
555                    if j > i {
556                        mask_data[i * seq_len + j] = f32::NEG_INFINITY;
557                    }
558                }
559            }
560            candle_core::Tensor::from_vec(mask_data, &[1, 1, seq_len, seq_len], device)?
561        };
562
563        let masked_scores = scaled_scores.broadcast_add(&causal_mask)?;
564
565        let attention_weights = candle_nn::ops::softmax_last_dim(&masked_scores)?;
566
567        let v_contiguous = v_expanded.contiguous()?;
568        let attn_output = attention_weights.matmul(&v_contiguous)?;
569
570        let attn_output = attn_output
571            .transpose(1, 2)?
572            .reshape(&[batch_size, seq_len, self.num_heads * self.head_dim])?;
573
574        let attn_output = Tensor::from_candle(attn_output);
575
576        let output = ops_fn::matmul(&attn_output, &self.dense)?;
577
578        Ok(output)
579    }
580
581    fn to_device(&mut self, device: &Device) -> Result<()> {
582        self.q_proj = self.q_proj.to_device(device)?;
583        self.k_proj = self.k_proj.to_device(device)?;
584        self.v_proj = self.v_proj.to_device(device)?;
585        self.dense = self.dense.to_device(device)?;
586        Ok(())
587    }
588}
589
590impl PhiMLP {
591    fn new(config: &PhiConfig, device: &Device) -> Result<Self> {
592        let fc1 = ops_fn::zeros(&[config.hidden_size, config.intermediate_size], DataType::Float32, device)?;
593        let fc2 = ops_fn::zeros(&[config.intermediate_size, config.hidden_size], DataType::Float32, device)?;
594
595        Ok(Self {
596            fc1,
597            fc2,
598            hidden_act: config.hidden_act.clone(),
599        })
600    }
601
602    fn forward(&self, hidden_states: &Tensor) -> Result<Tensor> {
603        // First projection with activation
604        let intermediate = ops_fn::matmul(hidden_states, &self.fc1)?;
605
606        // Apply activation (Phi typically uses GELU)
607        let activated = match self.hidden_act.as_str() {
608            "gelu" | "gelu_new" => ops_fn::gelu(&intermediate)?,
609            "silu" | "swish" => ops_fn::silu(&intermediate)?,
610            _ => return Err(anyhow::anyhow!("Unsupported activation: {}", self.hidden_act)),
611        };
612
613        // Second projection
614        let output = ops_fn::matmul(&activated, &self.fc2)?;
615
616        Ok(output)
617    }
618
619    fn to_device(&mut self, device: &Device) -> Result<()> {
620        self.fc1 = self.fc1.to_device(device)?;
621        self.fc2 = self.fc2.to_device(device)?;
622        Ok(())
623    }
624}
625
626#[cfg(test)]
627mod tests {
628    use super::*;
629
630    #[test]
631    fn test_phi_model_creation() {
632        let config = PhiConfig {
633            vocab_size: 1000,
634            hidden_size: 128,
635            intermediate_size: 512,
636            num_hidden_layers: 2,
637            num_attention_heads: 8,
638            num_key_value_heads: 8,
639            ..Default::default()
640        };
641
642        let model = PhiModelV2::new(config).unwrap();
643        assert_eq!(model.config().vocab_size(), 1000);
644        assert_eq!(model.config().hidden_size(), 128);
645        assert_eq!(model.config().num_layers(), 2);
646    }
647
648    #[test]
649    fn test_phi_forward_pass() {
650        let config = PhiConfig {
651            vocab_size: 100,
652            hidden_size: 64,
653            intermediate_size: 256,
654            num_hidden_layers: 1,
655            num_attention_heads: 4,
656            num_key_value_heads: 4,
657            ..Default::default()
658        };
659
660        let model = PhiModelV2::new(config).unwrap();
661        let input_ids = ops_fn::zeros(&[2, 8], DataType::Int64, &Device::CPU).unwrap();
662        let inputs = ModelInputs::text(input_ids);
663
664        let outputs = model.forward(&inputs).unwrap();
665        match outputs {
666            ModelOutputs::Logits { logits, .. } => {
667                assert_eq!(logits.shape(), &[2, 8, 100]);
668            }
669            _ => panic!("Expected logits output"),
670        }
671    }
672
673    #[test]
674    fn test_phi_generation() {
675        let config = PhiConfig {
676            vocab_size: 256,
677            hidden_size: 64,
678            intermediate_size: 256,
679            num_hidden_layers: 1,
680            num_attention_heads: 4,
681            num_key_value_heads: 4,
682            ..Default::default()
683        };
684        let model = PhiModelV2::new(config).unwrap();
685        let gen_config = GenerationConfig {
686            max_new_tokens: 5,
687            ..Default::default()
688        };
689
690        let output = model.generate("Hello", &gen_config).unwrap();
691        assert!(!output.is_empty());
692    }
693}