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, pos);
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 &format!("h.{layer_id}"),
30 result,
31 );
32 }
33
34 result = layernorm(builder, config.layer_norm_epsilon, "ln_f", result);
35
36 if tokens > 1 {
38 result = narrow(builder, 1, tokens - 1, 1, result);
39 }
40
41 linear_no_bias(
43 builder,
44 config.hidden_size,
45 config.vocab_size,
46 "wte",
47 result,
48 )
49 }
50}
51
52impl Model {
53 fn gpt_linear(builder: &Builder, in_dim: usize, out_dim: usize, name: &str, x: Var) -> Var {
56 let w_type = NdArrayType::new(Shape(vec![in_dim, out_dim]), x.label.dtype);
57 let b_type = NdArrayType::new(Shape(vec![out_dim]), x.label.dtype);
58
59 let w = parameter(builder, w_type, format!("{name}.weight"));
60 let b = parameter(builder, b_type, format!("{name}.bias"));
61
62 let mut w_t = w;
64
65 if x.label.shape.0.len() == 3 {
66 let batch_size = x.label.shape.0[0];
67 w_t = expand(builder, Shape(vec![batch_size, in_dim, out_dim]), w_t);
68 }
69
70 let m = mat_mul(builder, x, w_t);
71 let bb = expand(builder, m.label.shape.clone(), b);
72 m + bb
73 }
74
75 pub fn embeddings(builder: &Builder, config: &Config, x: Var, pos: usize) -> Var {
76 let t = NdArrayType::new(
77 Shape(vec![config.vocab_size, config.hidden_size]),
78 Dtype::F32,
79 );
80 let weights = parameter(builder, t, "wte.weight".to_string());
81 let we = embedding(builder, x.clone(), weights);
82
83 let t = NdArrayType::new(
84 Shape(vec![config.max_position_embeddings, config.hidden_size]),
85 Dtype::F32,
86 );
87 let pos = range_indices(builder, pos, pos + x.label.size());
88 let pos = expand(builder, x.label.shape, pos);
89 let weights = parameter(builder, t, "wpe.weight".to_string());
90 let pe = embedding(builder, pos, weights);
91
92 we + pe
93 }
94
95 pub fn attention(
96 builder: &Builder,
97 layer_id: usize,
98 config: &Config,
99 cache: &mut Cache,
100 name: &str,
101 x: Var,
102 ) -> Var {
103 let dim = config.hidden_size;
104 let num_heads = config.num_attention_heads;
105 let head_dim = dim / num_heads;
106
107 let b = x.label.shape.0[0];
108 let s = x.label.shape.0[1];
109
110 let c_attn = Model::gpt_linear(builder, dim, 3 * dim, &format!("{name}.c_attn"), x);
111
112 let a = split(builder, 2, 3, c_attn);
113 let q = a[0].clone();
114 let k = a[1].clone();
115 let v = a[2].clone();
116
117 let q = reshape(builder, Shape(vec![b, s, num_heads, head_dim]), q);
118 let k = reshape(builder, Shape(vec![b, s, num_heads, head_dim]), k);
119 let v = reshape(builder, Shape(vec![b, s, num_heads, head_dim]), v);
120
121 let q = transpose(builder, 1, 2, q);
122 let k = transpose(builder, 1, 2, k);
123 let v = transpose(builder, 1, 2, v);
124
125 let (k, v) = cache.update_kv_cache(builder, layer_id, k, v);
126
127 let tk = transpose(builder, 2, 3, k);
128 let attn = mat_mul(builder, q, tk);
129 let denom = constant(builder, attn.label.clone(), f32::sqrt(head_dim as f32));
130 let mut attn = attn / denom;
131
132 if s > 1 {
133 let mask = causal_mask(builder, s);
134 let mask = expand(builder, attn.label.shape.clone(), mask);
135 attn = attn + mask;
136 }
137
138 let attn = softmax(builder, attn);
139 let attn = mat_mul(builder, attn, v);
140
141 let attn = transpose(builder, 1, 2, attn);
142 let attn = reshape(builder, Shape(vec![b, s, dim]), attn);
143
144 Model::gpt_linear(builder, dim, dim, &format!("{name}.c_proj"), attn)
145 }
146
147 pub fn mlp(builder: &Builder, dim: usize, name: &str, x: Var) -> Var {
148 let x = Model::gpt_linear(builder, dim, dim * 4, &format!("{name}.c_fc"), x);
149 let x = gelu(builder, x);
150
151 Model::gpt_linear(builder, dim * 4, dim, &format!("{name}.c_proj"), x)
152 }
153
154 pub fn layer(
155 builder: &Builder,
156 layer_id: usize,
157 config: &Config,
158 cache: &mut Cache,
159 name: &str,
160 x: Var,
161 ) -> Var {
162 let res = x.clone();
163 let x = layernorm(
164 builder,
165 config.layer_norm_epsilon,
166 &format!("{name}.ln_1"),
167 x,
168 );
169 let x = Model::attention(builder, layer_id, config, cache, &format!("{name}.attn"), x);
170 let x = res + x;
171 let res = x.clone();
172 let x = layernorm(
173 builder,
174 config.layer_norm_epsilon,
175 &format!("{name}.ln_2"),
176 x,
177 );
178 let x = Model::mlp(builder, config.hidden_size, &format!("{name}.mlp"), x);
179 x + res
180 }
181}