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(&self, token_id: usize, pos: usize, state: &mut Self::State) -> Vec<f32> {
251        let hp = &self.hp;
252        let mut hidden = self.weights.token_embd.dequant_row(token_id);
253        let emb_scale = (hp.hidden_dim as f32).sqrt();
254        for x in hidden.iter_mut() {
255            *x *= emb_scale;
256        }
257
258        let per_layer_in = self.project_per_layer_inputs(token_id, &hidden);
259
260        for (il, layer) in self.weights.layers.iter().enumerate() {
261            let head_dim = hp.head_dim(il);
262            let n_heads = hp.n_heads;
263            let n_kv = hp.n_kv_heads;
264            let attn_in = rms_norm(&hidden, &layer.attn_norm, hp.rms_norm_eps);
265
266            let mut q = layer.attn.q_proj.apply(&attn_in);
267            q = rms_norm_per_head(&q, &layer.attn.q_norm, head_dim, hp.rms_norm_eps);
268
269            let theta = hp.rope_theta_for(il);
270            let freq = if hp.is_swa_layer(il) {
271                None
272            } else {
273                self.weights.rope_freqs.as_deref()
274            };
275            for h in 0..n_heads {
276                let slice = &mut q[h * head_dim..(h + 1) * head_dim];
277                match freq {
278                    Some(f) => apply_rope_with_freq_factors(slice, pos, theta, f),
279                    None => apply_rope(slice, pos, theta),
280                }
281            }
282            // Gemma-4 attention_scale = 1.0: compensate kernel's 1/sqrt(d).
283            let compensate = hp.attention_scale * (head_dim as f32).sqrt();
284            for v in q.iter_mut() {
285                *v *= compensate;
286            }
287
288            let attn_out = if hp.has_kv(il) {
289                let k_proj = layer
290                    .attn
291                    .k_proj
292                    .as_ref()
293                    .expect("has_kv layer missing attn_k");
294                let mut k = k_proj.apply(&attn_in);
295                let k_norm = layer
296                    .attn
297                    .k_norm
298                    .as_ref()
299                    .expect("has_kv layer missing attn_k_norm");
300                k = rms_norm_per_head(&k, k_norm, head_dim, hp.rms_norm_eps);
301                for h in 0..n_kv {
302                    let slice = &mut k[h * head_dim..(h + 1) * head_dim];
303                    match freq {
304                        Some(f) => apply_rope_with_freq_factors(slice, pos, theta, f),
305                        None => apply_rope(slice, pos, theta),
306                    }
307                }
308
309                let mut v = match layer.attn.v_proj.as_ref() {
310                    Some(vp) => vp.apply(&attn_in),
311                    None => k.clone(),
312                };
313                // ggml_rms_norm without weight on V
314                v = rms_norm_per_head(&v, &vec![1.0; head_dim], head_dim, hp.rms_norm_eps);
315
316                let cache = state.kv[il].as_mut().expect("kv slot");
317                cache
318                    .push(&k, &v)
319                    .expect("unbounded KvCache growth is infallible");
320
321                if hp.is_swa_layer(il) {
322                    causal_gqa_attention_windowed(
323                        &q,
324                        &cache.k,
325                        &cache.v,
326                        n_heads,
327                        n_kv,
328                        head_dim,
329                        cache.seq_len,
330                        hp.sliding_window,
331                    )
332                } else {
333                    causal_gqa_attention(
334                        &q,
335                        &cache.k,
336                        &cache.v,
337                        n_heads,
338                        n_kv,
339                        head_dim,
340                        cache.seq_len,
341                    )
342                }
343            } else {
344                let reuse = hp.kv_reuse_layer(il);
345                let cache = state.kv[reuse]
346                    .as_ref()
347                    .expect("reuse kv layer missing cache");
348                // Shared-KV layers may have a different head_dim than the
349                // reused cache (SWA vs full). Attention only works when dims match.
350                assert_eq!(
351                    cache.head_dim, head_dim,
352                    "gemma4 shared-KV reuse head_dim mismatch layer {il} -> {reuse}"
353                );
354                if hp.is_swa_layer(il) {
355                    causal_gqa_attention_windowed(
356                        &q,
357                        &cache.k,
358                        &cache.v,
359                        n_heads,
360                        n_kv,
361                        head_dim,
362                        cache.seq_len,
363                        hp.sliding_window,
364                    )
365                } else {
366                    causal_gqa_attention(
367                        &q,
368                        &cache.k,
369                        &cache.v,
370                        n_heads,
371                        n_kv,
372                        head_dim,
373                        cache.seq_len,
374                    )
375                }
376            };
377
378            let mut attn_proj = layer.attn.o_proj.apply(&attn_out);
379            attn_proj = rms_norm(&attn_proj, &layer.attn.post_attn_norm, hp.rms_norm_eps);
380            let mut attn_out_res = hidden;
381            for (a, p) in attn_out_res.iter_mut().zip(attn_proj.iter()) {
382                *a += p;
383            }
384
385            let ffn_in = rms_norm(&attn_out_res, &layer.ffn_norm, hp.rms_norm_eps);
386            let gate = layer.ffn_gate.apply(&ffn_in);
387            let up = layer.ffn_up.apply(&ffn_in);
388            let mut ffn_out = layer.ffn_down.apply(&geglu(&gate, &up));
389            ffn_out = rms_norm(&ffn_out, &layer.ffn_post_norm, hp.rms_norm_eps);
390
391            let mut cur = attn_out_res;
392            for (c, f) in cur.iter_mut().zip(ffn_out.iter()) {
393                *c += f;
394            }
395
396            if let (Some(gate_w), Some(proj_w), Some(post_n), Some(pl_in)) = (
397                layer.per_layer_inp_gate.as_ref(),
398                layer.per_layer_proj.as_ref(),
399                layer.per_layer_post_norm.as_ref(),
400                per_layer_in.as_ref(),
401            ) {
402                let pe_in = cur.clone();
403                let mut g = gate_w.apply(&cur);
404                for x in g.iter_mut() {
405                    *x = gelu(*x);
406                }
407                for (gx, p) in g.iter_mut().zip(pl_in[il].iter()) {
408                    *gx *= *p;
409                }
410                let mut pe = proj_w.apply(&g);
411                pe = rms_norm(&pe, post_n, hp.rms_norm_eps);
412                cur = pe_in;
413                for (c, p) in cur.iter_mut().zip(pe.iter()) {
414                    *c += p;
415                }
416            }
417
418            if let Some(s) = layer.out_scale {
419                for x in cur.iter_mut() {
420                    *x *= s;
421                }
422            }
423            hidden = cur;
424        }
425
426        let mut logits = self.weights.output_head.apply(&rms_norm(
427            &hidden,
428            &self.weights.output_norm,
429            hp.rms_norm_eps,
430        ));
431        if let Some(sc) = hp.final_logit_softcap {
432            softcap_inplace(&mut logits, sc);
433        }
434        logits
435    }
436}
437
438#[cfg(test)]
439mod tests {
440    use super::*;
441
442    fn sample_hp() -> Gemma4Hparams {
443        // E2B-like: 10 layers, shared_kv=4 → n_kv_from_start=6; SWA pattern 4+1
444        let n_layer = 10;
445        let mut is_swa = Vec::new();
446        for i in 0..n_layer {
447            is_swa.push((i + 1) % 5 != 0);
448        }
449        Gemma4Hparams {
450            arch: "gemma4".into(),
451            n_layer,
452            hidden_dim: 64,
453            ffn_dims: vec![128; n_layer],
454            n_heads: 4,
455            n_kv_heads: 1,
456            head_dim_full: 32,
457            head_dim_swa: 16,
458            sliding_window: 8,
459            is_swa,
460            n_layer_kv_from_start: 6,
461            embd_per_layer: 8,
462            rms_norm_eps: 1e-6,
463            rope_theta: 1_000_000.0,
464            rope_theta_swa: 10_000.0,
465            final_logit_softcap: Some(30.0),
466            attention_scale: 1.0,
467        }
468    }
469
470    #[test]
471    fn layer_routing_swa_and_shared_kv() {
472        let hp = sample_hp();
473        assert!(hp.is_swa_layer(0));
474        assert!(!hp.is_swa_layer(4));
475        assert_eq!(hp.head_dim(0), 16);
476        assert_eq!(hp.head_dim(4), 32);
477        assert!(hp.has_kv(5));
478        assert!(!hp.has_kv(6));
479        // layer 6 SWA → reuse 6-2=4; layer 9 full → reuse 6-1=5
480        assert_eq!(hp.kv_reuse_layer(6), 4);
481        assert_eq!(hp.kv_reuse_layer(9), 5);
482        assert_eq!(hp.rope_theta_for(0), 10_000.0);
483        assert_eq!(hp.rope_theta_for(4), 1_000_000.0);
484    }
485}