Skip to main content

cortiq_engine/
textenc.rs

1//! Gemma-2 text ENCODER forward (the Lumina-Image 2.0 prompt encoder):
2//! token ids → per-token hidden states of every layer.
3//!
4//! Second increment of the image-generation runtime
5//! (docs/GENERATIVE.ru.md). A deliberately standalone, f32,
6//! full-sequence forward — no KV cache, no sampling — loaded from a
7//! diffusers/HF `text_encoder/` directory (config.json + sharded
8//! safetensors). Gemma-2 specifics carried exactly: embedding scale
9//! √hidden, RMSNorm with (1 + w), sandwich norms around attention AND
10//! the GeGLU MLP, GQA with `query_pre_attn_scalar` scaling, attention
11//! logit softcapping tanh(x/50)·50, RoPE θ=10000. Sliding-window
12//! layers are exact for prompts shorter than the 4096 window (image
13//! prompts are), enforced by an assert.
14//!
15//! Parity: `python/gemma_ref.py` + `tests/textenc_parity.rs` on the
16//! real Lumina text-encoder weights.
17
18use 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, // [nh·hd, hidden]
30    k: Proj, // [nkv·hd, hidden]
31    v: Proj,
32    o: Proj, // [hidden, nh·hd]
33    post_attn_norm: Vec<f32>,
34    pre_ffn_norm: Vec<f32>,
35    gate: Proj, // [inter, hidden]
36    up: Proj,
37    down: Proj, // [hidden, inter]
38    post_ffn_norm: Vec<f32>,
39}
40
41/// Gemma-2 encoder: exact f32 from a diffusers directory, or
42/// CMF-quantized (mmap-resident, per-token embed dequant) from a
43/// packaged file.
44pub struct GemmaEncoder {
45    embed: QTensor, // [vocab, hidden]
46    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, // 1/√query_pre_attn_scalar
54    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; // √(2/π)
72    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        // Some exports (the Lumina one) drop the "model." prefix.
99        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    /// Load from a packaged imagegen .cmf (`te.*` tensors +
148    /// `te.config_json`). Quantized projections stay mmap-resident;
149    /// embeddings dequantize per token.
150    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    /// Full-sequence causal forward. Returns the FINAL-normed hidden
194    /// states `[n, hidden]` and, if `keep_layer_inputs`, the residual
195    /// stream entering every layer (what `hidden_states[i]` means in HF
196    /// — index nl equals the pre-final-norm stream).
197    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            // ── attention (pre-norm, sandwich post-norm) ──
223            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            // RoPE over the first hd dims of every head (full-dim rope).
238            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            // ── GeGLU MLP (pre-norm, sandwich post-norm) ──
299            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}