1use crate::model_config;
9use super::traits::*;
10use anyhow::Result;
11use serde::{Serialize, Deserialize};
12
13model_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 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
53pub struct PhiModelV2 {
55 config: PhiConfig,
56 device: Device,
57
58 embed_tokens: Tensor,
60 layers: Vec<PhiLayer>,
61 norm: Tensor,
62 lm_head: Tensor,
63}
64
65pub struct PhiLayer {
67 self_attn: PhiAttention,
68 mlp: PhiMLP,
69 input_layernorm: Tensor,
70}
71
72pub struct PhiAttention {
74 q_proj: Tensor,
75 k_proj: Tensor,
76 v_proj: Tensor,
77 dense: Tensor, num_heads: usize,
79 num_key_value_heads: usize,
80 head_dim: usize,
81 scale: f32,
82 partial_rotary_factor: f32,
83}
84
85pub 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 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 if let Some(embed_weights) = weights.get("model.embed_tokens.weight") {
141 model.embed_tokens = embed_weights.clone();
142 }
143
144 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 if let Some(lm_head_weights) = weights.get("lm_head.weight") {
153 model.lm_head = ops_fn::transpose(lm_head_weights)?;
154 }
155
156 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 let mut hidden_states = ops_fn::embedding(input_ids, &self.embed_tokens)?;
169
170 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 hidden_states = ops_fn::layer_norm(&hidden_states, &self.norm, None, self.config.rms_norm_eps)?;
177
178 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 let residual = hidden_states.clone();
338
339 let normed = ops_fn::layer_norm(hidden_states, &self.input_layernorm, None, 1e-5)?;
341
342 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 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 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 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 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 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
397fn 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 let rotary_dim = ((head_dim as f32 * partial_rotary_factor) as usize / 2) * 2; let half_rotary_dim = rotary_dim / 2;
412
413 if rotary_dim == 0 {
414 return Ok((q.clone(), k.clone()));
416 }
417
418 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 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 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 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 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 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 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 let intermediate = ops_fn::matmul(hidden_states, &self.fc1)?;
605
606 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 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}