Skip to content

Commit bea2c5c

Browse files
committed
tests (and way to mark slow tests)
1 parent dc32d79 commit bea2c5c

2 files changed

Lines changed: 84 additions & 8 deletions

File tree

pyproject.toml

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,6 @@
11
[project]
22
name = "repeng"
3-
version = "0.4.0"
3+
version = "0.5.0"
44
description = "representation engineering / control vectors"
55
authors = [{ name = "Theia Vogel", email = "theia@vgel.me" }]
66
readme = "README.md"
@@ -19,6 +19,7 @@ dev = ["pytest>=8.0.2", "ruff>=0.8.3"]
1919

2020
[tool.pytest.ini_options]
2121
python_files = ["tests.py"]
22+
markers = ["slow: marks tests as slow (deselect with '-m \"not slow\"')"]
2223

2324
[build-system]
2425
requires = ["hatchling"]

repeng/tests.py

Lines changed: 82 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -3,19 +3,14 @@
33
import pathlib
44
import tempfile
55

6+
import pytest
67
from transformers import AutoModelForCausalLM, AutoTokenizer, PreTrainedTokenizerBase
78

89
from . import ControlModel, ControlVector, DatasetEntry
910
from .control import model_layer_list
1011

1112

12-
def test_layer_list():
13-
_, gpt2 = load_gpt2_model()
14-
assert len(model_layer_list(gpt2)) == 12
15-
_, lts = load_llama_tinystories_model()
16-
assert len(model_layer_list(lts)) == 4
17-
18-
13+
@pytest.mark.slow
1914
def test_round_trip_gguf():
2015
tokenizer, model = load_llama_tinystories_model()
2116
suffixes = load_suffixes()[:50] # truncate to train vector faster
@@ -36,6 +31,7 @@ def test_round_trip_gguf():
3631
assert mushroom_cat_vector == read
3732

3833

34+
@pytest.mark.slow
3935
def test_train_gpt2():
4036
tokenizer, model = load_gpt2_model()
4137
suffixes = load_suffixes()[:50] # truncate to train vector faster
@@ -81,6 +77,7 @@ def gen(vector: ControlVector | None, strength_coeff: float | None = None):
8177
assert sad == gen(-(happy_vector * 50))
8278

8379

80+
@pytest.mark.slow
8481
def test_train_llama_tinystories():
8582
tokenizer, model = load_llama_tinystories_model()
8683
suffixes = load_suffixes()[:50] # truncate to train vector faster
@@ -119,6 +116,84 @@ def gen(vector: ControlVector | None, strength_coeff: float | None = None):
119116
assert cat.removeprefix(prompt) == " cat Bud guitar"
120117

121118

119+
@pytest.mark.slow
120+
def test_layer_list_real():
121+
# test on some real models
122+
_, gpt2 = load_gpt2_model()
123+
assert len(model_layer_list(gpt2)) == 12
124+
_, lts = load_llama_tinystories_model()
125+
assert len(model_layer_list(lts)) == 4
126+
127+
128+
def test_layer_list_override():
129+
import torch
130+
from transformers.models.llama import LlamaForCausalLM, LlamaConfig
131+
132+
fake_layers = torch.nn.ModuleList([])
133+
model = LlamaForCausalLM(
134+
LlamaConfig(vocab_size=100, hidden_size=32, intermediate_size=32)
135+
)
136+
model.repeng_layers = fake_layers
137+
138+
layers = model_layer_list(model)
139+
assert layers == fake_layers
140+
141+
142+
def test_layer_list_dummy_llama():
143+
from transformers.models.llama import LlamaForCausalLM, LlamaConfig
144+
145+
model = LlamaForCausalLM(
146+
LlamaConfig(vocab_size=100, hidden_size=32, intermediate_size=32)
147+
)
148+
layers = model_layer_list(model)
149+
assert layers == model.model.layers
150+
151+
152+
def test_layer_list_dummy_mistral():
153+
from transformers.models.mistral import MistralForCausalLM, MistralConfig
154+
155+
model = MistralForCausalLM(
156+
MistralConfig(vocab_size=100, hidden_size=32, intermediate_size=32)
157+
)
158+
layers = model_layer_list(model)
159+
assert layers == model.model.layers
160+
161+
162+
def test_layer_list_dummy_gemma():
163+
from transformers.models.gemma import GemmaForCausalLM, GemmaConfig
164+
165+
model = GemmaForCausalLM(
166+
GemmaConfig(vocab_size=100, hidden_size=32, intermediate_size=32)
167+
)
168+
layers = model_layer_list(model)
169+
assert layers == model.model.layers
170+
171+
172+
def test_layer_list_dummy_qwen():
173+
from transformers.models.qwen2 import Qwen2ForCausalLM, Qwen2Config
174+
175+
model = Qwen2ForCausalLM(
176+
Qwen2Config(vocab_size=100, hidden_size=32, intermediate_size=32)
177+
)
178+
layers = model_layer_list(model)
179+
assert layers == model.model.layers
180+
181+
182+
def test_attention_type_dummy_qwen():
183+
# tests that 'attention_type' is forwarded through getattr correctly for
184+
# qwen inference
185+
import torch
186+
from transformers.models.qwen2 import Qwen2ForCausalLM, Qwen2Config
187+
188+
model = Qwen2ForCausalLM(
189+
Qwen2Config(vocab_size=100, hidden_size=32, intermediate_size=32)
190+
)
191+
model = ControlModel(model, list(range(32)))
192+
assert model_layer_list(model)[15].attention_type == "full_attention"
193+
194+
model.forward(input_ids=torch.tensor([[0]], dtype=torch.long))
195+
196+
122197
################################################################################
123198
# Helpers
124199
################################################################################

0 commit comments

Comments
 (0)