import tempfile import unittest from pathlib import Path from types import SimpleNamespace import torch from owmi.backends import HFBackend, TransformerLensBackend from owmi.benchmarks.runner import BenchmarkRunner from owmi.interventions import ( LinearSAE, apply_intervention, apply_sae_feature_intervention, load_linear_sae, make_forward_hook, make_head_prehook, ) from owmi.types import ExperimentConfig, InterventionSpec from owmi.validation import ( head_slice_isolation_report, validate_head_intervention, validate_sae_roundtrip, ) class _TinyTokenizer: pad_token = eos_token = '' eos_token_id = 0 def apply_chat_template(self, messages, tokenize=False, add_generation_prompt=True): return 'x' def __call__(self, prompt, return_tensors='pt'): class Encoded(dict): def to(self, device): return self return Encoded(input_ids=torch.tensor([[1, 2, 3]])) class _TinyBlock(torch.nn.Module): """Decoder block with a synthetic attention layout: o_proj input == block hidden.""" def __init__(self, width): super().__init__() attn = torch.nn.Module() attn.o_proj = torch.nn.Linear(width, width, bias=False) with torch.no_grad(): attn.o_proj.weight.copy_(torch.eye(width)) self.self_attn = attn def forward(self, hidden): return (self.self_attn.o_proj(hidden),) class _TinyAttnModel(torch.nn.Module): def __init__(self, n_heads=2, head_dim=2, kv_heads=None): super().__init__() self.width = n_heads * head_dim self.config = SimpleNamespace( num_attention_heads=n_heads, num_key_value_heads=kv_heads if kv_heads is not None else n_heads, head_dim=head_dim, ) inner = torch.nn.Module() inner.layers = torch.nn.ModuleList([_TinyBlock(self.width)]) self.model = inner def hidden_for(self, seq_len): return torch.arange(seq_len * self.width, dtype=torch.float32).reshape(1, seq_len, self.width) + 1.0 def forward(self, input_ids, use_cache=False): hidden = self.hidden_for(int(input_ids.shape[1])) out = self.model.layers[0](hidden)[0] return type('Output', (), {'logits': out})() def _hf_backend(model): backend = HFBackend.__new__(HFBackend) backend.model = model backend.config = ExperimentConfig(model_name='fake') backend.tokenizer = _TinyTokenizer() backend.device = 'cpu' return backend class _FakeBackend: def __init__(self): self.intervention_calls = [] self.generated = [] def generate(self, prompt): self.generated.append(prompt) return '{"detected": true, "confidence": 1}' def run_with_intervention(self, prompt, intervention): self.intervention_calls.append((prompt, intervention)) return '{"detected": true, "confidence": 1}' def output_divergence(self, prompt, intervention): return {"mean_js": 0.25, "max_js": 0.5, "n_positions": 3.0} def _runner_row(obj, output_dir='test-results'): return { 'run_id': 'routing-run', 'suite_name': 'test', 'model': {'name': 'fake', 'hook_temporal_scope': 'active_through_probe'}, 'benchmark': {'name': 'toy', 'schema': 'multiple_choice'}, 'benchmark_example': { 'benchmark': 'toy', 'item_id': '1', 'schema': 'multiple_choice', 'prompt': 'Question?', 'reference': {'answer_letter': 'A'}, 'metadata': {}, }, 'object': obj, 'intervention': {'kind': 'intervention', 'condition': 'intervention', 'mode': 'zero', 'strength': 1.0}, 'probe': {'task': 'detection', 'condition': 'model'}, 'seed': 7, 'output_dir': output_dir, } def _identity_sae_state(width=4): return { 'W_enc': torch.eye(width), 'b_enc': torch.zeros(width), 'W_dec': torch.eye(width), 'b_dec': torch.zeros(width), } class ChannelSliceOperatorTests(unittest.TestCase): def test_zero_restricted_to_channel_window_and_positions(self): hidden = torch.ones((1, 3, 4)) spec = InterventionSpec(layer_index=0, token_positions=[0, 1], mode='zero') out = apply_intervention(hidden, spec, channel_slice=(2, 4)) self.assertTrue(torch.equal(out[:, :2, 2:4], torch.zeros((1, 2, 2)))) self.assertTrue(torch.equal(out[:, :2, 0:2], torch.ones((1, 2, 2)))) self.assertTrue(torch.equal(out[:, 2], hidden[:, 2])) def test_scale_noise_replace_random_leave_off_window_untouched(self): torch.manual_seed(0) hidden = torch.randn((1, 3, 4)) for mode, reference_mode in (('scale', 'none'), ('noise', 'none'), ('replace', 'random')): spec = InterventionSpec(layer_index=0, mode=mode, strength=0.5, reference_mode=reference_mode) out = apply_intervention(hidden, spec, channel_slice=(0, 2)) self.assertTrue(torch.equal(out[..., 2:], hidden[..., 2:]), mode) self.assertFalse(torch.equal(out[..., :2], hidden[..., :2]), mode) def test_channel_slice_out_of_range_raises(self): hidden = torch.ones((1, 3, 4)) spec = InterventionSpec(layer_index=0, mode='zero') with self.assertRaises(ValueError): apply_intervention(hidden, spec, channel_slice=(2, 5)) class HeadPrehookTests(unittest.TestCase): def test_head_prehook_modifies_prefill_slice_but_never_decode_steps(self): spec = InterventionSpec(layer_index=0, mode='zero', object_kind='attention_head', head_index=1) prehook = make_head_prehook(spec, head_dim=2, prompt_length=3) prefill = torch.ones((1, 3, 4)) modified = prehook(torch.nn.Identity(), (prefill,))[0] self.assertTrue(torch.equal(modified[..., 2:4], torch.zeros((1, 3, 2)))) self.assertTrue(torch.equal(modified[..., 0:2], torch.ones((1, 3, 2)))) decode = torch.ones((1, 1, 4)) self.assertIsNone(prehook(torch.nn.Identity(), (decode,))) def test_head_prehook_requires_head_index(self): spec = InterventionSpec(layer_index=0, mode='zero', object_kind='attention_head') with self.assertRaises(ValueError): make_head_prehook(spec, head_dim=2) class HFHeadRoutingTests(unittest.TestCase): def test_head_zero_changes_only_target_head_slice_under_gqa(self): backend = _hf_backend(_TinyAttnModel(n_heads=4, head_dim=2, kv_heads=2)) spec = InterventionSpec(layer_index=0, mode='zero', object_kind='attention_head', head_index=3) report = validate_head_intervention(backend, 'x', spec) self.assertEqual(report['head_dim'], 2) self.assertEqual(report['head_slice'], [6, 8]) self.assertTrue(report['changed']) self.assertTrue(report['isolated']) def test_head_index_addresses_query_heads_not_kv_heads(self): backend = _hf_backend(_TinyAttnModel(n_heads=4, head_dim=2, kv_heads=2)) spec = InterventionSpec(layer_index=0, mode='zero', object_kind='attention_head', head_index=2) report = validate_head_intervention(backend, 'x', spec) self.assertEqual(report['head_slice'], [4, 6]) out_of_range = InterventionSpec(layer_index=0, mode='zero', object_kind='attention_head', head_index=4) with self.assertRaises(ValueError): validate_head_intervention(backend, 'x', out_of_range) def test_head_replace_baseline_reproduces_baseline_at_site(self): backend = _hf_backend(_TinyAttnModel()) spec = InterventionSpec( layer_index=0, mode='replace', reference_mode='baseline', object_kind='attention_head', head_index=0, ) report = validate_head_intervention(backend, 'x', spec) self.assertFalse(report['changed']) self.assertTrue(report['isolated']) def test_head_output_divergence_routes_through_head_hook(self): backend = _hf_backend(_TinyAttnModel()) spec = InterventionSpec(layer_index=0, mode='zero', object_kind='attention_head', head_index=1) result = backend.output_divergence('x', spec) self.assertEqual(result['n_positions'], 3.0) self.assertGreater(result['mean_js'], 0.0) def test_head_slice_isolation_report_flags_off_slice_changes(self): baseline = torch.zeros((1, 2, 4)) leaked = torch.zeros((1, 2, 4)) leaked[..., 0] = 1.0 report = head_slice_isolation_report(baseline, leaked, head_index=1, head_dim=2) self.assertFalse(report['changed']) self.assertFalse(report['isolated']) class SAEFeatureTests(unittest.TestCase): def setUp(self): tmpdir = tempfile.TemporaryDirectory() self.addCleanup(tmpdir.cleanup) self.sae_path = str(Path(tmpdir.name) / 'sae.pt') torch.save(_identity_sae_state(4), self.sae_path) self.hidden = torch.arange(12, dtype=torch.float32).reshape(1, 3, 4) + 1.0 def _spec(self, mode='zero', **kwargs): defaults = dict( layer_index=0, mode=mode, object_kind='sae_feature', feature_id=2, sae_weights_path=self.sae_path, ) defaults.update(kwargs) return InterventionSpec(**defaults) def test_load_linear_sae_validates_keys_and_normalizes_transposed_weights(self): bad_path = str(Path(self.sae_path).parent / 'bad.pt') state = _identity_sae_state(4) del state['W_dec'] torch.save(state, bad_path) with self.assertRaises(ValueError): load_linear_sae(bad_path) transposed_path = str(Path(self.sae_path).parent / 'transposed.pt') torch.save({ 'W_enc': torch.eye(4)[:3, :], # [d_sae, d_model], must be normalized 'b_enc': torch.zeros(3), 'W_dec': torch.eye(4)[:, :3], # [d_model, d_sae], must be normalized 'b_dec': torch.zeros(4), }, transposed_path) sae = load_linear_sae(transposed_path) self.assertEqual((sae.d_model, sae.d_sae), (4, 3)) self.assertEqual(tuple(sae.W_enc.shape), (4, 3)) self.assertEqual(tuple(sae.W_dec.shape), (3, 4)) def test_zero_feature_changes_only_that_feature_at_target_positions(self): sae = load_linear_sae(self.sae_path) spec = self._spec(token_positions=[0, 2]) out = apply_sae_feature_intervention(self.hidden, spec, sae, position_seq_len=3) self.assertTrue(torch.equal(out[:, [0, 2], 2], torch.zeros((1, 2)))) self.assertTrue(torch.equal(out[:, 1], self.hidden[:, 1])) for channel in (0, 1, 3): self.assertTrue(torch.equal(out[..., channel], self.hidden[..., channel])) def test_scale_at_strength_one_is_exact_identity(self): sae = load_linear_sae(self.sae_path) out = apply_sae_feature_intervention(self.hidden, self._spec(mode='scale', strength=1.0), sae) self.assertTrue(torch.equal(out, self.hidden)) def test_replace_baseline_uses_reference_feature_and_requires_reference(self): sae = load_linear_sae(self.sae_path) spec = self._spec(mode='replace', reference_mode='baseline') reference = self.hidden * 2.0 out = apply_sae_feature_intervention(self.hidden, spec, sae, reference_hidden=reference) self.assertTrue(torch.equal(out[..., 2], reference[..., 2])) for channel in (0, 1, 3): self.assertTrue(torch.equal(out[..., channel], self.hidden[..., channel])) with self.assertRaises(ValueError): apply_sae_feature_intervention(self.hidden, spec, sae) def test_feature_id_out_of_range_raises(self): sae = load_linear_sae(self.sae_path) with self.assertRaises(ValueError): apply_sae_feature_intervention(self.hidden, self._spec(feature_id=4), sae) def test_forward_hook_sae_path_is_prefill_only(self): sae = load_linear_sae(self.sae_path) hook = make_forward_hook(self._spec(), prompt_length=3, sae=sae) modified = hook(torch.nn.Identity(), (), (self.hidden,))[0] self.assertTrue(torch.equal(modified[..., 2], torch.zeros((1, 3)))) decode = torch.ones((1, 1, 4)) unchanged = hook(torch.nn.Identity(), (), (decode,))[0] self.assertTrue(torch.equal(unchanged, decode)) def test_hf_backend_routes_sae_hook_to_block_and_modifies_only_feature(self): model = _TinyAttnModel(n_heads=2, head_dim=2) backend = _hf_backend(model) handle = backend._register_intervention_hook(self._spec(), prompt_length=3) module = backend.resolve_block_module(0) captured = {} def capture(_, __, output): captured['tensor'] = output[0].detach().clone() return None capture_handle = module.register_forward_hook(capture) try: with torch.no_grad(): _ = backend.model(input_ids=torch.tensor([[1, 2, 3]]), use_cache=False) finally: capture_handle.remove() handle.remove() baseline = model.hidden_for(3) self.assertTrue(torch.equal(captured['tensor'][..., 2], torch.zeros((1, 3)))) for channel in (0, 1, 3): self.assertTrue(torch.equal(captured['tensor'][..., channel], baseline[..., channel])) def test_sae_roundtrip_validation_identity_and_lossy(self): identity = validate_sae_roundtrip(self.sae_path, self.hidden, self._spec()) self.assertEqual(identity['reconstruction_mse'], 0.0) self.assertTrue(identity['feature_delta_confirmed']) self.assertTrue(identity['delta_in_feature_direction']) lossy = LinearSAE( W_enc=torch.eye(4)[:, :3], b_enc=torch.zeros(3), W_dec=torch.eye(4)[:3, :], b_dec=torch.zeros(4), ) report = validate_sae_roundtrip(lossy, self.hidden, self._spec(feature_id=1)) self.assertGreater(report['reconstruction_mse'], 0.0) self.assertTrue(report['feature_delta_confirmed']) self.assertTrue(report['delta_in_feature_direction']) class RunnerRoutingTests(unittest.TestCase): def setUp(self): tmpdir = tempfile.TemporaryDirectory() self.addCleanup(tmpdir.cleanup) self.tmp = Path(tmpdir.name) self.sae_path = str(self.tmp / 'sae.pt') torch.save(_identity_sae_state(4), self.sae_path) def test_core_intervention_populates_head_routing_fields(self): obj = {'kind': 'attention_head', 'layer_index': 1, 'head_index': 2, 'token_positions': [0]} spec = BenchmarkRunner(_runner_row(obj))._core_intervention() self.assertEqual(spec.object_kind, 'attention_head') self.assertEqual(spec.head_index, 2) spec.validate() def test_core_intervention_reads_sae_path_from_direction_path_or_metadata(self): for obj in ( {'kind': 'sae_feature', 'layer_index': 1, 'feature_id': 3, 'direction_path': self.sae_path}, {'kind': 'sae_feature', 'layer_index': 1, 'feature_id': 3, 'metadata': {'sae_weights_path': self.sae_path}}, ): spec = BenchmarkRunner(_runner_row(obj))._core_intervention() self.assertEqual(spec.object_kind, 'sae_feature') self.assertEqual(spec.feature_id, 3) self.assertEqual(spec.sae_weights_path, self.sae_path) spec.validate() def test_unroutable_object_configs_still_raise_not_implemented(self): cases = [ {'kind': 'attention_head', 'layer_index': 1}, {'kind': 'sae_feature', 'layer_index': 1}, {'kind': 'sae_feature', 'layer_index': 1, 'feature_id': 0}, ] for obj in cases: with self.assertRaises(NotImplementedError): BenchmarkRunner(_runner_row(obj))._core_intervention() with self.assertRaisesRegex(NotImplementedError, 'direction_path'): BenchmarkRunner(_runner_row(cases[2]))._core_intervention() def test_runner_end_to_end_head_row_with_fake_backend(self): backend = _FakeBackend() obj = {'kind': 'attention_head', 'layer_index': 1, 'head_index': 2, 'token_positions': [0]} result = BenchmarkRunner(_runner_row(obj, output_dir=str(self.tmp / 'out'))).run(backend=backend) self.assertEqual(result.extra['intervention_spec']['object_kind'], 'attention_head') self.assertEqual(result.extra['intervention_spec']['head_index'], 2) for _, spec in backend.intervention_calls: self.assertEqual(spec.object_kind, 'attention_head') self.assertEqual(spec.head_index, 2) def test_runtime_spec_validation_for_object_kinds(self): with self.assertRaises(ValueError): InterventionSpec(layer_index=0, object_kind='attention_head').validate() with self.assertRaises(ValueError): InterventionSpec(layer_index=0, object_kind='sae_feature', feature_id=1).validate() InterventionSpec( layer_index=0, object_kind='sae_feature', feature_id=1, sae_weights_path=self.sae_path ).validate() def test_transformerlens_backend_rejects_head_and_sae_kinds(self): backend = TransformerLensBackend.__new__(TransformerLensBackend) backend.model = object() backend.config = ExperimentConfig(model_name='fake') specs = [ InterventionSpec(layer_index=0, object_kind='attention_head', head_index=0), InterventionSpec(layer_index=0, object_kind='sae_feature', feature_id=0, sae_weights_path=self.sae_path), ] for spec in specs: with self.assertRaises(NotImplementedError): backend.run_with_intervention('x', spec) with self.assertRaises(NotImplementedError): backend.output_divergence('x', spec) if __name__ == '__main__': unittest.main()