"""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)