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#[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 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 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
121pub trait ModelBuilder {
123 fn build(
125 &self,
126 builder: &Builder,
127 config: &Config,
128 cache: &mut Cache,
129 pos: usize,
130 x: Var,
131 ) -> Var;
132 fn post_load(&mut self, _tensors: &mut HashMap<String, TaggedNdArray>) {}
134}