Skip to main content

candle_transformers/models/
phi3.rs

1//! Microsoft Phi-3 model implementation
2//!
3//! See Phi model details at:
4//! - [Phi-3 Model](https://huggingface.co/microsoft/phi-3)
5//!
6//! The Phi series are decoder-only transformers designed for code and language tasks.
7//! Key characteristics:
8//! - Decoder-only transformer architecture
9//! - RoPE embeddings
10//! - Layer normalization
11//! - QK normalization
12//! - Mixed activation functions
13//! - Improved context window handling
14//!
15//! References:
16//! - [Hugging Face Implementation](https://huggingface.co/microsoft/phi-3)
17//! - [Alternative Implementation](https://huggingface.co/microsoft/phi-3/tree/main)
18//!
19
20// This implementation is based on:
21// https://huggingface.co/microsoft/Phi-3-mini-4k-instruct/blob/main/modeling_phi3.py
22use crate::models::with_tracing::{linear_no_bias as linear, Linear, RmsNorm};
23use candle::{DType, Device, IndexOp, Module, Result, Tensor, D};
24use candle_nn::VarBuilder;
25use std::sync::Arc;
26
27#[derive(Debug, Clone, serde::Deserialize)]
28pub enum RopeScalingType {
29    #[serde(rename = "longrope")]
30    LongRope,
31}
32
33#[derive(Debug, Clone, serde::Deserialize)]
34pub struct RopeScaling {
35    pub short_factor: Vec<f32>,
36    pub long_factor: Vec<f32>,
37    #[serde(rename = "type")]
38    pub type_: RopeScalingType,
39}
40
41// https://huggingface.co/microsoft/Phi-3-mini-4k-instruct/blob/main/config.json
42#[derive(Debug, Clone, serde::Deserialize)]
43pub struct Config {
44    pub vocab_size: usize,
45    pub hidden_act: candle_nn::Activation,
46    pub hidden_size: usize,
47    pub intermediate_size: usize,
48    pub num_hidden_layers: usize,
49    pub num_attention_heads: usize,
50    pub num_key_value_heads: usize,
51    pub rms_norm_eps: f64,
52    pub rope_theta: f64,
53    pub bos_token_id: Option<u32>,
54    pub eos_token_id: Option<u32>,
55    pub rope_scaling: Option<RopeScaling>,
56    pub max_position_embeddings: usize,
57    pub original_max_position_embeddings: Option<usize>,
58    pub partial_rotary_factor: Option<f64>,
59    #[serde(default)]
60    pub tie_word_embeddings: bool,
61}
62
63impl Config {
64    pub fn head_dim(&self) -> usize {
65        self.hidden_size / self.num_attention_heads
66    }
67}
68
69#[derive(Debug, Clone)]
70pub struct RotaryEmbedding {
71    partial_dim: Option<usize>,
72    sin: Tensor,
73    cos: Tensor,
74}
75
76impl RotaryEmbedding {
77    pub fn new(dtype: DType, cfg: &Config, dev: &Device) -> Result<Self> {
78        let partial_dim = cfg
79            .partial_rotary_factor
80            .as_ref()
81            .map(|v| (v * cfg.head_dim() as f64) as usize);
82        let dim = partial_dim.unwrap_or(cfg.head_dim());
83        let freqs = match cfg.rope_scaling.as_ref() {
84            None => {
85                let max_seq_len = cfg.max_position_embeddings;
86                let inv_freq: Vec<_> = (0..dim)
87                    .step_by(2)
88                    .map(|i| 1f32 / cfg.rope_theta.powf(i as f64 / dim as f64) as f32)
89                    .collect();
90                let inv_freq = Tensor::from_vec(inv_freq, (1, ()), dev)?.to_dtype(dtype)?;
91                let t = Tensor::arange(0u32, max_seq_len as u32, dev)?
92                    .to_dtype(dtype)?
93                    .reshape((max_seq_len, 1))?;
94                t.matmul(&inv_freq)?
95            }
96            Some(rope_scaling) => {
97                let inv_freq_s: Vec<_> = (0..dim)
98                    .step_by(2)
99                    .zip(rope_scaling.short_factor.iter())
100                    .map(|(i, &f)| f / cfg.rope_theta.powf(i as f64 / dim as f64) as f32)
101                    .collect();
102                let inv_freq_s = Tensor::from_vec(inv_freq_s, (1, ()), dev)?.to_dtype(dtype)?;
103                let max_seq_len = cfg.max_position_embeddings;
104                match cfg.original_max_position_embeddings {
105                    None => {
106                        let t = Tensor::arange(0u32, max_seq_len as u32, dev)?
107                            .to_dtype(dtype)?
108                            .reshape((max_seq_len, 1))?;
109                        t.matmul(&inv_freq_s)?
110                    }
111                    Some(original_max_seq_len) => {
112                        let t_s = Tensor::arange(0u32, original_max_seq_len as u32, dev)?
113                            .to_dtype(dtype)?
114                            .reshape((original_max_seq_len, 1))?;
115                        let freq_s = t_s.matmul(&inv_freq_s)?;
116                        let inv_freq_l: Vec<_> = (0..dim)
117                            .step_by(2)
118                            .zip(rope_scaling.long_factor.iter())
119                            .map(|(i, &f)| f / cfg.rope_theta.powf(i as f64 / dim as f64) as f32)
120                            .collect();
121                        let inv_freq_l =
122                            Tensor::from_vec(inv_freq_l, (1, ()), dev)?.to_dtype(dtype)?;
123                        let t_l =
124                            Tensor::arange(original_max_seq_len as u32, max_seq_len as u32, dev)?
125                                .to_dtype(dtype)?
126                                .reshape(((), 1))?;
127                        let freq_l = t_l.matmul(&inv_freq_l)?;
128                        Tensor::cat(&[&freq_s, &freq_l], 0)?
129                    }
130                }
131            }
132        };
133        Ok(Self {
134            partial_dim,
135            sin: freqs.sin()?,
136            cos: freqs.cos()?,
137        })
138    }
139
140    fn rope(&self, xs: &Tensor, cos: &Tensor, sin: &Tensor) -> Result<Tensor> {
141        let x = match self.partial_dim {
142            None => candle_nn::rotary_emb::rope(&xs.contiguous()?, cos, sin)?,
143            Some(dim) => {
144                let xs_rot = xs.i((.., .., .., ..dim))?.contiguous()?;
145                let xs_pass = xs.i((.., .., .., dim..))?;
146                let xs_rot = candle_nn::rotary_emb::rope(&xs_rot, cos, sin)?;
147                Tensor::cat(&[&xs_rot, &xs_pass], D::Minus1)?.contiguous()?
148            }
149        };
150        Ok(x)
151    }
152
153    pub fn apply_rotary_emb_qkv(
154        &self,
155        q: &Tensor,
156        k: &Tensor,
157        seqlen_offset: usize,
158    ) -> Result<(Tensor, Tensor)> {
159        let (_b_sz, _h, seq_len, _n_embd) = q.dims4()?;
160        let cos = self.cos.narrow(0, seqlen_offset, seq_len)?;
161        let sin = self.sin.narrow(0, seqlen_offset, seq_len)?;
162        let q_embed = self.rope(&q.contiguous()?, &cos, &sin)?;
163        let k_embed = self.rope(&k.contiguous()?, &cos, &sin)?;
164        Ok((q_embed, k_embed))
165    }
166}
167
168#[derive(Debug, Clone)]
169struct Attention {
170    qkv_proj: Linear,
171    o_proj: Linear,
172    num_heads: usize,
173    num_kv_heads: usize,
174    num_kv_groups: usize,
175    head_dim: usize,
176    rotary_emb: Arc<RotaryEmbedding>,
177    kv_cache: Option<(Tensor, Tensor)>,
178}
179
180impl Attention {
181    fn new(rotary_emb: Arc<RotaryEmbedding>, cfg: &Config, vb: VarBuilder) -> Result<Self> {
182        let num_heads = cfg.num_attention_heads;
183        let num_kv_heads = cfg.num_key_value_heads;
184        let head_dim = cfg.head_dim();
185        let op_size = num_heads * head_dim + 2 * num_kv_heads * head_dim;
186        let qkv_proj = linear(cfg.hidden_size, op_size, vb.pp("qkv_proj"))?;
187        let o_proj = linear(num_heads * head_dim, cfg.hidden_size, vb.pp("o_proj"))?;
188        Ok(Self {
189            qkv_proj,
190            o_proj,
191            rotary_emb,
192            kv_cache: None,
193            num_heads,
194            num_kv_heads,
195            num_kv_groups: num_heads / num_kv_heads,
196            head_dim,
197        })
198    }
199
200    fn forward(
201        &mut self,
202        xs: &Tensor,
203        attention_mask: Option<&Tensor>,
204        seqlen_offset: usize,
205    ) -> Result<Tensor> {
206        let (b_sz, q_len, _) = xs.dims3()?;
207
208        let qkv = self.qkv_proj.forward(xs)?;
209        let query_pos = self.num_heads * self.head_dim;
210        let query_states = qkv.narrow(D::Minus1, 0, query_pos)?;
211        let key_states = qkv.narrow(D::Minus1, query_pos, self.num_kv_heads * self.head_dim)?;
212        let value_states = qkv.narrow(
213            D::Minus1,
214            query_pos + self.num_kv_heads * self.head_dim,
215            self.num_kv_heads * self.head_dim,
216        )?;
217
218        let query_states = query_states
219            .reshape((b_sz, q_len, self.num_heads, self.head_dim))?
220            .transpose(1, 2)?;
221        let key_states = key_states
222            .reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))?
223            .transpose(1, 2)?;
224        let value_states = value_states
225            .reshape((b_sz, q_len, self.num_kv_heads, self.head_dim))?
226            .transpose(1, 2)?;
227
228        let (query_states, key_states) =
229            self.rotary_emb
230                .apply_rotary_emb_qkv(&query_states, &key_states, seqlen_offset)?;
231
232        let (key_states, value_states) = match &self.kv_cache {
233            None => (key_states, value_states),
234            Some((prev_k, prev_v)) => {
235                let key_states = Tensor::cat(&[prev_k, &key_states], 2)?;
236                let value_states = Tensor::cat(&[prev_v, &value_states], 2)?;
237                (key_states, value_states)
238            }
239        };
240        self.kv_cache = Some((key_states.clone(), value_states.clone()));
241
242        let key_states = crate::utils::repeat_kv(key_states, self.num_kv_groups)?.contiguous()?;
243        let value_states =
244            crate::utils::repeat_kv(value_states, self.num_kv_groups)?.contiguous()?;
245
246        let attn_output = {
247            let scale = 1f64 / f64::sqrt(self.head_dim as f64);
248            let attn_weights = (query_states.matmul(&key_states.transpose(2, 3)?)? * scale)?;
249
250            let attn_weights = match attention_mask {
251                None => attn_weights,
252                Some(mask) => attn_weights.broadcast_add(mask)?,
253            };
254            let attn_weights = candle_nn::ops::softmax_last_dim(&attn_weights)?;
255            attn_weights.matmul(&value_states)?
256        };
257        attn_output
258            .transpose(1, 2)?
259            .reshape((b_sz, q_len, ()))?
260            .apply(&self.o_proj)
261    }
262
263    fn clear_kv_cache(&mut self) {
264        self.kv_cache = None
265    }
266}
267
268#[derive(Debug, Clone)]
269struct Mlp {
270    gate_up_proj: Linear,
271    down_proj: Linear,
272    act_fn: candle_nn::Activation,
273    i_size: usize,
274}
275
276impl Mlp {
277    fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
278        let hidden_size = cfg.hidden_size;
279        let i_size = cfg.intermediate_size;
280        let gate_up_proj = linear(hidden_size, 2 * i_size, vb.pp("gate_up_proj"))?;
281        let down_proj = linear(i_size, hidden_size, vb.pp("down_proj"))?;
282        Ok(Self {
283            gate_up_proj,
284            down_proj,
285            act_fn: cfg.hidden_act,
286            i_size,
287        })
288    }
289}
290
291impl Module for Mlp {
292    fn forward(&self, xs: &Tensor) -> Result<Tensor> {
293        let up_states = xs.apply(&self.gate_up_proj)?;
294        let gate = up_states.narrow(D::Minus1, 0, self.i_size)?;
295        let up_states = up_states.narrow(D::Minus1, self.i_size, self.i_size)?;
296        let up_states = (up_states * gate.apply(&self.act_fn))?;
297        up_states.apply(&self.down_proj)
298    }
299}
300
301#[derive(Debug, Clone)]
302struct DecoderLayer {
303    self_attn: Attention,
304    mlp: Mlp,
305    input_layernorm: RmsNorm,
306    post_attention_layernorm: RmsNorm,
307}
308
309impl DecoderLayer {
310    fn new(rotary_emb: Arc<RotaryEmbedding>, cfg: &Config, vb: VarBuilder) -> Result<Self> {
311        let self_attn = Attention::new(rotary_emb, cfg, vb.pp("self_attn"))?;
312        let mlp = Mlp::new(cfg, vb.pp("mlp"))?;
313        let input_layernorm =
314            RmsNorm::new(cfg.hidden_size, cfg.rms_norm_eps, vb.pp("input_layernorm"))?;
315        let post_attention_layernorm = RmsNorm::new(
316            cfg.hidden_size,
317            cfg.rms_norm_eps,
318            vb.pp("post_attention_layernorm"),
319        )?;
320        Ok(Self {
321            self_attn,
322            mlp,
323            input_layernorm,
324            post_attention_layernorm,
325        })
326    }
327
328    fn forward(
329        &mut self,
330        xs: &Tensor,
331        attention_mask: Option<&Tensor>,
332        seqlen_offset: usize,
333    ) -> Result<Tensor> {
334        let residual = xs;
335        let xs = self.input_layernorm.forward(xs)?;
336        let xs = self.self_attn.forward(&xs, attention_mask, seqlen_offset)?;
337        let xs = (xs + residual)?;
338        let residual = &xs;
339        let xs = xs.apply(&self.post_attention_layernorm)?.apply(&self.mlp)?;
340        residual + xs
341    }
342
343    fn clear_kv_cache(&mut self) {
344        self.self_attn.clear_kv_cache()
345    }
346}
347
348#[derive(Debug, Clone)]
349pub struct Model {
350    embed_tokens: candle_nn::Embedding,
351    layers: Vec<DecoderLayer>,
352    norm: RmsNorm,
353    lm_head: Linear,
354    device: Device,
355    dtype: DType,
356}
357
358impl Model {
359    pub fn new(cfg: &Config, vb: VarBuilder) -> Result<Self> {
360        let vb_m = vb.pp("model");
361        let embed_tokens =
362            candle_nn::embedding(cfg.vocab_size, cfg.hidden_size, vb_m.pp("embed_tokens"))?;
363        let rotary_emb = Arc::new(RotaryEmbedding::new(vb.dtype(), cfg, vb_m.device())?);
364        let mut layers = Vec::with_capacity(cfg.num_hidden_layers);
365        let vb_l = vb_m.pp("layers");
366        for layer_idx in 0..cfg.num_hidden_layers {
367            let layer = DecoderLayer::new(rotary_emb.clone(), cfg, vb_l.pp(layer_idx))?;
368            layers.push(layer)
369        }
370        let norm = RmsNorm::new(cfg.hidden_size, cfg.rms_norm_eps, vb_m.pp("norm"))?;
371        let lm_head = if cfg.tie_word_embeddings {
372            Linear::from_weights(embed_tokens.embeddings().clone(), None)
373        } else {
374            linear(cfg.hidden_size, cfg.vocab_size, vb.pp("lm_head"))?
375        };
376        Ok(Self {
377            embed_tokens,
378            layers,
379            norm,
380            lm_head,
381            device: vb.device().clone(),
382            dtype: vb.dtype(),
383        })
384    }
385
386    fn prepare_decoder_attention_mask(
387        &self,
388        b_size: usize,
389        tgt_len: usize,
390        seqlen_offset: usize,
391    ) -> Result<Tensor> {
392        let mask: Vec<_> = (0..tgt_len)
393            .flat_map(|i| (0..tgt_len).map(move |j| if i < j { f32::NEG_INFINITY } else { 0. }))
394            .collect();
395        let mask = Tensor::from_slice(&mask, (tgt_len, tgt_len), &self.device)?;
396        let mask = if seqlen_offset > 0 {
397            let mask0 = Tensor::zeros((tgt_len, seqlen_offset), DType::F32, &self.device)?;
398            Tensor::cat(&[&mask0, &mask], D::Minus1)?
399        } else {
400            mask
401        };
402        mask.expand((b_size, 1, tgt_len, tgt_len + seqlen_offset))?
403            .to_dtype(self.dtype)
404    }
405
406    pub fn forward(&mut self, input_ids: &Tensor, seqlen_offset: usize) -> Result<Tensor> {
407        let (b_size, seq_len) = input_ids.dims2()?;
408        let attention_mask = if seq_len <= 1 {
409            None
410        } else {
411            let mask = self.prepare_decoder_attention_mask(b_size, seq_len, seqlen_offset)?;
412            Some(mask)
413        };
414        let mut xs = self.embed_tokens.forward(input_ids)?;
415        for layer in self.layers.iter_mut() {
416            xs = layer.forward(&xs, attention_mask.as_ref(), seqlen_offset)?
417        }
418        xs.narrow(1, seq_len - 1, 1)?
419            .apply(&self.norm)?
420            .apply(&self.lm_head)
421    }
422
423    pub fn clear_kv_cache(&mut self) {
424        for layer in self.layers.iter_mut() {
425            layer.clear_kv_cache()
426        }
427    }
428}