px-explorer-v4 / tests /test_gemma4_e2b_patch_refactor.py
BuildBot
push_hf: sparse-branch für HF-Push (nur Code, 0 LFS)
9644d0b
Raw
History Blame Contribute Delete
12.9 kB
"""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)