1use crate::dit::Proj;
19use crate::pool::Pool;
20use crate::qtensor::QTensor;
21use crate::vae::{StTensor, read_safetensors};
22use cortiq_core::CmfModel;
23use std::collections::HashMap;
24use std::path::Path;
25use std::sync::Arc;
26
27struct Layer {
28 input_norm: Vec<f32>,
29 q: Proj, k: Proj, v: Proj,
32 o: Proj, post_attn_norm: Vec<f32>,
34 pre_ffn_norm: Vec<f32>,
35 gate: Proj, up: Proj,
37 down: Proj, post_ffn_norm: Vec<f32>,
39}
40
41pub struct GemmaEncoder {
45 embed: QTensor, layers: Vec<Layer>,
47 final_norm: Vec<f32>,
48 pool: Option<Arc<Pool>>,
49 pub hidden: usize,
50 nh: usize,
51 nkv: usize,
52 hd: usize,
53 scale: f32, softcap: f32,
55 theta: f32,
56 eps: f64,
57 window: usize,
58}
59
60fn rms_norm_gemma(x: &[f32], w: &[f32], eps: f64) -> Vec<f32> {
61 let n = x.len() as f64;
62 let ss = x.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / n;
63 let inv = 1.0 / (ss + eps).sqrt();
64 x.iter()
65 .zip(w)
66 .map(|(&v, &g)| ((v as f64 * inv) as f32) * (1.0 + g))
67 .collect()
68}
69
70fn gelu_tanh(v: f32) -> f32 {
71 const C: f32 = 0.797_884_6; 0.5 * v * (1.0 + (C * (v + 0.044715 * v * v * v)).tanh())
73}
74
75impl GemmaEncoder {
76 pub fn load_dir(dir: &Path) -> Result<Self, String> {
77 let cfg: serde_json::Value = serde_json::from_slice(
78 &std::fs::read(dir.join("config.json")).map_err(|e| format!("config.json: {e}"))?,
79 )
80 .map_err(|e| format!("config.json: {e}"))?;
81 let idx: serde_json::Value = serde_json::from_slice(
82 &std::fs::read(dir.join("model.safetensors.index.json"))
83 .map_err(|e| format!("index: {e}"))?,
84 )
85 .map_err(|e| format!("index: {e}"))?;
86 let mut shards: Vec<String> = idx["weight_map"]
87 .as_object()
88 .ok_or("weight_map")?
89 .values()
90 .filter_map(|v| v.as_str().map(String::from))
91 .collect();
92 shards.sort();
93 shards.dedup();
94 let mut t: HashMap<String, StTensor> = HashMap::new();
95 for sh in &shards {
96 t.extend(read_safetensors(&dir.join(sh))?);
97 }
98 let take = |n: &str| -> Result<Vec<f32>, String> {
100 t.get(n)
101 .or_else(|| t.get(n.strip_prefix("model.").unwrap_or(n)))
102 .map(|v| v.data.clone())
103 .ok_or_else(|| format!("missing tensor {n}"))
104 };
105 let nl = cfg["num_hidden_layers"].as_u64().ok_or("layers")? as usize;
106 let hidden = cfg["hidden_size"].as_u64().ok_or("hidden")? as usize;
107 let mut layers = Vec::with_capacity(nl);
108 for l in 0..nl {
109 let p = format!("model.layers.{l}");
110 let o = take(&format!("{p}.self_attn.o_proj.weight"))?;
111 let o_cols = o.len() / hidden;
112 let down = take(&format!("{p}.mlp.down_proj.weight"))?;
113 let inter = down.len() / hidden;
114 layers.push(Layer {
115 input_norm: take(&format!("{p}.input_layernorm.weight"))?,
116 q: Proj::f32(take(&format!("{p}.self_attn.q_proj.weight"))?, hidden),
117 k: Proj::f32(take(&format!("{p}.self_attn.k_proj.weight"))?, hidden),
118 v: Proj::f32(take(&format!("{p}.self_attn.v_proj.weight"))?, hidden),
119 o: Proj::f32(o, o_cols),
120 post_attn_norm: take(&format!("{p}.post_attention_layernorm.weight"))?,
121 pre_ffn_norm: take(&format!("{p}.pre_feedforward_layernorm.weight"))?,
122 gate: Proj::f32(take(&format!("{p}.mlp.gate_proj.weight"))?, hidden),
123 up: Proj::f32(take(&format!("{p}.mlp.up_proj.weight"))?, hidden),
124 down: Proj::f32(down, inter),
125 post_ffn_norm: take(&format!("{p}.post_feedforward_layernorm.weight"))?,
126 });
127 }
128 let embed = take("model.embed_tokens.weight")?;
129 let vocab = embed.len() / hidden;
130 Ok(Self {
131 embed: QTensor::from_f32(embed, vocab, hidden),
132 layers,
133 final_norm: take("model.norm.weight")?,
134 pool: Pool::from_env(),
135 hidden,
136 nh: cfg["num_attention_heads"].as_u64().ok_or("nh")? as usize,
137 nkv: cfg["num_key_value_heads"].as_u64().ok_or("nkv")? as usize,
138 hd: cfg["head_dim"].as_u64().ok_or("hd")? as usize,
139 scale: 1.0 / (cfg["query_pre_attn_scalar"].as_f64().unwrap_or(256.0) as f32).sqrt(),
140 softcap: cfg["attn_logit_softcapping"].as_f64().unwrap_or(0.0) as f32,
141 theta: cfg["rope_theta"].as_f64().unwrap_or(10000.0) as f32,
142 eps: cfg["rms_norm_eps"].as_f64().unwrap_or(1e-6),
143 window: cfg["sliding_window"].as_u64().unwrap_or(4096) as usize,
144 })
145 }
146
147 pub fn from_cmf(model: &Arc<CmfModel>) -> Result<Self, String> {
151 let cfg: serde_json::Value = serde_json::from_slice(
152 model
153 .tensor_bytes("te.config_json")
154 .map_err(|e| e.to_string())?,
155 )
156 .map_err(|e| format!("te.config_json: {e}"))?;
157 let f32v = |n: &str| -> Result<Vec<f32>, String> { crate::dit::cmf_f32(model, n) };
158 let nl = cfg["num_hidden_layers"].as_u64().ok_or("layers")? as usize;
159 let mut layers = Vec::with_capacity(nl);
160 for l in 0..nl {
161 let p = format!("te.layers.{l}");
162 layers.push(Layer {
163 input_norm: f32v(&format!("{p}.input_layernorm.weight"))?,
164 q: Proj::from_model(model, &format!("{p}.self_attn.q_proj.weight"))?,
165 k: Proj::from_model(model, &format!("{p}.self_attn.k_proj.weight"))?,
166 v: Proj::from_model(model, &format!("{p}.self_attn.v_proj.weight"))?,
167 o: Proj::from_model(model, &format!("{p}.self_attn.o_proj.weight"))?,
168 post_attn_norm: f32v(&format!("{p}.post_attention_layernorm.weight"))?,
169 pre_ffn_norm: f32v(&format!("{p}.pre_feedforward_layernorm.weight"))?,
170 gate: Proj::from_model(model, &format!("{p}.mlp.gate_proj.weight"))?,
171 up: Proj::from_model(model, &format!("{p}.mlp.up_proj.weight"))?,
172 down: Proj::from_model(model, &format!("{p}.mlp.down_proj.weight"))?,
173 post_ffn_norm: f32v(&format!("{p}.post_feedforward_layernorm.weight"))?,
174 });
175 }
176 Ok(Self {
177 embed: QTensor::from_model(model, "te.embed_tokens.weight")?,
178 layers,
179 final_norm: f32v("te.norm.weight")?,
180 pool: Pool::from_env(),
181 hidden: cfg["hidden_size"].as_u64().ok_or("hidden")? as usize,
182 nh: cfg["num_attention_heads"].as_u64().ok_or("nh")? as usize,
183 nkv: cfg["num_key_value_heads"].as_u64().ok_or("nkv")? as usize,
184 hd: cfg["head_dim"].as_u64().ok_or("hd")? as usize,
185 scale: 1.0 / (cfg["query_pre_attn_scalar"].as_f64().unwrap_or(256.0) as f32).sqrt(),
186 softcap: cfg["attn_logit_softcapping"].as_f64().unwrap_or(0.0) as f32,
187 theta: cfg["rope_theta"].as_f64().unwrap_or(10000.0) as f32,
188 eps: cfg["rms_norm_eps"].as_f64().unwrap_or(1e-6),
189 window: cfg["sliding_window"].as_u64().unwrap_or(4096) as usize,
190 })
191 }
192
193 pub fn encode(&self, ids: &[u32], keep_layer_inputs: bool) -> (Vec<f32>, Vec<Vec<f32>>) {
198 let n = ids.len();
199 assert!(
200 n < self.window,
201 "prompt of {n} tokens exceeds the sliding window {}",
202 self.window
203 );
204 let hs = self.hidden;
205 let pool = self.pool.as_deref();
206 let emb_scale = (hs as f32).sqrt();
207 let mut h = vec![0f32; n * hs];
208 for (i, &id) in ids.iter().enumerate() {
209 let row = &mut h[i * hs..(i + 1) * hs];
210 self.embed.row_f32(id as usize, row);
211 for v in row.iter_mut() {
212 *v *= emb_scale;
213 }
214 }
215 let mut streams = Vec::new();
216 let (nh, nkv, hd) = (self.nh, self.nkv, self.hd);
217 let hpk = nh / nkv;
218 for layer in &self.layers {
219 if keep_layer_inputs {
220 streams.push(h.clone());
221 }
222 let mut q_all = vec![0f32; n * nh * hd];
224 let mut k_all = vec![0f32; n * nkv * hd];
225 let mut v_all = vec![0f32; n * nkv * hd];
226 let mut xn_all = vec![0f32; n * hs];
227 for p in 0..n {
228 xn_all[p * hs..(p + 1) * hs].copy_from_slice(&rms_norm_gemma(
229 &h[p * hs..(p + 1) * hs],
230 &layer.input_norm,
231 self.eps,
232 ));
233 }
234 layer.q.matmat(&xn_all, n, &mut q_all, pool);
235 layer.k.matmat(&xn_all, n, &mut k_all, pool);
236 layer.v.matmat(&xn_all, n, &mut v_all, pool);
237 for (all, heads) in [(&mut q_all, nh), (&mut k_all, nkv)] {
239 for p in 0..n {
240 for hh in 0..heads {
241 let v = &mut all[(p * heads + hh) * hd..(p * heads + hh + 1) * hd];
242 for i in 0..hd / 2 {
243 let freq = 1.0 / self.theta.powf(2.0 * i as f32 / hd as f32);
244 let (sin, cos) = (p as f32 * freq).sin_cos();
245 let (a, b) = (v[i], v[i + hd / 2]);
246 v[i] = a * cos - b * sin;
247 v[i + hd / 2] = a * sin + b * cos;
248 }
249 }
250 }
251 }
252 let mut attn_out = vec![0f32; n * nh * hd];
253 let mut row = vec![0f32; n];
254 for hh in 0..nh {
255 let kv = hh / hpk;
256 for p in 0..n {
257 let qv = &q_all[(p * nh + hh) * hd..(p * nh + hh + 1) * hd];
258 for (j, r) in row[..=p].iter_mut().enumerate() {
259 let kvv = &k_all[(j * nkv + kv) * hd..(j * nkv + kv + 1) * hd];
260 let mut d = 0f32;
261 for (a, b) in qv.iter().zip(kvv) {
262 d += a * b;
263 }
264 let mut s = d * self.scale;
265 if self.softcap > 0.0 {
266 s = self.softcap * (s / self.softcap).tanh();
267 }
268 *r = s;
269 }
270 let mx = row[..=p].iter().cloned().fold(f32::MIN, f32::max);
271 let mut den = 0f32;
272 for r in row[..=p].iter_mut() {
273 *r = (*r - mx).exp();
274 den += *r;
275 }
276 let inv = 1.0 / den;
277 let out = &mut attn_out[(p * nh + hh) * hd..(p * nh + hh + 1) * hd];
278 for (j, &rw) in row[..=p].iter().enumerate() {
279 let vv = &v_all[(j * nkv + kv) * hd..(j * nkv + kv + 1) * hd];
280 for (o, s) in out.iter_mut().zip(vv) {
281 *o += rw * inv * s;
282 }
283 }
284 }
285 }
286 let mut proj_all = vec![0f32; n * hs];
287 layer.o.matmat(&attn_out, n, &mut proj_all, pool);
288 for p in 0..n {
289 let post = rms_norm_gemma(
290 &proj_all[p * hs..(p + 1) * hs],
291 &layer.post_attn_norm,
292 self.eps,
293 );
294 for (dst, v) in h[p * hs..(p + 1) * hs].iter_mut().zip(&post) {
295 *dst += v;
296 }
297 }
298 let inter = layer.gate.rows();
300 for p in 0..n {
301 xn_all[p * hs..(p + 1) * hs].copy_from_slice(&rms_norm_gemma(
302 &h[p * hs..(p + 1) * hs],
303 &layer.pre_ffn_norm,
304 self.eps,
305 ));
306 }
307 let mut g_all = vec![0f32; n * inter];
308 let mut u_all = vec![0f32; n * inter];
309 layer.gate.matmat(&xn_all, n, &mut g_all, pool);
310 layer.up.matmat(&xn_all, n, &mut u_all, pool);
311 for (g, u) in g_all.iter_mut().zip(&u_all) {
312 *g = gelu_tanh(*g) * u;
313 }
314 let mut d_all = vec![0f32; n * hs];
315 layer.down.matmat(&g_all, n, &mut d_all, pool);
316 for p in 0..n {
317 let post =
318 rms_norm_gemma(&d_all[p * hs..(p + 1) * hs], &layer.post_ffn_norm, self.eps);
319 for (dst, v) in h[p * hs..(p + 1) * hs].iter_mut().zip(&post) {
320 *dst += v;
321 }
322 }
323 }
324 if keep_layer_inputs {
325 streams.push(h.clone());
326 }
327 let mut out = Vec::with_capacity(n * hs);
328 for p in 0..n {
329 out.extend(rms_norm_gemma(
330 &h[p * hs..(p + 1) * hs],
331 &self.final_norm,
332 self.eps,
333 ));
334 }
335 (out, streams)
336 }
337}