1use 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 let normalizer = constant(
24 builder,
25 result.label.clone(),
26 (config.hidden_size as f32).sqrt(),
27 );
28
29 result = result * normalizer;
30
31 for i in 0..config.num_hidden_layers {
32 result = Model::layer(
33 builder,
34 i,
35 config,
36 cache,
37 pos,
38 &format!("model.layers.{i}"),
39 result,
40 );
41 }
42
43 result = Model::rmsnorm(builder, config.rms_norm_eps, "model.norm", result);
44
45 if tokens > 1 {
47 result = narrow(builder, 1, tokens - 1, 1, result);
48 }
49
50 linear_no_bias(
52 builder,
53 config.hidden_size,
54 config.vocab_size,
55 "model.embed_tokens",
56 result,
57 )
58 }
59}
60
61impl Model {
62 fn rmsnorm(builder: &Builder, eps: f32, name: &str, x: Var) -> Var {
64 let shape = vec![x.label.shape.0[x.label.shape.0.len() - 1]];
65 let t = NdArrayType::new(Shape(shape), x.label.dtype);
66 let gamma = parameter(builder, t, format!("{name}.weight"));
67 let lr = rmsnorm_raw(builder, eps, x);
68 let gamma = expand(builder, lr.label.shape.clone(), gamma);
69 lr * increment(builder, gamma)
71 }
72
73 pub fn embeddings(builder: &Builder, config: &Config, x: Var) -> Var {
74 let t = NdArrayType::new(
75 Shape(vec![config.vocab_size, config.hidden_size]),
76 Dtype::F32,
77 );
78 let weights = parameter(builder, t, "model.embed_tokens.weight".to_string());
79 embedding(builder, x, weights)
80 }
81
82 pub fn attention(
83 builder: &Builder,
84 layer_id: usize,
85 config: &Config,
86 cache: &mut Cache,
87 pos: usize,
88 name: &str,
89 x: Var,
90 ) -> Var {
91 let dim = config.hidden_size;
92 let num_heads = config.num_attention_heads;
93 let num_kv_heads = config.num_key_value_heads;
94 let rep = num_heads / num_kv_heads;
95 let head_dim = config.head_dim;
96 let b = x.label.shape.0[0];
97 let s = x.label.shape.0[1];
98
99 let q = linear_no_bias(
100 builder,
101 dim,
102 num_heads * head_dim,
103 &format!("{name}.q_proj"),
104 x.clone(),
105 );
106 let k = linear_no_bias(
107 builder,
108 dim,
109 num_kv_heads * head_dim,
110 &format!("{name}.k_proj"),
111 x.clone(),
112 );
113 let v = linear_no_bias(
114 builder,
115 dim,
116 num_kv_heads * head_dim,
117 &format!("{name}.v_proj"),
118 x,
119 );
120
121 let q = reshape(builder, Shape(vec![b, s, num_heads, head_dim]), q);
122 let k = reshape(builder, Shape(vec![b, s, num_kv_heads, head_dim]), k);
123 let v = reshape(builder, Shape(vec![b, s, num_kv_heads, head_dim]), v);
124
125 let q = transpose(builder, 1, 2, q);
126 let k = transpose(builder, 1, 2, k);
127 let v = transpose(builder, 1, 2, v);
128
129 let q = reshape(builder, Shape(vec![b * s * num_heads, head_dim]), q);
131 let k = reshape(builder, Shape(vec![b * s * num_kv_heads, head_dim]), k);
132 let q = Model::rmsnorm(builder, config.rms_norm_eps, &format!("{name}.q_norm"), q);
133 let k = Model::rmsnorm(builder, config.rms_norm_eps, &format!("{name}.k_norm"), k);
134 let q = reshape(builder, Shape(vec![b, num_heads, s, head_dim]), q);
135 let k = reshape(builder, Shape(vec![b, num_kv_heads, s, head_dim]), k);
136
137 let theta = if (layer_id + 1) % config.sliding_window_pattern > 0 {
140 config.rope_local_base_freq
141 } else {
142 config.rope_theta
143 };
144 let q = rope(builder, theta, pos, s, q);
145 let k = rope(builder, theta, pos, s, k);
146
147 let (k, v) = cache.update_kv_cache(builder, layer_id, k, v);
148
149 let k = repeat_kv(builder, rep, k);
150 let v = repeat_kv(builder, rep, v);
151
152 let tk = transpose(builder, 2, 3, k);
153 let attn = mat_mul(builder, q, tk);
154 let denom = constant(builder, attn.label.clone(), f32::sqrt(head_dim as f32));
155 let attn = attn / denom;
156
157 let mask = causal_mask(builder, s);
158 let mask = expand(builder, attn.label.shape.clone(), mask);
159 let attn = attn + mask;
160
161 let attn = softmax(builder, attn);
162 let attn = mat_mul(builder, attn, v);
163 let x = transpose(builder, 1, 2, attn);
164 let x = reshape(builder, Shape(vec![b, s, num_heads * head_dim]), x);
165
166 linear_no_bias(
167 builder,
168 num_heads * head_dim,
169 dim,
170 &format!("{name}.o_proj"),
171 x,
172 )
173 }
174
175 pub fn mlp(builder: &Builder, config: &Config, name: &str, x: Var) -> Var {
176 let gated = linear_no_bias(
177 builder,
178 config.hidden_size,
179 config.intermediate_size,
180 &format!("{name}.gate_proj"),
181 x.clone(),
182 );
183 let up = linear_no_bias(
184 builder,
185 config.hidden_size,
186 config.intermediate_size,
187 &format!("{name}.up_proj"),
188 x,
189 );
190 let x = gelu(builder, gated) * up;
191
192 linear_no_bias(
193 builder,
194 config.intermediate_size,
195 config.hidden_size,
196 &format!("{name}.down_proj"),
197 x,
198 )
199 }
200
201 pub fn layer(
202 builder: &Builder,
203 layer_id: usize,
204 config: &Config,
205 cache: &mut Cache,
206 pos: usize,
207 name: &str,
208 x: Var,
209 ) -> Var {
210 let res = x.clone();
211 let x = Model::rmsnorm(
212 builder,
213 config.rms_norm_eps,
214 &format!("{name}.input_layernorm"),
215 x,
216 );
217 let x = Model::attention(
218 builder,
219 layer_id,
220 config,
221 cache,
222 pos,
223 &format!("{name}.self_attn"),
224 x,
225 );
226 let x = Model::rmsnorm(
227 builder,
228 config.rms_norm_eps,
229 &format!("{name}.post_attention_layernorm"),
230 x,
231 );
232 let x = res + x;
233 let res = x.clone();
234 let x = Model::rmsnorm(
235 builder,
236 config.rms_norm_eps,
237 &format!("{name}.pre_feedforward_layernorm"),
238 x,
239 );
240 let x = Model::mlp(builder, config, &format!("{name}.mlp"), x);
241 let x = Model::rmsnorm(
242 builder,
243 config.rms_norm_eps,
244 &format!("{name}.post_feedforward_layernorm"),
245 x,
246 );
247 x + res
248 }
249}