Skip to main content

catgrad_llm/models/
phi.rs

1// Phi-3 model description
2
3use super::utils::{Cache, Config, ModelBuilder};
4use catgrad::backend::cpu::eval::Builder;
5use catgrad::core::nn::layers::*;
6use catgrad::core::{Dtype, NdArrayType, Shape, Var};
7
8pub struct Model;
9
10impl ModelBuilder for Model {
11    fn build(
12        &self,
13        builder: &Builder,
14        config: &Config,
15        cache: &mut Cache,
16        pos: usize,
17        x: Var,
18    ) -> Var {
19        let tokens = x.label.shape.0[1];
20        let emb = Model::embeddings(builder, config, x);
21        let mut result = emb;
22
23        for layer_id in 0..config.num_hidden_layers {
24            result = Model::layer(
25                builder,
26                layer_id,
27                config,
28                cache,
29                pos,
30                &format!("model.layers.{layer_id}"),
31                result,
32            );
33        }
34
35        // Get the logits for the last token only
36        if tokens > 1 {
37            result = narrow(builder, 1, tokens - 1, 1, result);
38        }
39
40        result = rmsnorm(builder, config.rms_norm_eps, "model.norm", result);
41
42        let lm_head_weights = if config.tie_word_embeddings {
43            "model.embed_tokens"
44        } else {
45            "lm_head"
46        };
47        linear_no_bias(
48            builder,
49            config.hidden_size,
50            config.vocab_size,
51            lm_head_weights,
52            result,
53        )
54    }
55}
56
57impl Model {
58    pub fn embeddings(builder: &Builder, config: &Config, x: Var) -> Var {
59        let t = NdArrayType::new(
60            Shape(vec![config.vocab_size, config.hidden_size]),
61            Dtype::F32,
62        );
63        let weights = parameter(builder, t, "model.embed_tokens.weight".to_string());
64        embedding(builder, x, weights)
65    }
66
67    pub fn attention(
68        builder: &Builder,
69        layer_id: usize,
70        config: &Config,
71        cache: &mut Cache,
72        pos: usize,
73        name: &str,
74        x: Var,
75    ) -> Var {
76        let dim = config.hidden_size;
77        let num_heads = config.num_attention_heads;
78        let num_kv_heads = config.num_key_value_heads;
79        let rep = num_heads / num_kv_heads;
80        let head_dim = config.hidden_size / num_heads;
81        let b = x.label.shape.0[0];
82        let s = x.label.shape.0[1];
83
84        let qkv = linear_no_bias(builder, dim, 3 * dim, &format!("{name}.qkv_proj"), x);
85
86        let q = narrow(builder, 2, 0, num_heads * head_dim, qkv.clone());
87        let k = narrow(
88            builder,
89            2,
90            num_heads * head_dim,
91            num_kv_heads * head_dim,
92            qkv.clone(),
93        );
94        let v = narrow(
95            builder,
96            2,
97            (num_heads + num_kv_heads) * head_dim,
98            num_kv_heads * head_dim,
99            qkv,
100        );
101
102        let q = reshape(builder, Shape(vec![b, s, num_heads, head_dim]), q);
103        let k = reshape(builder, Shape(vec![b, s, num_kv_heads, head_dim]), k);
104        let v = reshape(builder, Shape(vec![b, s, num_kv_heads, head_dim]), v);
105
106        let q = transpose(builder, 1, 2, q);
107        let k = transpose(builder, 1, 2, k);
108        let v = transpose(builder, 1, 2, v);
109
110        let q = apply_rope_embedding(builder, pos, cache.cos.clone(), cache.sin.clone(), q);
111        let k = apply_rope_embedding(builder, pos, cache.cos.clone(), cache.sin.clone(), k);
112
113        let (k, v) = cache.update_kv_cache(builder, layer_id, k, v);
114
115        let k = repeat_kv(builder, rep, k);
116        let v = repeat_kv(builder, rep, v);
117
118        let tk = transpose(builder, 2, 3, k);
119        let attn = mat_mul(builder, q, tk);
120        let denom = constant(builder, attn.label.clone(), f32::sqrt(head_dim as f32));
121        let attn = attn / denom;
122
123        let mask = causal_mask(builder, s);
124        let mask = expand(builder, attn.label.shape.clone(), mask);
125        let attn = attn + mask;
126
127        let attn = softmax(builder, attn);
128        let attn = mat_mul(builder, attn, v);
129        let x = transpose(builder, 1, 2, attn);
130        let x = reshape(builder, Shape(vec![b, s, dim]), x);
131
132        linear_no_bias(builder, dim, dim, &format!("{name}.o_proj"), x)
133    }
134
135    pub fn mlp(builder: &Builder, config: &Config, name: &str, x: Var) -> Var {
136        let gate_up = linear_no_bias(
137            builder,
138            config.hidden_size,
139            config.intermediate_size * 2,
140            &format!("{name}.gate_up_proj"),
141            x,
142        );
143
144        let gate_up = split(builder, 2, 2, gate_up);
145
146        let gate = gate_up[0].clone();
147        let up = gate_up[1].clone();
148
149        let x = silu(builder, gate) * up; // SwiGLU
150
151        linear_no_bias(
152            builder,
153            config.intermediate_size,
154            config.hidden_size,
155            &format!("{name}.down_proj"),
156            x,
157        )
158    }
159
160    pub fn layer(
161        builder: &Builder,
162        layer_id: usize,
163        config: &Config,
164        cache: &mut Cache,
165        pos: usize,
166        name: &str,
167        x: Var,
168    ) -> Var {
169        let res = x.clone();
170        let x = rmsnorm(
171            builder,
172            config.rms_norm_eps,
173            &format!("{name}.input_layernorm"),
174            x,
175        );
176        let x = Model::attention(
177            builder,
178            layer_id,
179            config,
180            cache,
181            pos,
182            &format!("{name}.self_attn"),
183            x,
184        );
185        let x = res + x;
186        let res = x.clone();
187        let x = rmsnorm(
188            builder,
189            config.rms_norm_eps,
190            &format!("{name}.post_attention_layernorm"),
191            x,
192        );
193        let x = Model::mlp(builder, config, &format!("{name}.mlp"), x);
194        x + res
195    }
196}