Skip to main content

catgrad_llm/models/
utils.rs

1use catgrad::{
2    backend::cpu::{eval::Builder, ndarray::TaggedNdArray},
3    core::{
4        Dtype, NdArrayType, Shape, Var,
5        nn::layers::{concat, rope_tables},
6    },
7};
8
9use std::collections::HashMap;
10
11#[derive(Debug, Clone, serde::Deserialize)]
12#[serde(untagged)]
13pub enum EosTokenId {
14    Single(i32),
15    Multiple(Vec<i32>),
16}
17
18// This configuration contains the union of relevant fields from all supported models.
19// Models ignore fields they don't need. The aliases are for GPT-2 alternative names.
20#[derive(Debug, Clone, Default, serde::Deserialize)]
21#[serde(default)]
22pub struct Config {
23    #[serde(alias = "n_embd")]
24    pub hidden_size: usize,
25    pub intermediate_size: usize,
26    #[serde(alias = "n_layer")]
27    pub num_hidden_layers: usize,
28    #[serde(alias = "n_head")]
29    pub num_attention_heads: usize,
30    pub num_key_value_heads: usize,
31    pub head_dim: usize,
32    pub rope_theta: f32,
33    pub sliding_window_pattern: usize,
34    pub rope_local_base_freq: f32,
35    #[serde(alias = "n_positions")]
36    pub max_position_embeddings: usize,
37    pub no_rope_layer_interval: usize,
38    pub layer_norm_epsilon: f32,
39    pub rms_norm_eps: f32,
40    pub tie_word_embeddings: bool,
41    pub eos_token_id: Option<EosTokenId>,
42    pub vocab_size: usize,
43    pub architectures: Vec<String>,
44}
45
46impl Config {
47    // Sometimes the head_dim fields is missing
48    pub fn get_head_dim(&self) -> usize {
49        if self.head_dim == 0 {
50            self.hidden_size / self.num_attention_heads
51        } else {
52            self.head_dim
53        }
54    }
55
56    pub fn get_num_kv_heads(&self) -> usize {
57        if self.num_key_value_heads == 0 {
58            self.num_attention_heads
59        } else {
60            self.num_key_value_heads
61        }
62    }
63
64    pub fn get_eos_token_ids(&self) -> Vec<i32> {
65        match self.eos_token_id {
66            Some(EosTokenId::Single(id)) => vec![id],
67            Some(EosTokenId::Multiple(ref ids)) => ids.clone(),
68            None => vec![],
69        }
70    }
71}
72
73pub struct Cache {
74    pub cos: Var,
75    pub sin: Var,
76    pub use_kv_cache: bool,
77    pub in_kv_cache: Vec<(Var, Var)>,
78    pub out_kv_cache: Vec<(Var, Var)>,
79}
80
81impl Cache {
82    pub fn init(builder: &Builder, config: &Config, positions: usize, use_kv_cache: bool) -> Self {
83        let (cos, sin) = rope_tables(builder, config.rope_theta, positions, config.get_head_dim());
84
85        // Empty KV Cache of the correct shape
86        let kv_cache_type = NdArrayType::new(
87            Shape(vec![1, config.get_num_kv_heads(), 0, config.get_head_dim()]),
88            Dtype::F32,
89        );
90        let empty = Var::new(builder.clone(), kv_cache_type);
91        Self {
92            cos,
93            sin,
94            use_kv_cache,
95            in_kv_cache: vec![(empty.clone(), empty.clone()); config.num_hidden_layers],
96            out_kv_cache: vec![(empty.clone(), empty); config.num_hidden_layers],
97        }
98    }
99
100    pub fn update_kv_cache(
101        &mut self,
102        builder: &Builder,
103        layer_id: usize,
104        k: Var,
105        v: Var,
106    ) -> (Var, Var) {
107        let (mut k, mut v) = (k, v);
108        if self.use_kv_cache {
109            let cached_k = self.in_kv_cache[layer_id].0.clone();
110            let cached_v = self.in_kv_cache[layer_id].1.clone();
111
112            k = concat(builder, 2, cached_k, k);
113            v = concat(builder, 2, cached_v, v);
114
115            self.out_kv_cache[layer_id] = (k.clone(), v.clone());
116        }
117        (k, v)
118    }
119}
120
121// Trait for model builders for various architectures (llama, qwen, gpt2, etc.)
122pub trait ModelBuilder {
123    // Build the model architecture graph for a given input shape
124    fn build(
125        &self,
126        builder: &Builder,
127        config: &Config,
128        cache: &mut Cache,
129        pos: usize,
130        x: Var,
131    ) -> Var;
132    // Optional post-processing of loaded weights (renaming, reshaping, etc.)
133    fn post_load(&mut self, _tensors: &mut HashMap<String, TaggedNdArray>) {}
134}