Skip to main content

catgrad_llm/models/
gemma.rs

1// Gemma-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        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        // Get the logits for the last token only
46        if tokens > 1 {
47            result = narrow(builder, 1, tokens - 1, 1, result);
48        }
49
50        // Gemma uses weight tying so lm_head is the same as embed_tokens
51        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    // Gemma uses a non-standard RMSNorm
63    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        // this is different for Gemma, standard RMSNorm multiplies by gamma
70        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        // Norm
130        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        // Rope
138        // Every 6th layer uses global attention, otherwise local attention
139        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}