1use 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
20pub const GEMMA4_ARCHES: &[&str] = &["gemma4", "gemma4-assistant"];
22
23#[derive(Debug, Clone)]
25pub struct Gemma4Hparams {
26 pub arch: String,
27 pub n_layer: usize,
28 pub hidden_dim: usize,
29 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 pub is_swa: Vec<bool>,
38 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 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 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 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 pub kv: Vec<Option<KvCache>>,
132}
133
134impl Gemma4Engine {
135 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 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 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 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 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 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 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 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}