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 for i in 0..config.num_hidden_layers {
24 result = Model::layer(
25 builder,
26 i,
27 config,
28 cache,
29 pos,
30 &format!("model.layers.{i}"),
31 result,
32 );
33 }
34
35 result = rmsnorm(builder, config.rms_norm_eps, "model.norm", result);
36
37 if tokens > 1 {
39 result = narrow(builder, 1, tokens - 1, 1, result);
40 }
41
42 linear_no_bias(
43 builder,
44 config.hidden_size,
45 config.vocab_size,
46 "lm_head",
47 result,
48 )
49 }
50}
51
52impl Model {
53 pub fn embeddings(builder: &Builder, config: &Config, x: Var) -> Var {
54 let t = NdArrayType::new(
55 Shape(vec![config.vocab_size, config.hidden_size]),
56 Dtype::F32,
57 );
58 let weights = parameter(builder, t, "model.embed_tokens.weight".to_string());
59 embedding(builder, x, weights)
60 }
61
62 pub fn attention(
63 builder: &Builder,
64 layer_id: usize,
65 config: &Config,
66 cache: &mut Cache,
67 pos: usize,
68 name: &str,
69 x: Var,
70 ) -> Var {
71 let dim = config.hidden_size;
72 let num_heads = config.num_attention_heads;
73 let num_kv_heads = config.num_key_value_heads;
74 let head_dim = config.hidden_size / num_heads;
75 let b = x.label.shape.0[0];
76 let s = x.label.shape.0[1];
77
78 let q = linear_no_bias(builder, dim, dim, &format!("{name}.q_proj"), x.clone());
79 let k = linear_no_bias(
80 builder,
81 dim,
82 dim * num_kv_heads / num_heads,
83 &format!("{name}.k_proj"),
84 x.clone(),
85 );
86 let v = linear_no_bias(
87 builder,
88 dim,
89 dim * num_kv_heads / num_heads,
90 &format!("{name}.v_proj"),
91 x,
92 );
93
94 let q = rmsnorm(builder, config.rms_norm_eps, &format!("{name}.q_norm"), q);
95 let k = rmsnorm(builder, config.rms_norm_eps, &format!("{name}.k_norm"), k);
96
97 let q = reshape(builder, Shape(vec![b, s, num_heads, head_dim]), q);
98 let k = reshape(builder, Shape(vec![b, s, num_kv_heads, head_dim]), k);
99 let v = reshape(builder, Shape(vec![b, s, num_kv_heads, head_dim]), v);
100
101 let q = transpose(builder, 1, 2, q);
102 let k = transpose(builder, 1, 2, k);
103 let v = transpose(builder, 1, 2, v);
104
105 let q = apply_rope_embedding(builder, pos, cache.cos.clone(), cache.sin.clone(), q);
106 let k = apply_rope_embedding(builder, pos, cache.cos.clone(), cache.sin.clone(), k);
107
108 let (k, v) = cache.update_kv_cache(builder, layer_id, k, v);
109
110 let tk = transpose(builder, 2, 3, k);
111 let attn = mat_mul(builder, q, tk);
112 let denom = constant(builder, attn.label.clone(), f32::sqrt(head_dim as f32));
113 let attn = attn / denom;
114
115 let mask = causal_mask(builder, s);
116 let mask = expand(builder, attn.label.shape.clone(), mask);
117 let attn = attn + mask;
118
119 let attn = softmax(builder, attn);
120 let attn = mat_mul(builder, attn, v);
121 let x = transpose(builder, 1, 2, attn);
122 let x = reshape(builder, Shape(vec![b, s, dim]), x);
123
124 linear_no_bias(builder, dim, dim, &format!("{name}.o_proj"), x)
125 }
126
127 pub fn mlp(builder: &Builder, config: &Config, name: &str, x: Var) -> Var {
128 let gated = linear_no_bias(
129 builder,
130 config.hidden_size,
131 config.intermediate_size,
132 &format!("{name}.gate_proj"),
133 x.clone(),
134 );
135 let up = linear_no_bias(
136 builder,
137 config.hidden_size,
138 config.intermediate_size,
139 &format!("{name}.up_proj"),
140 x,
141 );
142 let x = silu(builder, gated) * up; linear_no_bias(
145 builder,
146 config.intermediate_size,
147 config.hidden_size,
148 &format!("{name}.down_proj"),
149 x,
150 )
151 }
152
153 pub fn layer(
154 builder: &Builder,
155 layer_id: usize,
156 config: &Config,
157 cache: &mut Cache,
158 pos: usize,
159 name: &str,
160 x: Var,
161 ) -> Var {
162 let res = x.clone();
163 let x = Model::attention(
164 builder,
165 layer_id,
166 config,
167 cache,
168 pos,
169 &format!("{name}.self_attn"),
170 x,
171 );
172 let x = rmsnorm(
173 builder,
174 config.rms_norm_eps,
175 &format!("{name}.post_attention_layernorm"),
176 x,
177 );
178 let x = res + x;
179
180 let res = x.clone();
181 let x = Model::mlp(builder, config, &format!("{name}.mlp"), x);
182 let x = rmsnorm(
183 builder,
184 config.rms_norm_eps,
185 &format!("{name}.post_feedforward_layernorm"),
186 x,
187 );
188 x + res
189 }
190}