Skip to main content

ferrox_models/
gemma4_engine.rs

1//! Gemma-4 dedicated text engine (E2B / E4B-style GGUF).
2//!
3//! Distinct from generic Gemma-2/3 [`crate::decoder::Decoder`]:
4//! - per-layer token embeddings (`per_layer_token_embd` + proj)
5//! - shared KV layers (`attention.shared_kv_layers`)
6//! - SWA vs full layers with different head dims (`key_length` /
7//!   `key_length_swa`) and a bool SWA pattern array
8//!
9//! Graph mirrors `.scratch/llama.cpp/src/models/gemma4.cpp`.
10
11use ferrox_core::attention::{
12    apply_rope, apply_rope_with_freq_factors, causal_gqa_attention, causal_gqa_attention_windowed,
13};
14use ferrox_core::cache::KvCache;
15use ferrox_core::matmul::{geglu, gelu, rms_norm, rms_norm_per_head, softcap_inplace};
16use ferrox_core::weight_matrix::WeightMatrix;
17
18use crate::engine::Engine;
19
20/// Architectures served by this engine.
21pub const GEMMA4_ARCHES: &[&str] = &["gemma4", "gemma4-assistant"];
22
23/// Hyperparameters from `{arch}.*` GGUF metadata.
24#[derive(Debug, Clone)]
25pub struct Gemma4Hparams {
26    pub arch: String,
27    pub n_layer: usize,
28    pub hidden_dim: usize,
29    /// Per-layer FFN intermediate sizes (`feed_forward_length` array or scalar).
30    pub ffn_dims: Vec<usize>,
31    pub n_heads: usize,
32    pub n_kv_heads: usize,
33    pub head_dim_full: usize,
34    pub head_dim_swa: usize,
35    pub sliding_window: usize,
36    /// `true` = SWA layer; length == `n_layer`.
37    pub is_swa: Vec<bool>,
38    /// First this many layers own a KV cache; later layers reuse.
39    pub n_layer_kv_from_start: usize,
40    pub embd_per_layer: usize,
41    pub rms_norm_eps: f32,
42    pub rope_theta: f32,
43    pub rope_theta_swa: f32,
44    pub final_logit_softcap: Option<f32>,
45    /// Attention score scale passed to the kernel (Gemma-4: 1.0).
46    pub attention_scale: f32,
47}
48
49impl Gemma4Hparams {
50    pub fn is_swa_layer(&self, il: usize) -> bool {
51        self.is_swa.get(il).copied().unwrap_or(false)
52    }
53
54    pub fn head_dim(&self, il: usize) -> usize {
55        if self.is_swa_layer(il) {
56            self.head_dim_swa
57        } else {
58            self.head_dim_full
59        }
60    }
61
62    pub fn has_kv(&self, il: usize) -> bool {
63        il < self.n_layer_kv_from_start
64    }
65
66    /// Layer whose KV cache is reused when `!has_kv(il)` (llama.cpp GEMMA4 reuse).
67    pub fn kv_reuse_layer(&self, il: usize) -> usize {
68        debug_assert!(!self.has_kv(il));
69        self.n_layer_kv_from_start - if self.is_swa_layer(il) { 2 } else { 1 }
70    }
71
72    pub fn rope_theta_for(&self, il: usize) -> f32 {
73        if self.is_swa_layer(il) {
74            self.rope_theta_swa
75        } else {
76            self.rope_theta
77        }
78    }
79
80    pub fn ffn_dim(&self, il: usize) -> usize {
81        self.ffn_dims
82            .get(il)
83            .copied()
84            .unwrap_or_else(|| *self.ffn_dims.last().unwrap_or(&self.hidden_dim))
85    }
86}
87
88pub struct Gemma4AttnWeights {
89    pub q_proj: WeightMatrix,
90    pub k_proj: Option<WeightMatrix>,
91    pub v_proj: Option<WeightMatrix>,
92    pub o_proj: WeightMatrix,
93    pub q_norm: Vec<f32>,
94    pub k_norm: Option<Vec<f32>>,
95    pub post_attn_norm: Vec<f32>,
96}
97
98pub struct Gemma4LayerWeights {
99    pub attn_norm: Vec<f32>,
100    pub attn: Gemma4AttnWeights,
101    pub ffn_norm: Vec<f32>,
102    pub ffn_gate: WeightMatrix,
103    pub ffn_up: WeightMatrix,
104    pub ffn_down: WeightMatrix,
105    pub ffn_post_norm: Vec<f32>,
106    pub per_layer_inp_gate: Option<WeightMatrix>,
107    pub per_layer_proj: Option<WeightMatrix>,
108    pub per_layer_post_norm: Option<Vec<f32>>,
109    pub out_scale: Option<f32>,
110}
111
112pub struct Gemma4Weights {
113    pub token_embd: WeightMatrix,
114    pub per_layer_token_embd: Option<WeightMatrix>,
115    pub per_layer_model_proj: Option<WeightMatrix>,
116    pub per_layer_proj_norm: Option<Vec<f32>>,
117    pub layers: Vec<Gemma4LayerWeights>,
118    pub output_norm: Vec<f32>,
119    pub output_head: WeightMatrix,
120    /// Full-attn RoPE frequency factors (`rope_freqs.weight`), length `head_dim_full/2`.
121    pub rope_freqs: Option<Vec<f32>>,
122}
123
124pub struct Gemma4Engine {
125    pub weights: Gemma4Weights,
126    pub hp: Gemma4Hparams,
127}
128
129pub struct Gemma4DecodeState {
130    /// One cache per layer that `has_kv`; indexed by layer id (holes for shared layers).
131    pub kv: Vec<Option<KvCache>>,
132}
133
134impl Gemma4Engine {
135    /// See [`crate::decoder::Decoder::probe_kernels`]. Called at the end
136    /// of loading, before the registry is sealed.
137    ///
138    /// This engine implements [`Engine::forward_token`] and nothing
139    /// else: there is no batched path, so a 512-token prefill is 512
140    /// sequential single-token passes and every projection is a
141    /// `batch = 1` matvec. On Metal that is one command buffer, one
142    /// commit and one wait per projection per token, which is why
143    /// Gemma-4-E2B measures *slower* on Metal than on CPU while
144    /// producing correct output. The registry records that as a missing
145    /// engine capability rather than leaving it to be rediscovered from
146    /// a benchmark.
147    pub fn probe_kernels(&self) {
148        use ferrox_core::kernel_registry as reg;
149
150        if !reg::enabled() {
151            return;
152        }
153        let w = &self.weights;
154        w.token_embd.probe_kernels("token_embd");
155        w.output_head.probe_kernels("output_head");
156        if let Some(m) = &w.per_layer_token_embd {
157            m.probe_kernels("per_layer_token_embd");
158        }
159        if let Some(m) = &w.per_layer_model_proj {
160            m.probe_kernels("per_layer_model_proj");
161        }
162        for layer in &w.layers {
163            layer.attn.q_proj.probe_kernels("attn_q");
164            if let Some(m) = &layer.attn.k_proj {
165                m.probe_kernels("attn_k");
166            }
167            if let Some(m) = &layer.attn.v_proj {
168                m.probe_kernels("attn_v");
169            }
170            layer.attn.o_proj.probe_kernels("attn_o");
171            layer.ffn_gate.probe_kernels("ffn_gate");
172            layer.ffn_up.probe_kernels("ffn_up");
173            layer.ffn_down.probe_kernels("ffn_down");
174            if let Some(m) = &layer.per_layer_inp_gate {
175                m.probe_kernels("per_layer_inp_gate");
176            }
177            if let Some(m) = &layer.per_layer_proj {
178                m.probe_kernels("per_layer_proj");
179            }
180        }
181        reg::record_build(
182            reg::Lookup::new(
183                ferrox_core::weight_matrix::active_backend(),
184                reg::op::ENGINE_PREFILL_BATCH,
185                None,
186            )
187            .with_role("gemma4_engine"),
188            reg::Outcome::slow_path("sequential forward_token (batch=1 matvec per projection)"),
189        );
190    }
191
192    pub fn new_state(&self) -> Gemma4DecodeState {
193        let kv = (0..self.hp.n_layer)
194            .map(|il| {
195                if self.hp.has_kv(il) {
196                    let hd = self.hp.head_dim(il);
197                    Some(KvCache::new(self.hp.n_kv_heads, hd))
198                } else {
199                    None
200                }
201            })
202            .collect();
203        Gemma4DecodeState { kv }
204    }
205
206    fn project_per_layer_inputs(&self, token_id: usize, hidden: &[f32]) -> Option<Vec<Vec<f32>>> {
207        let pl_embd = self.weights.per_layer_token_embd.as_ref()?;
208        let pl_proj = self.weights.per_layer_model_proj.as_ref()?;
209        let pl_norm = self.weights.per_layer_proj_norm.as_ref()?;
210        let n = self.hp.embd_per_layer;
211        let n_layer = self.hp.n_layer;
212        let scale_tok = (n as f32).sqrt();
213        let mut per_tok = pl_embd.dequant_row(token_id);
214        for x in per_tok.iter_mut() {
215            *x *= scale_tok;
216        }
217        // per_tok: [n_layer * n] contiguous as layer-major chunks of n
218        let proj_scale = 1.0 / (self.hp.hidden_dim as f32).sqrt();
219        let mut from_model = pl_proj.apply(hidden);
220        for x in from_model.iter_mut() {
221            *x *= proj_scale;
222        }
223        // from_model: [n_layer * n]
224        let mut out = Vec::with_capacity(n_layer);
225        let input_scale = 1.0 / 2f32.sqrt();
226        for il in 0..n_layer {
227            let start = il * n;
228            let mut chunk: Vec<f32> = from_model[start..start + n].to_vec();
229            chunk = rms_norm(&chunk, pl_norm, self.hp.rms_norm_eps);
230            for (c, t) in chunk.iter_mut().zip(per_tok[start..start + n].iter()) {
231                *c = (*c + *t) * input_scale;
232            }
233            out.push(chunk);
234        }
235        Some(out)
236    }
237}
238
239impl Engine for Gemma4Engine {
240    type State = Gemma4DecodeState;
241
242    fn new_state(&self) -> Gemma4DecodeState {
243        Gemma4Engine::new_state(self)
244    }
245
246    fn vocab_size(&self) -> usize {
247        self.weights.output_head.rows()
248    }
249
250    fn forward_token_on_worker(
251        &self,
252        token_id: usize,
253        pos: usize,
254        state: &mut Self::State,
255    ) -> Vec<f32> {
256        let hp = &self.hp;
257        let mut hidden = self.weights.token_embd.dequant_row(token_id);
258        let emb_scale = (hp.hidden_dim as f32).sqrt();
259        for x in hidden.iter_mut() {
260            *x *= emb_scale;
261        }
262
263        let per_layer_in = self.project_per_layer_inputs(token_id, &hidden);
264
265        for (il, layer) in self.weights.layers.iter().enumerate() {
266            let head_dim = hp.head_dim(il);
267            let n_heads = hp.n_heads;
268            let n_kv = hp.n_kv_heads;
269            let attn_in = rms_norm(&hidden, &layer.attn_norm, hp.rms_norm_eps);
270
271            let mut q = layer.attn.q_proj.apply(&attn_in);
272            q = rms_norm_per_head(&q, &layer.attn.q_norm, head_dim, hp.rms_norm_eps);
273
274            let theta = hp.rope_theta_for(il);
275            let freq = if hp.is_swa_layer(il) {
276                None
277            } else {
278                self.weights.rope_freqs.as_deref()
279            };
280            for h in 0..n_heads {
281                let slice = &mut q[h * head_dim..(h + 1) * head_dim];
282                match freq {
283                    Some(f) => apply_rope_with_freq_factors(slice, pos, theta, f),
284                    None => apply_rope(slice, pos, theta),
285                }
286            }
287            // Gemma-4 attention_scale = 1.0: compensate kernel's 1/sqrt(d).
288            let compensate = hp.attention_scale * (head_dim as f32).sqrt();
289            for v in q.iter_mut() {
290                *v *= compensate;
291            }
292
293            let attn_out = if hp.has_kv(il) {
294                let k_proj = layer
295                    .attn
296                    .k_proj
297                    .as_ref()
298                    .expect("has_kv layer missing attn_k");
299                let mut k = k_proj.apply(&attn_in);
300                let k_norm = layer
301                    .attn
302                    .k_norm
303                    .as_ref()
304                    .expect("has_kv layer missing attn_k_norm");
305                k = rms_norm_per_head(&k, k_norm, head_dim, hp.rms_norm_eps);
306                for h in 0..n_kv {
307                    let slice = &mut k[h * head_dim..(h + 1) * head_dim];
308                    match freq {
309                        Some(f) => apply_rope_with_freq_factors(slice, pos, theta, f),
310                        None => apply_rope(slice, pos, theta),
311                    }
312                }
313
314                let mut v = match layer.attn.v_proj.as_ref() {
315                    Some(vp) => vp.apply(&attn_in),
316                    None => k.clone(),
317                };
318                // ggml_rms_norm without weight on V
319                v = rms_norm_per_head(&v, &vec![1.0; head_dim], head_dim, hp.rms_norm_eps);
320
321                let cache = state.kv[il].as_mut().expect("kv slot");
322                cache
323                    .push(&k, &v)
324                    .expect("unbounded KvCache growth is infallible");
325
326                if hp.is_swa_layer(il) {
327                    causal_gqa_attention_windowed(
328                        &q,
329                        &cache.k,
330                        &cache.v,
331                        n_heads,
332                        n_kv,
333                        head_dim,
334                        cache.rows(),
335                        hp.sliding_window,
336                    )
337                } else {
338                    causal_gqa_attention(
339                        &q,
340                        &cache.k,
341                        &cache.v,
342                        n_heads,
343                        n_kv,
344                        head_dim,
345                        cache.rows(),
346                    )
347                }
348            } else {
349                let reuse = hp.kv_reuse_layer(il);
350                let cache = state.kv[reuse]
351                    .as_ref()
352                    .expect("reuse kv layer missing cache");
353                // Shared-KV layers may have a different head_dim than the
354                // reused cache (SWA vs full). Attention only works when dims match.
355                assert_eq!(
356                    cache.head_dim, head_dim,
357                    "gemma4 shared-KV reuse head_dim mismatch layer {il} -> {reuse}"
358                );
359                if hp.is_swa_layer(il) {
360                    causal_gqa_attention_windowed(
361                        &q,
362                        &cache.k,
363                        &cache.v,
364                        n_heads,
365                        n_kv,
366                        head_dim,
367                        cache.rows(),
368                        hp.sliding_window,
369                    )
370                } else {
371                    causal_gqa_attention(
372                        &q,
373                        &cache.k,
374                        &cache.v,
375                        n_heads,
376                        n_kv,
377                        head_dim,
378                        cache.rows(),
379                    )
380                }
381            };
382
383            let mut attn_proj = layer.attn.o_proj.apply(&attn_out);
384            attn_proj = rms_norm(&attn_proj, &layer.attn.post_attn_norm, hp.rms_norm_eps);
385            let mut attn_out_res = hidden;
386            for (a, p) in attn_out_res.iter_mut().zip(attn_proj.iter()) {
387                *a += p;
388            }
389
390            let ffn_in = rms_norm(&attn_out_res, &layer.ffn_norm, hp.rms_norm_eps);
391            let gate = layer.ffn_gate.apply(&ffn_in);
392            let up = layer.ffn_up.apply(&ffn_in);
393            let mut ffn_out = layer.ffn_down.apply(&geglu(&gate, &up));
394            ffn_out = rms_norm(&ffn_out, &layer.ffn_post_norm, hp.rms_norm_eps);
395
396            let mut cur = attn_out_res;
397            for (c, f) in cur.iter_mut().zip(ffn_out.iter()) {
398                *c += f;
399            }
400
401            if let (Some(gate_w), Some(proj_w), Some(post_n), Some(pl_in)) = (
402                layer.per_layer_inp_gate.as_ref(),
403                layer.per_layer_proj.as_ref(),
404                layer.per_layer_post_norm.as_ref(),
405                per_layer_in.as_ref(),
406            ) {
407                let pe_in = cur.clone();
408                let mut g = gate_w.apply(&cur);
409                for x in g.iter_mut() {
410                    *x = gelu(*x);
411                }
412                for (gx, p) in g.iter_mut().zip(pl_in[il].iter()) {
413                    *gx *= *p;
414                }
415                let mut pe = proj_w.apply(&g);
416                pe = rms_norm(&pe, post_n, hp.rms_norm_eps);
417                cur = pe_in;
418                for (c, p) in cur.iter_mut().zip(pe.iter()) {
419                    *c += p;
420                }
421            }
422
423            if let Some(s) = layer.out_scale {
424                for x in cur.iter_mut() {
425                    *x *= s;
426                }
427            }
428            hidden = cur;
429        }
430
431        let mut logits = self.weights.output_head.apply(&rms_norm(
432            &hidden,
433            &self.weights.output_norm,
434            hp.rms_norm_eps,
435        ));
436        if let Some(sc) = hp.final_logit_softcap {
437            softcap_inplace(&mut logits, sc);
438        }
439        logits
440    }
441}
442
443#[cfg(test)]
444mod tests {
445    use super::*;
446
447    fn sample_hp() -> Gemma4Hparams {
448        // E2B-like: 10 layers, shared_kv=4 → n_kv_from_start=6; SWA pattern 4+1
449        let n_layer = 10;
450        let mut is_swa = Vec::new();
451        for i in 0..n_layer {
452            is_swa.push((i + 1) % 5 != 0);
453        }
454        Gemma4Hparams {
455            arch: "gemma4".into(),
456            n_layer,
457            hidden_dim: 64,
458            ffn_dims: vec![128; n_layer],
459            n_heads: 4,
460            n_kv_heads: 1,
461            head_dim_full: 32,
462            head_dim_swa: 16,
463            sliding_window: 8,
464            is_swa,
465            n_layer_kv_from_start: 6,
466            embd_per_layer: 8,
467            rms_norm_eps: 1e-6,
468            rope_theta: 1_000_000.0,
469            rope_theta_swa: 10_000.0,
470            final_logit_softcap: Some(30.0),
471            attention_scale: 1.0,
472        }
473    }
474
475    #[test]
476    fn layer_routing_swa_and_shared_kv() {
477        let hp = sample_hp();
478        assert!(hp.is_swa_layer(0));
479        assert!(!hp.is_swa_layer(4));
480        assert_eq!(hp.head_dim(0), 16);
481        assert_eq!(hp.head_dim(4), 32);
482        assert!(hp.has_kv(5));
483        assert!(!hp.has_kv(6));
484        // layer 6 SWA → reuse 6-2=4; layer 9 full → reuse 6-1=5
485        assert_eq!(hp.kv_reuse_layer(6), 4);
486        assert_eq!(hp.kv_reuse_layer(9), 5);
487        assert_eq!(hp.rope_theta_for(0), 10_000.0);
488        assert_eq!(hp.rope_theta_for(4), 1_000_000.0);
489    }
490}