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 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 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; 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}