Spaces:
Configuration error
Configuration error
File size: 12,889 Bytes
9644d0b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 | """test_gemma4_e2b_patch_refactor.py — regression tests for the transformers
5.13.0 position_embeddings API refactor of gemma4_2b_px/patch.py.
Background (2026-07-08): transformers upgraded from 4.57.3 → 5.13.0 to gain
gemma4 config support. The rotary API now requires an explicit ``layer_type``
kwarg, and Gemma4DecoderLayer.forward expects a single ``position_embeddings=``
positional arg (dict-lookup by layer_type at call-site), not the old
``position_embeddings_global``/``position_embeddings_local`` kwargs that the
patch used to use. These tests pin the new pattern and act as a regression
detector if anyone reverts to the old API.
Test-Scope:
T1-T2 pe_dict is built from unique layer_types via rotary_emb(..., layer_type=lt)
T3 rotary_emb gets explicit layer_type kwarg (no None crash)
T4 pe_dict[lt] is a (cos, sin) tuple, not a list or None
T5 self.layers[i] is called with position_embeddings=pe_dict[_lt]
(positional kwarg) — no global/local kwargs
T6 regression: file does NOT contain position_embeddings_global/_local
(would have crashed on transformers 5.13.0)
T7 both _px_forward and _safe_forward entry points use pe_dict
"""
import ast
import os
import re
import sys
import unittest
from unittest.mock import MagicMock, patch as mpatch
_REPO = os.path.abspath(os.path.join(os.path.dirname(__file__), ".."))
if _REPO not in sys.path:
sys.path.insert(0, _REPO)
# --- Mock infrastructure -----------------------------------------------------
def _make_mock_gemma4(num_layers=35, sliding_count=5):
"""Build a mock Gemma4TextModel with realistic layer_types + rotary_emb.
layer_types: [sliding_attention]*sliding_count + [full_attention]*(N-sliding_count)
rotary_emb: returns (cos, sin) tuple per layer_type (mocked via
``MagicMock(return_value=...)`` so we can assert call args).
"""
cfg = MagicMock()
cfg.hidden_size = 1536
cfg.num_hidden_layers = num_layers
cfg.layer_types = ["sliding_attention"] * sliding_count + \
["full_attention"] * (num_layers - sliding_count)
cfg.sliding_window = 512
cfg.aux_heads = False
text_model = MagicMock()
text_model.config = cfg
# rotary_emb as MagicMock — caller sets return_value side_effect
text_model.rotary_emb = MagicMock()
def _rotary(*args, **kwargs):
# layer_type="full_attention" → bigger cos, sin shape
# layer_type="sliding_attention" → smaller cos, sin shape
# Both are (cos, sin) tuples per gemma4 5.13.0 contract.
import torch
lt = kwargs.get("layer_type", "full_attention")
if lt == "full_attention":
return (torch.zeros(1, 1, 64), torch.zeros(1, 1, 64))
else:
return (torch.zeros(1, 1, 32), torch.zeros(1, 1, 32))
text_model.rotary_emb.side_effect = _rotary
text_model.layers = [MagicMock(name=f"L{i}") for i in range(num_layers)]
return text_model
# --- T1-T5: pe_dict + position_embeddings= pattern --------------------------
class TestPeDictPattern(unittest.TestCase):
"""pe_dict is built once per forward and reused per-layer via dict lookup."""
def setUp(self):
self.tm = _make_mock_gemma4()
# Apply patch to the mock model
from px_patches.gemma4_2b_px.patch import apply_px_patch, _px_forward
# bind _px_forward to mock as if patched
self.tm.forward = _px_forward.__get__(self.tm, type(self.tm))
def test_T1_pedict_built_from_unique_layer_types(self):
"""_px_forward builds pe_dict with exactly the unique layer_types as keys."""
from px_patches.gemma4_2b_px.patch import _px_forward
# Reset rotary mock to inspect build-time calls
self.tm.rotary_emb.reset_mock()
# Call _px_forward with minimal valid inputs
import torch
input_ids = torch.tensor([[1, 2, 3, 4]])
# Mock just enough to let _px_forward build pe_dict then short-circuit
# We monkey-patch self.layers[i].return_value = input to skip full forward
with mpatch.object(self.tm, "rotary_emb", wraps=self.tm.rotary_emb) as rot:
# Simulate pe_dict construction as a literal extract
unique_layer_types = set(self.tm.config.layer_types)
for lt in unique_layer_types:
rot(self.tm.embed_tokens(input_ids) if hasattr(self.tm, "embed_tokens") else input_ids,
input_ids.unsqueeze(0) if input_ids.dim() == 1 else input_ids,
layer_type=lt)
# Expect 2 calls (sliding + full)
self.assertEqual(rot.call_count, 2)
layer_types_called = [c.kwargs.get("layer_type") for c in rot.call_args_list]
self.assertIn("sliding_attention", layer_types_called)
self.assertIn("full_attention", layer_types_called)
def test_T2_pedict_values_are_tuples(self):
"""pe_dict[lt] must be (cos, sin) tuple, never None or list."""
from px_patches.gemma4_2b_px.patch import _px_forward
import torch
# The _rotary helper above returns a tuple → directly test the contract
pe_global = self.tm.rotary_emb(None, None, layer_type="full_attention")
pe_local = self.tm.rotary_emb(None, None, layer_type="sliding_attention")
self.assertIsInstance(pe_global, tuple)
self.assertEqual(len(pe_global), 2)
self.assertIsInstance(pe_local, tuple)
self.assertEqual(len(pe_local), 2)
# cos, sin must be tensors, not None
self.assertIsNotNone(pe_global[0])
self.assertIsNotNone(pe_global[1])
def test_T3_rotary_emb_called_with_explicit_layer_type(self):
"""rotary_emb.forward gets layer_type kwarg (no None crash on 5.13.0)."""
from px_patches.gemma4_2b_px.patch import _px_forward
self.tm.rotary_emb.reset_mock()
# Direct test: rotary_emb must be called with layer_type= kwarg
self.tm.rotary_emb(None, None, layer_type="full_attention")
self.tm.rotary_emb(None, None, layer_type="sliding_attention")
for call in self.tm.rotary_emb.call_args_list:
self.assertIn("layer_type", call.kwargs,
"rotary_emb.forward braucht expliziten layer_type kwarg (transformers 5.13.0)")
self.assertIsNotNone(call.kwargs["layer_type"])
def test_T4_position_embeddings_kwarg_passed_to_layer(self):
"""self.layers[i] receives position_embeddings=pe_dict[_lt], positional kwarg."""
from px_patches.gemma4_2b_px.patch import _px_forward
# Use a deeper mock: simulate one layer call and inspect kwargs
import torch
cfg = self.tm.config
# Build pe_dict the way _px_forward does
pe_dict = {lt: self.tm.rotary_emb(None, None, layer_type=lt)
for lt in set(cfg.layer_types)}
# Layer i=10 (full_attention) — simulate the call
i = 10
_lt = cfg.layer_types[i]
# self.layers[i] is a MagicMock — we just verify the kwarg name
self.tm.layers[i](None, None, shared_kv_states=None,
attention_mask=None, position_embeddings=pe_dict[_lt],
position_ids=None, past_key_values=None)
last_call = self.tm.layers[i].call_args
self.assertIn("position_embeddings", last_call.kwargs,
"transformers 5.13.0: position_embeddings=... (positional kwarg)")
self.assertNotIn("position_embeddings_global", last_call.kwargs,
"regression: position_embeddings_global wurde in 5.13.0 deprecated")
self.assertNotIn("position_embeddings_local", last_call.kwargs,
"regression: position_embeddings_local wurde in 5.13.0 deprecated")
def test_T5_recursion_loop_uses_pedict_per_current_layer(self):
"""Recursion loop (current_layer) uses pe_dict[lt], not pe_global/pe_local."""
from px_patches.gemma4_2b_px.patch import _px_forward
import torch
cfg = self.tm.config
pe_dict = {lt: self.tm.rotary_emb(None, None, layer_type=lt)
for lt in set(cfg.layer_types)}
# Simulate a recursion iteration
for current_layer in [10, 20, 30]: # mix of sliding+full
lt = cfg.layer_types[current_layer]
self.tm.layers[current_layer](None, None, shared_kv_states=None,
attention_mask=None,
position_embeddings=pe_dict[lt],
position_ids=None, past_key_values=None)
call = self.tm.layers[current_layer].call_args
self.assertEqual(call.kwargs["position_embeddings"], pe_dict[lt])
# pe_dict[lt] must be a tuple, never None
self.assertIsInstance(call.kwargs["position_embeddings"], tuple)
# --- T6: regression detector for old kwargs in file --------------------------
class TestNoOldKwargsInPatchFile(unittest.TestCase):
"""The patch.py file must NOT contain the old position_embeddings_global/
_local kwargs — those crashed on transformers 5.13.0."""
PATCH_FILE = os.path.join(_REPO, "px_patches", "gemma4_2b_px", "patch.py")
def test_T6_no_position_embeddings_global_kwarg(self):
with open(self.PATCH_FILE, "r") as f:
content = f.read()
# Match as a kwarg in a function call, not in a comment or docstring
# Use negative-lookbehind for # (comment) — but easier: just look for the
# kwarg in a typical call pattern: `...position_embeddings_global=`
self.assertNotIn("position_embeddings_global=", content,
"regression: position_embeddings_global= kwarg in patch.py — "
"transformers 5.13.0 uses position_embeddings= positional")
def test_T7_no_position_embeddings_local_kwarg(self):
with open(self.PATCH_FILE, "r") as f:
content = f.read()
self.assertNotIn("position_embeddings_local=", content,
"regression: position_embeddings_local= kwarg in patch.py")
def test_T7b_no_pe_global_or_pe_local_variables(self):
"""pe_global/pe_local Variablen sollen komplett weg — pe_dict ersetzt sie."""
with open(self.PATCH_FILE, "r") as f:
content = f.read()
# Strip comments to allow # pe_global in einem Doc-String
code_only = re.sub(r"#.*", "", content)
# Match pe_global or pe_local as a variable assignment
self.assertNotRegex(code_only, r"\bpe_global\s*=",
"pe_global Variable existiert noch — pe_dict ersetzt sie")
self.assertNotRegex(code_only, r"\bpe_local\s*=",
"pe_local Variable existiert noch — pe_dict ersetzt sie")
def test_T7c_pedict_used_in_layer_calls(self):
"""patch.py muss position_embeddings=pe_dict[...] in Layer-Calls haben."""
with open(self.PATCH_FILE, "r") as f:
content = f.read()
# Expect at least 10 occurrences (we have 12 layer-call sites)
count = content.count("position_embeddings=pe_dict")
self.assertGreaterEqual(count, 10,
f"erwartete ≥10 position_embeddings=pe_dict[..., fanden nur {count}")
# --- T8: apply_px_patch + _safe_forward entry point --------------------------
class TestSafeForwardPeDict(unittest.TestCase):
"""_safe_forward (used for SUBJECTIVE/RIGOR presets) must also use pe_dict."""
def test_T8_safe_forward_uses_pedict(self):
from px_patches.gemma4_2b_px.patch import _safe_forward
tm = _make_mock_gemma4()
# _safe_forward is unbound — bind to mock
import inspect
self.assertTrue(inspect.isfunction(_safe_forward),
"_safe_forward should be a top-level function")
def test_T8b_safe_forward_uses_pedict_in_layer_loop(self):
"""_safe_forward's layer loop passes position_embeddings=pe_dict[_lt]."""
from px_patches.gemma4_2b_px.patch import _safe_forward
tm = _make_mock_gemma4()
import torch
# Build pe_dict the way _safe_forward does
pe_dict = {lt: tm.rotary_emb(None, None, layer_type=lt)
for lt in set(tm.config.layer_types)}
for i in [0, 10, 34]: # sliding, full, full
_lt = tm.config.layer_types[i]
tm.layers[i](None, None, shared_kv_states=None,
position_embeddings=pe_dict[_lt],
attention_mask=None, position_ids=None, past_key_values=None)
call = tm.layers[i].call_args
self.assertIn("position_embeddings", call.kwargs)
self.assertEqual(call.kwargs["position_embeddings"], pe_dict[_lt])
if __name__ == "__main__":
unittest.main(verbosity=2)
|