Repository navigation
Expand file tree
/
Copy pathtest_interventions.py
More file actions
59 lines (46 loc) · 2.37 KB
/
Copy pathtest_interventions.py
File metadata and controls
59 lines (46 loc) · 2.37 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
import pytest
import torch
from thinker_r.interventions import apply_intervention
def test_normal_and_zero_do_not_mutate_input() -> None:
latents = torch.arange(24, dtype=torch.float32).reshape(2, 3, 4)
normal = apply_intervention(latents, "normal")
zero = apply_intervention(latents, "zero")
assert torch.equal(normal, latents)
assert normal.data_ptr() != latents.data_ptr()
assert torch.count_nonzero(zero) == 0
assert torch.count_nonzero(latents) > 0
def test_shuffle_preserves_each_sample_values() -> None:
latents = torch.arange(24, dtype=torch.float32).reshape(2, 3, 4)
generator = torch.Generator().manual_seed(42)
shuffled = apply_intervention(latents, "shuffle", generator=generator)
assert torch.equal(latents.sort(dim=1).values, shuffled.sort(dim=1).values)
def test_shuffle_features_preserves_values_but_changes_feature_order() -> None:
latents = torch.arange(24, dtype=torch.float32).reshape(2, 3, 4)
generator = torch.Generator().manual_seed(7)
shuffled = apply_intervention(latents, "shuffle_features", generator=generator)
assert torch.equal(latents.sort(dim=-1).values, shuffled.sort(dim=-1).values)
assert not torch.equal(latents, shuffled)
def test_cross_sample_rolls_batch() -> None:
latents = torch.arange(24, dtype=torch.float32).reshape(2, 3, 4)
swapped = apply_intervention(latents, "cross_sample")
assert torch.equal(swapped[0], latents[1])
assert torch.equal(swapped[1], latents[0])
def test_cross_sample_rejects_singleton_without_donor() -> None:
with pytest.raises(ValueError):
apply_intervention(torch.ones(1, 3, 4), "cross_sample")
def test_random_matched_is_deterministic_and_preserves_scale() -> None:
latents = torch.arange(2400, dtype=torch.float32).reshape(2, 30, 40)
first = apply_intervention(
latents, "random_matched", generator=torch.Generator().manual_seed(7)
)
second = apply_intervention(
latents, "random_matched", generator=torch.Generator().manual_seed(7)
)
assert torch.equal(first, second)
assert first.shape == latents.shape
assert torch.isfinite(first).all()
for sample in range(latents.shape[0]):
assert first[sample].mean() == pytest.approx(latents[sample].mean(), rel=0.05)
assert first[sample].std(unbiased=False) == pytest.approx(
latents[sample].std(unbiased=False), rel=0.05
)