Fix Gemma 4 KV-shared layers creating unused projections (#1158)

This commit is contained in:
glyphVault
2026-04-21 16:44:13 -07:00
committed by GitHub
parent 3cd9a52df2
commit 4f5cbd2a4f
2 changed files with 87 additions and 6 deletions
+14 -5
View File
@@ -180,6 +180,7 @@ class Attention(nn.Module):
self.layer_idx = layer_idx
self.layer_type = config.layer_types[layer_idx]
self.is_sliding = self.layer_type == "sliding_attention"
self.has_kv = layer_idx < config.num_hidden_layers - config.num_kv_shared_layers
self.head_dim = (
config.global_head_dim
@@ -202,14 +203,18 @@ class Attention(nn.Module):
self.scale = 1.0
self.q_proj = nn.Linear(dim, self.n_heads * self.head_dim, bias=False)
self.k_proj = nn.Linear(dim, self.n_kv_heads * self.head_dim, bias=False)
if not self.use_k_eq_v:
self.v_proj = nn.Linear(dim, self.n_kv_heads * self.head_dim, bias=False)
if self.has_kv:
self.k_proj = nn.Linear(dim, self.n_kv_heads * self.head_dim, bias=False)
if not self.use_k_eq_v:
self.v_proj = nn.Linear(
dim, self.n_kv_heads * self.head_dim, bias=False
)
self.o_proj = nn.Linear(self.n_heads * self.head_dim, dim, bias=False)
self.q_norm = nn.RMSNorm(self.head_dim, eps=config.rms_norm_eps)
self.k_norm = nn.RMSNorm(self.head_dim, eps=config.rms_norm_eps)
self.v_norm = RMSNormNoScale(self.head_dim, eps=config.rms_norm_eps)
if self.has_kv:
self.k_norm = nn.RMSNorm(self.head_dim, eps=config.rms_norm_eps)
self.v_norm = RMSNormNoScale(self.head_dim, eps=config.rms_norm_eps)
# RoPE (with partial rotation support)
layer_key = "sliding_attention" if self.is_sliding else "full_attention"
@@ -238,6 +243,10 @@ class Attention(nn.Module):
if shared_kv is not None:
keys, values = shared_kv
elif not self.has_kv:
raise ValueError(
f"Layer {self.layer_idx} is a KV-shared layer but received no shared_kv"
)
else:
keys = self.k_proj(x).reshape(B, L, self.n_kv_heads, self.head_dim)
values = keys
+73 -1
View File
@@ -5,7 +5,7 @@ import unittest
import mlx.core as mx
import mlx.nn as nn
from mlx.utils import tree_map
from mlx.utils import tree_flatten, tree_map
from mlx_lm.models import rope_utils
from mlx_lm.models.base import create_causal_mask, scaled_dot_product_attention
@@ -1547,6 +1547,78 @@ class TestModels(unittest.TestCase):
mx.allclose(logits, mx.ones((1, 1, 4), dtype=mx.float32) * 32.0)
)
def test_gemma4_kv_shared_layers_omit_kv_projections(self):
"""KV-shared layers must not create k_proj/v_proj/k_norm/v_norm so that
models saved without redundant weights (e.g. via transformers
save_pretrained) can be loaded with strict=True."""
from mlx_lm.models import gemma4_text
args = gemma4_text.ModelArgs(
model_type="gemma4_text",
hidden_size=128,
num_hidden_layers=10,
intermediate_size=256,
num_attention_heads=4,
head_dim=32,
global_head_dim=64,
rms_norm_eps=1e-6,
vocab_size=1000,
vocab_size_per_layer_input=1000,
num_key_value_heads=1,
num_kv_shared_layers=4,
hidden_size_per_layer_input=32,
sliding_window=8,
sliding_window_pattern=5,
final_logit_softcapping=30.0,
layer_types=[
"sliding_attention",
"sliding_attention",
"sliding_attention",
"sliding_attention",
"full_attention",
"sliding_attention",
"sliding_attention",
"sliding_attention",
"sliding_attention",
"full_attention",
],
rope_parameters={
"full_attention": {
"partial_rotary_factor": 0.25,
"rope_theta": 1000000.0,
},
"sliding_attention": {
"rope_theta": 10000.0,
},
},
)
model = gemma4_text.Model(args)
# Non-shared layers (0-5) should have KV projections
for i in range(6):
attn = model.model.layers[i].self_attn
self.assertTrue(attn.has_kv)
self.assertTrue(hasattr(attn, "k_proj"))
self.assertTrue(hasattr(attn, "k_norm"))
# Shared layers (6-9) should NOT have KV projections
for i in range(6, 10):
attn = model.model.layers[i].self_attn
self.assertFalse(attn.has_kv)
self.assertFalse(hasattr(attn, "k_proj"))
self.assertFalse(hasattr(attn, "k_norm"))
self.assertFalse(hasattr(attn, "v_proj"))
# Verify the model can load weights that omit shared-layer KV params
weights = dict(tree_flatten(model.parameters()))
kv_keys = [
k for k in weights if "k_proj" in k or "v_proj" in k or "k_norm" in k
]
for k in kv_keys:
# All KV keys should belong to non-shared layers (0-5)
layer_idx = int(k.split("layers.")[1].split(".")[0])
self.assertLess(layer_idx, 6)
def test_gemma4_input_embeddings_reconstruct_per_layer_inputs(self):
from mlx_lm.models import gemma4_text