J-space / tests /test_fitting.py
ayh015's picture
Upload folder using huggingface_hub
f6158c7 verified
Raw
History Blame Contribute Delete
2 kB
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)