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_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 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 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 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 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 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}