- src/custom_models torch modules, store, training, metrics, runner and CLI (package-relative imports for pip installability) - test suite (49 tests) + conftest with fast synthetic-store fixtures - uv-managed dev env (pyproject + CPU torch index), hatchling build - README: uv for the package, pip install into conda envs for the workflow
137 lines
4.9 KiB
Python
137 lines
4.9 KiB
Python
"""Torch module acceptance tests: budget resolution, geometry and shapes."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import pytest
|
|
import torch
|
|
|
|
from custom_models.config import ModelConfig
|
|
from custom_models.models import (
|
|
ConvEncoder,
|
|
GenusProbe,
|
|
LinearAttention,
|
|
LinearAttentionEncoder,
|
|
SignalPatchEmbed,
|
|
build_probe,
|
|
count_params,
|
|
)
|
|
from custom_models.seed import seed_everything
|
|
|
|
N_CLASSES = 11
|
|
BUDGETS = (1_000_000, 5_000_000, 20_000_000, 50_000_000, 100_000_000)
|
|
|
|
|
|
@pytest.mark.parametrize("n_classes", [4, 8])
|
|
@pytest.mark.parametrize("budget", BUDGETS)
|
|
@pytest.mark.parametrize("arch", ["cnn", "linatt"])
|
|
def test_build_probe_within_budget(arch: str, budget: int, n_classes: int) -> None:
|
|
cfg = ModelConfig(arch=arch, param_budget=budget, stride=4)
|
|
probe, realised = build_probe(cfg, n_classes)
|
|
ceiling = int(1.1 * budget)
|
|
assert realised.params_realised <= ceiling
|
|
assert realised.params_realised == count_params(probe)
|
|
assert realised.d_model is not None and realised.n_layers is not None
|
|
assert realised.param_budget == budget
|
|
assert probe.head.out_features == n_classes
|
|
|
|
|
|
def test_build_probe_cnn_five_million_geometry_is_stable() -> None:
|
|
cfg = ModelConfig(arch="cnn", param_budget=5_000_000, stride=4)
|
|
probe_a, realised_a = build_probe(cfg, 11)
|
|
_probe_b, realised_b = build_probe(cfg, 11)
|
|
assert realised_a.d_model == realised_b.d_model
|
|
assert realised_a.n_layers == realised_b.n_layers
|
|
assert realised_a.params_realised == realised_b.params_realised
|
|
assert realised_a.params_realised == count_params(probe_a)
|
|
|
|
|
|
def test_build_probe_deterministic_state() -> None:
|
|
seed_everything(0)
|
|
cfg = ModelConfig(arch="cnn", param_budget=1_000_000, stride=4)
|
|
probe_a, _ = build_probe(cfg, 4)
|
|
seed_everything(0)
|
|
probe_b, _ = build_probe(cfg, 4)
|
|
for (a_name, a_val), (b_name, b_val) in zip(
|
|
probe_a.state_dict().items(), probe_b.state_dict().items()
|
|
):
|
|
assert a_name == b_name
|
|
assert torch.equal(a_val, b_val)
|
|
|
|
|
|
def test_build_probe_no_candidate_raises() -> None:
|
|
cfg = ModelConfig(arch="cnn", param_budget=100, stride=4)
|
|
with pytest.raises(ValueError, match=r"no .+ candidate"):
|
|
build_probe(cfg, 4)
|
|
|
|
|
|
def test_genus_probe_rejects_single_class() -> None:
|
|
encoder = ConvEncoder(16, 1, 4, 4, 8)
|
|
with pytest.raises(ValueError, match="n_classes must be >= 2"):
|
|
GenusProbe(encoder, 8, 1)
|
|
|
|
|
|
@pytest.mark.parametrize(("stride", "expected_tokens"), [(4, 3000), (8, 1500)])
|
|
def test_patch_embed_token_count(stride: int, expected_tokens: int) -> None:
|
|
embed = SignalPatchEmbed(d_model=16, patch_len=min(4, stride), stride=stride)
|
|
x = torch.zeros(2, 12_000)
|
|
tokens = embed(x)
|
|
assert tokens.shape == (2, expected_tokens, 16)
|
|
|
|
|
|
@pytest.mark.parametrize("stride", [4, 8])
|
|
def test_forward_shapes_and_finiteness(stride: int) -> None:
|
|
cfg = ModelConfig(arch="cnn", param_budget=100_000, stride=stride)
|
|
probe, realised = build_probe(cfg, N_CLASSES)
|
|
x = torch.zeros(2, 1200)
|
|
logits = probe(x)
|
|
assert logits.shape == (2, N_CLASSES)
|
|
assert torch.isfinite(logits).all()
|
|
emb = probe.embed(x)
|
|
assert emb.shape == (2, realised.d_embed)
|
|
|
|
|
|
def test_encoder_rejects_indivisible_window() -> None:
|
|
embed = SignalPatchEmbed(d_model=16, patch_len=4, stride=8)
|
|
with pytest.raises(ValueError, match="not divisible"):
|
|
embed(torch.zeros(1, 12_002))
|
|
|
|
|
|
def test_no_batchnorm_and_uniform_pool_at_init() -> None:
|
|
cfg = ModelConfig(arch="cnn", param_budget=100_000, stride=4)
|
|
probe, _ = build_probe(cfg, N_CLASSES)
|
|
assert not any("BatchNorm" in type(m).__name__ for m in probe.modules())
|
|
encoder = probe.encoder
|
|
assert isinstance(encoder, ConvEncoder)
|
|
assert torch.count_nonzero(encoder.score.weight) == 0
|
|
assert torch.count_nonzero(encoder.score.bias) == 0
|
|
|
|
|
|
def test_conv_pool_starts_as_plain_mean() -> None:
|
|
encoder = ConvEncoder(d_model=16, n_layers=1, patch_len=4, stride=4, d_embed=8)
|
|
x = torch.randn(2, 1200)
|
|
h = encoder.embed(x)
|
|
for block in encoder.blocks:
|
|
h = block(h)
|
|
expected = encoder.proj(h.mean(dim=1))
|
|
assert torch.allclose(encoder(x), expected, atol=1e-5)
|
|
scores = torch.softmax(encoder.score(h), dim=1)
|
|
assert torch.allclose(scores, torch.full_like(scores, 1.0 / scores.shape[1]))
|
|
|
|
|
|
def test_linear_attention_heads_guard() -> None:
|
|
with pytest.raises(ValueError, match="divisible"):
|
|
LinearAttention(d_model=32, n_heads=5)
|
|
with pytest.raises(ValueError, match="n_heads must be >= 1"):
|
|
LinearAttention(d_model=32, n_heads=0)
|
|
|
|
|
|
def test_linear_attention_encoder_mean_pool() -> None:
|
|
encoder = LinearAttentionEncoder(16, 1, 4, 4, 8, n_heads=4)
|
|
x = torch.randn(2, 1200)
|
|
out = encoder(x)
|
|
assert out.shape == (2, 8)
|
|
h = encoder.embed(x)
|
|
for block in encoder.blocks:
|
|
h = block(h)
|
|
assert torch.allclose(encoder.proj(h.mean(dim=1)), out, atol=1e-5)
|