use super::gemma4_cache::Gemma4KvCache;
use super::gemma4_config::Gemma4Config;
use super::gemma4_loading::load_weights;
use super::gemma4_ops::{
gemma4_apply_rope, gemma4_geglu_mlp, gemma4_gelu_tanh, gemma4_logit_softcap, gemma4_rms_norm,
gemma4_rope_cos_sin, gemma4_rope_inv_freq, gemma4_scaled_embedding,
};
use super::gemma4_weights::Gemma4Weights;
use crate::error::InferenceError;
use crate::forward::cpu::{elementwise_mul, matmul_bt, rms_norm};
use crate::tokenizer::gemma_bpe::GemmaBpeTokenizer;
use crate::weights::SafetensorsFile;
use std::path::Path;
type LayerProbeTrace = Vec<(usize, Vec<f32>)>;
pub(crate) struct Gemma4Scratch {
hidden: Vec<f32>,
residual: Vec<f32>,
normed: Vec<f32>,
q: Vec<f32>,
k: Vec<f32>,
v: Vec<f32>,
scores: Vec<f32>,
context: Vec<f32>,
attn_out: Vec<f32>,
residual2: Vec<f32>,
normed2: Vec<f32>,
ffn_out: Vec<f32>,
residual3: Vec<f32>,
gate: Vec<f32>,
proj: Vec<f32>,
mlp_gate: Vec<f32>,
mlp_up: Vec<f32>,
v_norm_ones: Vec<f32>,
pub(crate) logits: Vec<f32>,
}
impl Gemma4Scratch {
pub(crate) fn new(cfg: &Gemma4Config) -> Self {
let hidden_size = cfg.hidden_size;
let widest_head = cfg.head_dim.max(cfg.global_head_dim);
let q_dim_max = cfg.num_attention_heads * widest_head;
let kv_dim_max = cfg.num_key_value_heads * widest_head;
let widest_mlp_intermediate = (0..cfg.num_hidden_layers)
.map(|layer_idx| cfg.mlp_intermediate_size(layer_idx))
.max()
.unwrap_or(cfg.intermediate_size);
Self {
hidden: vec![0f32; hidden_size],
residual: vec![0f32; hidden_size],
normed: vec![0f32; hidden_size],
q: vec![0f32; q_dim_max],
k: vec![0f32; kv_dim_max],
v: vec![0f32; kv_dim_max],
scores: Vec::new(),
context: vec![0f32; q_dim_max],
attn_out: vec![0f32; hidden_size],
residual2: vec![0f32; hidden_size],
normed2: vec![0f32; hidden_size],
ffn_out: vec![0f32; hidden_size],
residual3: vec![0f32; hidden_size],
gate: vec![0f32; cfg.hidden_size_per_layer_input],
proj: vec![0f32; hidden_size],
mlp_gate: vec![0f32; widest_mlp_intermediate],
mlp_up: vec![0f32; widest_mlp_intermediate],
v_norm_ones: vec![1.0f32; widest_head],
logits: vec![0f32; cfg.vocab_size],
}
}
fn ensure_scores_capacity(&mut self, n: usize) {
if self.scores.len() < n {
self.scores.resize(n, 0.0);
}
}
}
pub struct Gemma4Model {
pub(crate) config: Gemma4Config,
pub(crate) weights: Gemma4Weights,
pub(crate) tokenizer: GemmaBpeTokenizer,
local_inv_freq: Vec<f32>,
global_inv_freq: Vec<f32>,
}
impl Gemma4Model {
pub fn from_safetensors(path: &Path) -> Result<Self, InferenceError> {
let config = Gemma4Config::from_model_dir(path)?;
let model_path = path.join("model.safetensors");
if !model_path.exists() {
return Err(InferenceError::ModelNotFound(format!(
"missing model.safetensors in {}",
path.display()
)));
}
let mut source = SafetensorsFile::open(&model_path)?;
let weights = load_weights(&mut source, &config)?;
let tokenizer_path = path.join("tokenizer.json");
let tokenizer = GemmaBpeTokenizer::from_tokenizer_json(&tokenizer_path)?;
let local_inv_freq =
gemma4_rope_inv_freq(config.head_dim, config.rope_local_base_freq, None);
let global_inv_freq = gemma4_rope_inv_freq(
config.global_head_dim,
config.rope_theta,
Some(config.partial_rotary_factor),
);
Ok(Self {
config,
weights,
tokenizer,
local_inv_freq,
global_inv_freq,
})
}
pub fn config(&self) -> &Gemma4Config {
&self.config
}
pub fn tokenizer(&self) -> &GemmaBpeTokenizer {
&self.tokenizer
}
pub fn new_cache(&self, max_seq_len: usize) -> Result<Gemma4KvCache, InferenceError> {
Gemma4KvCache::new(&self.config, max_seq_len)
}
pub(crate) fn forward_step(
&self,
token_id: u32,
position: usize,
cache: &mut Gemma4KvCache,
scratch: &mut Gemma4Scratch,
capture_layers: &[usize],
) -> Result<LayerProbeTrace, InferenceError> {
let cfg = &self.config;
let hidden_size = cfg.hidden_size;
if token_id as usize >= cfg.vocab_size {
return Err(InferenceError::InvalidInput(format!(
"gemma4 forward: token_id {token_id} out of range (vocab_size={})",
cfg.vocab_size
)));
}
let per_layer_dim = cfg.hidden_size_per_layer_input;
let ple_packed_dim = cfg.num_hidden_layers * per_layer_dim;
gemma4_scaled_embedding(
&[token_id],
&self.weights.embed_tokens,
hidden_size,
&mut scratch.hidden[..hidden_size],
);
let per_layer_inputs = self.compute_per_layer_inputs(
&scratch.hidden[..hidden_size],
token_id,
per_layer_dim,
ple_packed_dim,
);
let (cos_local, sin_local) = gemma4_rope_cos_sin(&self.local_inv_freq, &[position as u32]);
let (cos_global, sin_global) =
gemma4_rope_cos_sin(&self.global_inv_freq, &[position as u32]);
let mut captured = Vec::with_capacity(capture_layers.len());
for layer_idx in 0..cfg.num_hidden_layers {
let lw = &self.weights.layers[layer_idx];
let is_global = cfg.is_global_layer(layer_idx);
let is_shared = cfg.is_kv_shared_layer(layer_idx);
let head_w = cfg.attn_head_dim(layer_idx);
let num_q_heads = cfg.num_attention_heads;
let num_kv_heads = cfg.num_key_value_heads;
let q_dim = num_q_heads * head_w;
let kv_dim = num_kv_heads * head_w;
let (cos, sin) = if is_global {
(&cos_global, &sin_global)
} else {
(&cos_local, &sin_local)
};
scratch.residual[..hidden_size].copy_from_slice(&scratch.hidden[..hidden_size]);
scratch.normed[..hidden_size].copy_from_slice(&scratch.hidden[..hidden_size]);
gemma4_rms_norm(
&mut scratch.normed[..hidden_size],
&lw.input_layernorm,
hidden_size,
cfg.rms_norm_eps,
);
matmul_bt(
&scratch.normed[..hidden_size],
&lw.q_proj,
&mut scratch.q[..q_dim],
1,
hidden_size,
q_dim,
);
for h in 0..num_q_heads {
let start = h * head_w;
gemma4_rms_norm(
&mut scratch.q[start..start + head_w],
&lw.q_norm,
head_w,
cfg.rms_norm_eps,
);
}
gemma4_apply_rope(&mut scratch.q[..q_dim], cos, sin, 1, num_q_heads, head_w);
if !is_shared {
let k_proj = lw.k_proj.as_ref().ok_or_else(|| {
InferenceError::Inference(format!(
"gemma4 forward: layer {layer_idx} is non-shared but has no k_proj weights"
))
})?;
let v_proj = lw.v_proj.as_ref().ok_or_else(|| {
InferenceError::Inference(format!(
"gemma4 forward: layer {layer_idx} is non-shared but has no v_proj weights"
))
})?;
let k_norm = lw.k_norm.as_ref().ok_or_else(|| {
InferenceError::Inference(format!(
"gemma4 forward: layer {layer_idx} is non-shared but has no k_norm weights"
))
})?;
matmul_bt(
&scratch.normed[..hidden_size],
k_proj,
&mut scratch.k[..kv_dim],
1,
hidden_size,
kv_dim,
);
matmul_bt(
&scratch.normed[..hidden_size],
v_proj,
&mut scratch.v[..kv_dim],
1,
hidden_size,
kv_dim,
);
for h in 0..num_kv_heads {
let start = h * head_w;
gemma4_rms_norm(
&mut scratch.k[start..start + head_w],
k_norm,
head_w,
cfg.rms_norm_eps,
);
}
gemma4_apply_rope(&mut scratch.k[..kv_dim], cos, sin, 1, num_kv_heads, head_w);
for h in 0..num_kv_heads {
let start = h * head_w;
rms_norm(
&mut scratch.v[start..start + head_w],
&scratch.v_norm_ones[..head_w],
head_w,
cfg.rms_norm_eps,
);
}
cache.append_kv(layer_idx, &scratch.k[..kv_dim], &scratch.v[..kv_dim])?;
}
let seq_len = cache.seq_len(layer_idx)?;
let k_view = cache.k_view(layer_idx)?;
let v_view = cache.v_view(layer_idx)?;
let groups = num_q_heads / num_kv_heads;
scratch.ensure_scores_capacity(seq_len);
for qh in 0..num_q_heads {
let kvh = qh / groups;
let q_head = &scratch.q[qh * head_w..(qh + 1) * head_w];
let scores = &mut scratch.scores[..seq_len];
for t in 0..seq_len {
let k_off = t * kv_dim + kvh * head_w;
let mut dot = 0.0f32;
for d in 0..head_w {
dot += q_head[d] * k_view[k_off + d];
}
scores[t] = dot;
}
softmax_row_fail_closed(scores);
let ctx_off = qh * head_w;
for d in 0..head_w {
let mut sum = 0.0f32;
for t in 0..seq_len {
let v_off = t * kv_dim + kvh * head_w;
sum += scratch.scores[t] * v_view[v_off + d];
}
scratch.context[ctx_off + d] = sum;
}
}
matmul_bt(
&scratch.context[..q_dim],
&lw.o_proj,
&mut scratch.attn_out[..hidden_size],
1,
q_dim,
hidden_size,
);
gemma4_rms_norm(
&mut scratch.attn_out[..hidden_size],
&lw.post_attention_layernorm,
hidden_size,
cfg.rms_norm_eps,
);
for i in 0..hidden_size {
scratch.hidden[i] = scratch.residual[i] + scratch.attn_out[i];
}
scratch.residual2[..hidden_size].copy_from_slice(&scratch.hidden[..hidden_size]);
scratch.normed2[..hidden_size].copy_from_slice(&scratch.hidden[..hidden_size]);
gemma4_rms_norm(
&mut scratch.normed2[..hidden_size],
&lw.pre_feedforward_layernorm,
hidden_size,
cfg.rms_norm_eps,
);
let mlp_dim = cfg.mlp_intermediate_size(layer_idx);
gemma4_geglu_mlp(
&scratch.normed2[..hidden_size],
&lw.gate_proj,
&lw.up_proj,
&lw.down_proj,
1,
hidden_size,
mlp_dim,
&mut scratch.mlp_gate[..mlp_dim],
&mut scratch.mlp_up[..mlp_dim],
&mut scratch.ffn_out[..hidden_size],
);
gemma4_rms_norm(
&mut scratch.ffn_out[..hidden_size],
&lw.post_feedforward_layernorm,
hidden_size,
cfg.rms_norm_eps,
);
for i in 0..hidden_size {
scratch.hidden[i] = scratch.residual2[i] + scratch.ffn_out[i];
}
scratch.residual3[..hidden_size].copy_from_slice(&scratch.hidden[..hidden_size]);
matmul_bt(
&scratch.hidden[..hidden_size],
&lw.per_layer_input_gate,
&mut scratch.gate[..per_layer_dim],
1,
hidden_size,
per_layer_dim,
);
gemma4_gelu_tanh(&mut scratch.gate[..per_layer_dim]);
let this_layer_input =
&per_layer_inputs[layer_idx * per_layer_dim..(layer_idx + 1) * per_layer_dim];
elementwise_mul(&mut scratch.gate[..per_layer_dim], this_layer_input);
matmul_bt(
&scratch.gate[..per_layer_dim],
&lw.per_layer_projection,
&mut scratch.proj[..hidden_size],
1,
per_layer_dim,
hidden_size,
);
gemma4_rms_norm(
&mut scratch.proj[..hidden_size],
&lw.post_per_layer_input_norm,
hidden_size,
cfg.rms_norm_eps,
);
for i in 0..hidden_size {
scratch.hidden[i] = scratch.residual3[i] + scratch.proj[i];
}
for v in scratch.hidden[..hidden_size].iter_mut() {
*v *= lw.layer_scalar;
}
if capture_layers.contains(&layer_idx) {
captured.push((layer_idx, scratch.hidden[..hidden_size].to_vec()));
}
}
gemma4_rms_norm(
&mut scratch.hidden[..hidden_size],
&self.weights.norm,
hidden_size,
cfg.rms_norm_eps,
);
matmul_bt(
&scratch.hidden[..hidden_size],
&self.weights.embed_tokens,
&mut scratch.logits[..cfg.vocab_size],
1,
hidden_size,
cfg.vocab_size,
);
gemma4_logit_softcap(
&mut scratch.logits[..cfg.vocab_size],
cfg.final_logit_softcapping,
);
Ok(captured)
}
fn compute_per_layer_inputs(
&self,
scaled_embed: &[f32],
token_id: u32,
per_layer_dim: usize,
ple_packed_dim: usize,
) -> Vec<f32> {
let cfg = &self.config;
let hidden_size = cfg.hidden_size;
let id_scale = (per_layer_dim as f32).sqrt();
let row_start = token_id as usize * ple_packed_dim;
let mut identity: Vec<f32> = self.weights.embed_tokens_per_layer
[row_start..row_start + ple_packed_dim]
.iter()
.map(|&v| v * id_scale)
.collect();
let mut ctx = vec![0f32; ple_packed_dim];
matmul_bt(
scaled_embed,
&self.weights.per_layer_model_projection,
&mut ctx,
1,
hidden_size,
ple_packed_dim,
);
let ctx_scale = 1.0 / (hidden_size as f32).sqrt();
for v in ctx.iter_mut() {
*v *= ctx_scale;
}
for layer in 0..cfg.num_hidden_layers {
let start = layer * per_layer_dim;
gemma4_rms_norm(
&mut ctx[start..start + per_layer_dim],
&self.weights.per_layer_projection_norm,
per_layer_dim,
cfg.rms_norm_eps,
);
}
let combine_scale = std::f32::consts::FRAC_1_SQRT_2;
for i in 0..ple_packed_dim {
identity[i] = (ctx[i] + identity[i]) * combine_scale;
}
identity
}
pub fn generate_greedy(
&self,
prompt_ids: &[u32],
max_new_tokens: usize,
max_seq_len: usize,
) -> Result<Vec<u32>, InferenceError> {
let mut cache = self.new_cache(max_seq_len)?;
let mut scratch = Gemma4Scratch::new(&self.config);
let mut generated = Vec::with_capacity(max_new_tokens);
let mut position = 0usize;
for &tok in prompt_ids {
self.forward_step(tok, position, &mut cache, &mut scratch, &[])?;
position += 1;
}
for _ in 0..max_new_tokens {
let next = argmax(&scratch.logits[..self.config.vocab_size]);
generated.push(next);
self.forward_step(next, position, &mut cache, &mut scratch, &[])?;
position += 1;
}
Ok(generated)
}
pub fn generate_greedy_with_probe(
&self,
prompt_ids: &[u32],
max_new_tokens: usize,
max_seq_len: usize,
probe_layers: &[usize],
) -> Result<(Vec<u32>, Vec<f32>, LayerProbeTrace), InferenceError> {
let mut cache = self.new_cache(max_seq_len)?;
let mut scratch = Gemma4Scratch::new(&self.config);
let mut generated = Vec::with_capacity(max_new_tokens);
let mut position = 0usize;
let mut probe = Vec::new();
for (i, &tok) in prompt_ids.iter().enumerate() {
let is_last = i + 1 == prompt_ids.len();
let layers: &[usize] = if is_last { probe_layers } else { &[] };
let captured = self.forward_step(tok, position, &mut cache, &mut scratch, layers)?;
if is_last {
probe = captured;
}
position += 1;
}
let final_logits = scratch.logits[..self.config.vocab_size].to_vec();
for _ in 0..max_new_tokens {
let next = argmax(&scratch.logits[..self.config.vocab_size]);
generated.push(next);
self.forward_step(next, position, &mut cache, &mut scratch, &[])?;
position += 1;
}
Ok((generated, final_logits, probe))
}
}
fn softmax_row_fail_closed(row: &mut [f32]) {
let max = row.iter().copied().fold(f32::NEG_INFINITY, f32::max);
if !max.is_finite() {
row.fill(0.0);
if !row.is_empty() {
row[row.len() - 1] = 1.0;
}
return;
}
let mut sum = 0.0f32;
for v in row.iter_mut() {
let e = (*v - max).exp();
*v = e;
sum += e;
}
if !sum.is_finite() || sum <= 0.0 {
row.fill(0.0);
if !row.is_empty() {
row[row.len() - 1] = 1.0;
}
return;
}
for v in row.iter_mut() {
*v /= sum;
}
}
fn argmax(logits: &[f32]) -> u32 {
let mut best_idx = 0usize;
let mut best_val = f32::NEG_INFINITY;
for (i, &v) in logits.iter().enumerate() {
if v > best_val {
best_val = v;
best_idx = i;
}
}
best_idx as u32
}
#[cfg(all(test, feature = "f16"))]
mod donor_mutation_tests {
use super::*;
fn model_dir() -> Option<std::path::PathBuf> {
let raw = std::env::var("LATTICE_GEMMA4_MODEL_DIR")
.unwrap_or_else(|_| "~/.lattice/models/gemma-4-e2b-it".to_string());
let path = if let Some(rest) = raw.strip_prefix("~/") {
std::path::PathBuf::from(std::env::var("HOME").ok()?).join(rest)
} else {
std::path::PathBuf::from(raw)
};
path.join("model.safetensors").exists().then_some(path)
}
fn run_greedy_with_probe_on_cache(
model: &Gemma4Model,
input_ids: &[u32],
max_new_tokens: usize,
cache: &mut Gemma4KvCache,
probe_layers: &[usize],
) -> (Vec<u32>, LayerProbeTrace) {
let mut scratch = Gemma4Scratch::new(&model.config);
let mut generated = Vec::with_capacity(max_new_tokens);
let mut position = 0usize;
let mut probe = Vec::new();
for (i, &tok) in input_ids.iter().enumerate() {
let is_last = i + 1 == input_ids.len();
let layers: &[usize] = if is_last { probe_layers } else { &[] };
let captured = model
.forward_step(tok, position, cache, &mut scratch, layers)
.expect("forward_step");
if is_last {
probe = captured;
}
position += 1;
}
for _ in 0..max_new_tokens {
let next = argmax(&scratch.logits[..model.config.vocab_size]);
generated.push(next);
model
.forward_step(next, position, cache, &mut scratch, &[])
.expect("forward_step");
position += 1;
}
(generated, probe)
}
#[test]
fn wrong_donor_mapping_diverges_layer34_probe_and_greedy_tokens() {
let Some(dir) = model_dir() else {
eprintln!("LATTICE_GEMMA4_MUTATION_TEST_SKIPPED reason=missing_checkpoint");
return;
};
let model = Gemma4Model::from_safetensors(&dir).expect("loading real checkpoint");
let input_ids: Vec<u32> = vec![2, 818, 5279, 529, 7001, 563];
let (baseline_greedy, _, baseline_probe) = model
.generate_greedy_with_probe(&input_ids, 3, 64, &[34])
.expect("baseline forward");
let mut mutated_cache = model.new_cache(64).expect("cache construction");
assert_eq!(
mutated_cache.layer_slot(34).unwrap(),
14,
"sanity: correct donor before mutation"
);
mutated_cache.override_layer_slot_for_test(34, 9);
assert_eq!(mutated_cache.layer_slot(34).unwrap(), 9);
let (mutated_greedy, mutated_probe) =
run_greedy_with_probe_on_cache(&model, &input_ids, 3, &mut mutated_cache, &[34]);
let baseline_hidden = &baseline_probe[0].1;
let mutated_hidden = &mutated_probe[0].1;
let diff = baseline_hidden
.iter()
.zip(mutated_hidden.iter())
.map(|(a, b)| (a - b).abs())
.fold(0f32, f32::max);
const E2E_GATE_TOLERANCE: f32 = 1e-3;
assert!(
diff > E2E_GATE_TOLERANCE * 10.0,
"wrong-donor mutation must blow through the e2e gate's own {E2E_GATE_TOLERANCE} \
tolerance by a wide margin (got diff {diff}) -- otherwise this test is decorative"
);
eprintln!(
"donor mutation: layer 34 hidden-state max-abs-diff={diff} (gate tolerance \
{E2E_GATE_TOLERANCE}, correct-donor baseline ~1e-5) -- fails the e2e gate's \
per-layer probe assertion. greedy tokens baseline={baseline_greedy:?} \
mutated={mutated_greedy:?} (top-1 margin at this prompt is wide enough, ~5 \
logit points, that this single-layer perturbation does not always flip argmax; \
the per-layer probe assertion is the gate this mutation is proven against)."
);
}
}