Skip to content

Commit 367dd3c

Browse files
committed
Restore tokenizer special-id patches in attribution tests
1 parent 1bbfc6b commit 367dd3c

2 files changed

Lines changed: 38 additions & 16 deletions

File tree

tests/test_attributions_gemma.py

Lines changed: 23 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,5 @@
11
import gc
2+
from contextlib import contextmanager
23
from functools import partial
34

45
import numpy as np
@@ -187,6 +188,18 @@ def verify_intervention(
187188
verify_intervention(expected_effects, layer, pos, feature_idx, new_activation)
188189

189190

191+
@contextmanager
192+
def patch_tokenizer_special_ids(model: TransformerLensReplacementModel, special_ids: list[int]):
193+
assert model.tokenizer is not None
194+
tokenizer_class = type(model.tokenizer)
195+
original_all_special_ids = tokenizer_class.all_special_ids # type: ignore
196+
try:
197+
tokenizer_class.all_special_ids = property(lambda self: special_ids) # type: ignore
198+
yield
199+
finally:
200+
tokenizer_class.all_special_ids = original_all_special_ids # type: ignore
201+
202+
190203
def load_dummy_gemma_model(cfg: HookedTransformerConfig) -> TransformerLensReplacementModel:
191204
transcoders = {
192205
layer_idx: SingleLayerTranscoder(
@@ -206,8 +219,6 @@ def load_dummy_gemma_model(cfg: HookedTransformerConfig) -> TransformerLensRepla
206219
model = ReplacementModel.from_config(cfg, transcoder_set)
207220
assert isinstance(model, TransformerLensReplacementModel)
208221

209-
type(model.tokenizer).all_special_ids = property(lambda self: [0]) # type: ignore
210-
211222
for _, param in model.named_parameters():
212223
nn.init.uniform_(param, a=-1, b=1)
213224

@@ -291,10 +302,12 @@ def test_small_gemma_model():
291302
}
292303
cfg = HookedTransformerConfig.from_dict(gemma_small_cfg)
293304
model = load_dummy_gemma_model(cfg)
294-
graph = attribute(s, model)
295305

296-
verify_token_and_error_edges(model, graph)
297-
verify_feature_edges(model, graph)
306+
with patch_tokenizer_special_ids(model, [0]):
307+
graph = attribute(s, model)
308+
309+
verify_token_and_error_edges(model, graph)
310+
verify_feature_edges(model, graph)
298311

299312

300313
def test_large_gemma_model():
@@ -386,10 +399,12 @@ def test_large_gemma_model():
386399
}
387400
cfg = HookedTransformerConfig.from_dict(gemma_large_cfg)
388401
model = load_dummy_gemma_model(cfg)
389-
graph = attribute(s, model)
390402

391-
verify_token_and_error_edges(model, graph)
392-
verify_feature_edges(model, graph)
403+
with patch_tokenizer_special_ids(model, [0]):
404+
graph = attribute(s, model)
405+
406+
verify_token_and_error_edges(model, graph)
407+
verify_feature_edges(model, graph)
393408

394409

395410
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available")

tests/test_attributions_llama.py

Lines changed: 15 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,7 @@
1515
from circuit_tracer.transcoder.activation_functions import TopK
1616
from circuit_tracer.utils import get_default_device
1717
from tests.test_attributions_gemma import (
18+
patch_tokenizer_special_ids,
1819
verify_feature_edges,
1920
verify_token_and_error_edges,
2021
)
@@ -46,8 +47,6 @@ def load_dummy_llama_model(cfg: HookedTransformerConfig, k: int) -> TransformerL
4647
model = ReplacementModel.from_config(cfg, transcoder_set)
4748
assert model.tokenizer is not None
4849

49-
ids = model.tokenizer.all_special_ids
50-
type(model.tokenizer).all_special_ids = property(lambda self: [0] + ids) # type: ignore
5150
for _, param in model.named_parameters():
5251
nn.init.uniform_(param, a=-1, b=1)
5352

@@ -129,10 +128,14 @@ def test_small_llama_model():
129128
cfg = HookedTransformerConfig.from_dict(llama_small_cfg)
130129
k = 4
131130
model = load_dummy_llama_model(cfg, k)
132-
graph = attribute(s, model)
131+
assert model.tokenizer is not None
132+
special_ids = [0] + model.tokenizer.all_special_ids
133133

134-
verify_token_and_error_edges(model, graph)
135-
verify_feature_edges(model, graph)
134+
with patch_tokenizer_special_ids(model, special_ids):
135+
graph = attribute(s, model)
136+
137+
verify_token_and_error_edges(model, graph)
138+
verify_feature_edges(model, graph)
136139

137140

138141
def test_large_llama_model():
@@ -208,10 +211,14 @@ def test_large_llama_model():
208211
cfg = HookedTransformerConfig.from_dict(llama_large_cfg)
209212
k = 16
210213
model = load_dummy_llama_model(cfg, k)
211-
graph = attribute(s, model)
214+
assert model.tokenizer is not None
215+
special_ids = [0] + model.tokenizer.all_special_ids
212216

213-
verify_token_and_error_edges(model, graph)
214-
verify_feature_edges(model, graph)
217+
with patch_tokenizer_special_ids(model, special_ids):
218+
graph = attribute(s, model)
219+
220+
verify_token_and_error_edges(model, graph)
221+
verify_feature_edges(model, graph)
215222

216223

217224
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available")

0 commit comments

Comments
 (0)