| from types import SimpleNamespace |
|
|
| import torch |
| from torch import nn |
|
|
| from math_jlens.fitting import jacobian_for_tokens, valid_position_mask |
| from math_jlens.corpus import response_window_start |
|
|
|
|
| class Block(nn.Module): |
| def __init__(self, width: int) -> None: |
| super().__init__() |
| self.linear = nn.Linear(width, width, bias=False) |
| with torch.no_grad(): |
| self.linear.weight.mul_(0.1) |
|
|
| def forward(self, hidden): |
| return hidden + self.linear(hidden) |
|
|
|
|
| class TinyModel(nn.Module): |
| def __init__(self, layers=4, width=8, vocab=32) -> None: |
| super().__init__() |
| torch.manual_seed(0) |
| self.n_layers = layers |
| self.d_model = width |
| self.embedding = nn.Embedding(vocab, width) |
| self.layers = nn.ModuleList(Block(width) for _ in range(layers)) |
| for parameter in self.parameters(): |
| parameter.requires_grad_(False) |
|
|
| def forward(self, input_ids): |
| hidden = self.embedding(input_ids) |
| for block in self.layers: |
| hidden = block(hidden) |
| return SimpleNamespace(last_hidden_state=hidden) |
|
|
|
|
| def test_position_mask(): |
| mask = valid_position_mask(12, skip_first=3) |
| assert mask.sum() == 8 |
| assert not mask[:3].any() |
| assert not mask[-1] |
|
|
|
|
| def test_four_stratified_response_windows(): |
| starts = [response_window_start(4096, 1024, stage) for stage in range(4)] |
| assert starts == [0, 1024, 2048, 3072] |
|
|
|
|
| def test_jacobian_shape_orientation_and_layers(): |
| model = TinyModel() |
| tokens = torch.arange(20).remainder(32).unsqueeze(0) |
| matrices, valid = jacobian_for_tokens( |
| model, tokens, source_layers=[0, 1, 2], target_layer=3, |
| dim_batch=4, skip_first=2, |
| ) |
| assert valid == 17 |
| assert set(matrices) == {0, 1, 2} |
| assert all(value.shape == (8, 8) for value in matrices.values()) |
| expected = torch.eye(8) + model.layers[3].linear.weight.detach() |
| torch.testing.assert_close(matrices[2], expected, rtol=0, atol=1e-5) |
|
|