Skip to main content

catgrad_llm/models/
olmo.rs

1// OLMo-2 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        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        // Get the logits for the last token only
38        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; // SwiGLU
143
144        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}