Initial release: E1 genus probe package moved out of the E01 workflow
- 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
This commit is contained in:
@@ -0,0 +1,50 @@
|
||||
"""Shared fixtures: fast synthetic stores (no pod5 needed) and determinism."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
|
||||
|
||||
from custom_models.seed import seed_everything
|
||||
from custom_models.store import synthetic_store
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _seed_every_test() -> None:
|
||||
"""Seed all RNGs before each test (the GATE ZERO §2.3 contract)."""
|
||||
seed_everything(0)
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def fast_store(tmp_path_factory: pytest.TempPathFactory) -> Path:
|
||||
"""A small 6-genus duty-cycle store for split/config tests."""
|
||||
out = tmp_path_factory.mktemp("store") / "main"
|
||||
synthetic_store(
|
||||
out_dir=out,
|
||||
n_genera=6,
|
||||
reads_per_genus=12,
|
||||
window_samples=1200,
|
||||
shard_windows=8,
|
||||
seed=0,
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
@pytest.fixture(scope="session")
|
||||
def trap_store(tmp_path_factory: pytest.TempPathFactory) -> Path:
|
||||
"""A small eval-only impostor store (and test-split source)."""
|
||||
out = tmp_path_factory.mktemp("store") / "trap"
|
||||
synthetic_store(
|
||||
out_dir=out,
|
||||
n_genera=2,
|
||||
reads_per_genus=6,
|
||||
window_samples=1200,
|
||||
shard_windows=8,
|
||||
seed=1,
|
||||
genus_prefix="trap_",
|
||||
)
|
||||
return out
|
||||
@@ -0,0 +1,136 @@
|
||||
"""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)
|
||||
@@ -0,0 +1,416 @@
|
||||
"""Pipeline acceptance tests: store, splits, metrics, training and the CLI."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import numpy as np
|
||||
import polars as pl
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from custom_models.cli import main as cli_main
|
||||
from custom_models.config import (
|
||||
DataConfig,
|
||||
ModelConfig,
|
||||
RunConfig,
|
||||
TrainConfig,
|
||||
load_config,
|
||||
run_id_for,
|
||||
save_config,
|
||||
)
|
||||
from custom_models.data import make_splits, permute_read_labels
|
||||
from custom_models.metrics import choose_tau_msp, closed_set_report, msp
|
||||
from custom_models.runner import evaluate_test, evaluate_val, train_run
|
||||
from custom_models.seed import seed_everything
|
||||
from custom_models.store import SignalStore, synthetic_store
|
||||
from custom_models.train import select_device
|
||||
|
||||
# ---------------------------------------------------------------- store
|
||||
|
||||
|
||||
def test_synthetic_store_roundtrip(tmp_path: Path) -> None:
|
||||
out = tmp_path / "s"
|
||||
synthetic_store(
|
||||
out_dir=out,
|
||||
n_genera=3,
|
||||
reads_per_genus=4,
|
||||
window_samples=1200,
|
||||
shard_windows=4,
|
||||
seed=0,
|
||||
)
|
||||
store = SignalStore(out)
|
||||
assert len(store) == 3 * 4
|
||||
assert store.window_len() == 1200
|
||||
meta = store.meta()
|
||||
assert meta["synthetic"] is True
|
||||
assert meta["n_genera"] == 3
|
||||
assert meta["n_windows"] == len(store)
|
||||
windows = [store.get(i) for i in range(len(store))]
|
||||
assert all(w.shape == (1200,) and w.dtype == np.float32 for w in windows)
|
||||
manifest = store.manifest()
|
||||
assert set(manifest.columns) >= {
|
||||
"window_id",
|
||||
"read_id",
|
||||
"genus",
|
||||
"species",
|
||||
"source_run",
|
||||
"window_index",
|
||||
"n_samples",
|
||||
}
|
||||
assert manifest["window_id"].to_list() == list(range(len(store)))
|
||||
genera = manifest["genus"].unique().sort().to_list()
|
||||
assert genera == [f"genus_{i:02d}" for i in range(3)]
|
||||
|
||||
|
||||
def test_store_labels_and_split_determinism(fast_store: Path) -> None:
|
||||
store = SignalStore(fast_store)
|
||||
manifest = store.manifest()
|
||||
selected = ("genus_00", "genus_01", "genus_02")
|
||||
cfg = DataConfig(genera=selected, min_test_windows_per_genus=0, seed=0)
|
||||
a = make_splits(manifest, cfg)
|
||||
b = make_splits(manifest, cfg)
|
||||
assert all(
|
||||
np.array_equal(x, y) for x, y in zip((a.train, a.val, a.test), (b.train, b.val, b.test))
|
||||
)
|
||||
assert a.genera == list(selected)
|
||||
# read-level purity: every window of a read shares one split
|
||||
read_arr = manifest["read_id"].to_numpy()
|
||||
wid_arr = manifest["window_id"].to_numpy()
|
||||
split_of = {int(w): 0 for w in a.train}
|
||||
split_of.update({int(w): 1 for w in a.val})
|
||||
split_of.update({int(w): 2 for w in a.test})
|
||||
by_read: dict[str, set[int]] = {}
|
||||
for wid, read_id in zip(wid_arr, read_arr):
|
||||
if int(wid) in split_of: # windows of unselected genera have no split
|
||||
by_read.setdefault(str(read_id), set()).add(split_of[int(wid)])
|
||||
assert all(len(v) == 1 for v in by_read.values())
|
||||
# a different data seed re-shuffles the read assignment
|
||||
c = make_splits(manifest, DataConfig(genera=selected, min_test_windows_per_genus=0, seed=1))
|
||||
assert not np.array_equal(a.train, c.train)
|
||||
# each read keeps its genus label; class ids map onto splits.genera
|
||||
assert len(a.labels) == len(manifest)
|
||||
assert (a.labels[a.train] >= 0).all()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- splits
|
||||
|
||||
|
||||
def test_make_splits_min_test_windows_guard(trap_store: Path) -> None:
|
||||
store = SignalStore(trap_store)
|
||||
with pytest.raises(ValueError, match="min_test_windows_per_genus"):
|
||||
make_splits(
|
||||
store.manifest(),
|
||||
DataConfig(genera=("trap_00", "trap_01"), min_test_windows_per_genus=10**6),
|
||||
)
|
||||
|
||||
|
||||
def test_make_splits_holdout_runs_guard(fast_store: Path) -> None:
|
||||
store = SignalStore(fast_store)
|
||||
manifest = store.manifest()
|
||||
# two source_runs, alternating per read: holdout picks news_run_*1 for test
|
||||
runs = [
|
||||
"news_run_0" if i % 2 == 0 else "news_run_1"
|
||||
for i, read_id in enumerate(manifest["read_id"].unique().sort().to_list())
|
||||
]
|
||||
split_map = dict(zip(manifest["read_id"].unique().sort().to_list(), runs))
|
||||
two_runs = manifest.with_columns(
|
||||
pl.col("read_id").replace_strict(split_map).alias("source_run")
|
||||
)
|
||||
assert (
|
||||
len(make_splits(two_runs, DataConfig(holdout_runs=True, min_test_windows_per_genus=1)).test)
|
||||
> 0
|
||||
)
|
||||
single_run = manifest.with_columns(pl.lit("only_run").alias("source_run"))
|
||||
with pytest.raises(ValueError, match="holdout_runs"):
|
||||
make_splits(single_run, DataConfig(holdout_runs=True, min_test_windows_per_genus=0))
|
||||
|
||||
|
||||
def test_permute_read_labels_preserves_counts(fast_store: Path) -> None:
|
||||
store = SignalStore(fast_store)
|
||||
manifest = store.manifest()
|
||||
pd = permute_read_labels(manifest, (), 3)
|
||||
counts_before = manifest["genus"].value_counts().sort("genus")
|
||||
counts_after = pd["genus"].value_counts().sort("genus")
|
||||
assert counts_before.equals(counts_after)
|
||||
# windows keep their read's label: read -> single genus before and after
|
||||
for frame in (manifest, pd):
|
||||
per_read = frame.group_by("read_id").agg(pl.col("genus").n_unique().alias("k"))
|
||||
assert per_read["k"].max() == 1
|
||||
pd2 = permute_read_labels(manifest, (), 3)
|
||||
assert pd.equals(pd2)
|
||||
with pytest.raises(ValueError, match="at least two reads"):
|
||||
permute_read_labels(manifest.filter(pl.col("read_id") == manifest["read_id"][0]), (), 0)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- configs
|
||||
|
||||
|
||||
def test_config_roundtrip(tmp_path: Path) -> None:
|
||||
cfg = RunConfig(
|
||||
run_id="cnn_1.0M_s4_g11",
|
||||
stage="arch_ladder",
|
||||
store_dir=Path("/data/store"),
|
||||
data=DataConfig(genera=("A", "B"), val_fraction=0.2, seed=7),
|
||||
model=ModelConfig(
|
||||
arch="cnn",
|
||||
param_budget=1_000_000,
|
||||
stride=4,
|
||||
d_model=64,
|
||||
n_layers=8,
|
||||
params_realised=912_345,
|
||||
),
|
||||
train=TrainConfig(lr=1e-3, seed=3),
|
||||
out_dir=Path("/runs"),
|
||||
genera=("A", "B"),
|
||||
notes="hello",
|
||||
)
|
||||
path = tmp_path / "config.json"
|
||||
save_config(cfg, path)
|
||||
assert load_config(path) == cfg
|
||||
|
||||
|
||||
def test_run_id_format() -> None:
|
||||
assert run_id_for("cnn", 1_000_000, 4, 11, "arch_ladder") == "cnn_1.0M_s4_g11"
|
||||
assert run_id_for("linatt", 5_000_000, 8, 6, "windows_ablation") == "linatt_5.0M_s8_g6"
|
||||
assert run_id_for("cnn", 1_000_000, 4, 11, "control") == "cnn_1.0M_s4_g11-shuf"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- metrics
|
||||
|
||||
|
||||
def test_msp_and_choose_tau() -> None:
|
||||
logits = np.array([[10.0, 0.0], [0.0, 10.0], [1.0, 1.0]])
|
||||
values = msp(logits)
|
||||
assert np.all(values >= 0.5) and values[2] == pytest.approx(0.5)
|
||||
with pytest.raises(ValueError):
|
||||
msp(np.array([1.0, 0.0]))
|
||||
with pytest.raises(ValueError):
|
||||
choose_tau_msp(np.array([]))
|
||||
arr = np.linspace(0.1, 0.9, 9)
|
||||
tau = choose_tau_msp(arr, target_known_recall=0.95)
|
||||
assert tau == pytest.approx(np.quantile(arr, 0.05))
|
||||
|
||||
|
||||
def test_closed_set_report_perfect_predictions() -> None:
|
||||
y = np.array([0, 0, 1, 1, 2])
|
||||
report = closed_set_report(y, y.copy(), ["a", "b", "c"], n_boot=50, seed=0)
|
||||
assert report.recall_macro == pytest.approx(1.0)
|
||||
assert report.recall_micro == pytest.approx(1.0)
|
||||
assert report.f1_macro == pytest.approx(1.0)
|
||||
assert (report.confusion == np.diag(np.diag(report.confusion))).all()
|
||||
assert report.ci["recall_macro"] == (1.0, 1.0)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- training
|
||||
|
||||
|
||||
def test_trainability_duty_cycle(tmp_path: Path) -> None:
|
||||
seed_everything(0)
|
||||
store_dir = tmp_path / "s"
|
||||
synthetic_store(
|
||||
store_dir, n_genera=2, reads_per_genus=12, window_samples=1200, shard_windows=12
|
||||
)
|
||||
cfg = RunConfig(
|
||||
run_id="",
|
||||
stage="arch_ladder",
|
||||
store_dir=store_dir,
|
||||
data=DataConfig(
|
||||
genera=(), val_fraction=0.25, test_fraction=0.25, min_test_windows_per_genus=1, seed=0
|
||||
),
|
||||
model=ModelConfig(arch="cnn", param_budget=100_000, stride=4),
|
||||
train=TrainConfig(
|
||||
max_epochs=25,
|
||||
patience=25,
|
||||
num_workers=0,
|
||||
amp="off",
|
||||
device="cpu",
|
||||
batch_size=8,
|
||||
),
|
||||
out_dir=tmp_path / "runs",
|
||||
)
|
||||
realised, _result, _splits = train_run(cfg)
|
||||
run_dir = realised.out_dir / realised.run_id
|
||||
assert (run_dir / "history.parquet").is_file()
|
||||
val_payload = evaluate_val(run_dir, n_boot=0)
|
||||
assert val_payload["closed"]["recall_macro"] >= 0.9
|
||||
|
||||
|
||||
def test_control_run_lands_at_chance(tmp_path: Path) -> None:
|
||||
seed_everything(1)
|
||||
store_dir = tmp_path / "s"
|
||||
synthetic_store(
|
||||
store_dir, n_genera=3, reads_per_genus=12, window_samples=1200, shard_windows=12
|
||||
)
|
||||
cfg = RunConfig(
|
||||
run_id="",
|
||||
stage="control",
|
||||
store_dir=store_dir,
|
||||
data=DataConfig(genera=(), val_fraction=0.25, min_test_windows_per_genus=1, seed=0),
|
||||
model=ModelConfig(arch="cnn", param_budget=100_000, stride=4),
|
||||
train=TrainConfig(
|
||||
max_epochs=4, patience=4, num_workers=0, amp="off", device="cpu", batch_size=8
|
||||
),
|
||||
out_dir=tmp_path / "runs",
|
||||
)
|
||||
realised, _result, _splits = train_run(cfg)
|
||||
assert realised.run_id.endswith("-shuf")
|
||||
payload = evaluate_test(realised.out_dir / realised.run_id, n_boot=0)
|
||||
chance = 1.0 / payload["n_classes"]
|
||||
assert 0.0 <= payload["closed"]["recall_macro"] <= 1.5 * chance
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- CLI
|
||||
|
||||
|
||||
def _cli_store_cmd(store_dir: Path) -> int:
|
||||
return cli_main(
|
||||
[
|
||||
"store",
|
||||
"--synthetic",
|
||||
"--out",
|
||||
str(store_dir),
|
||||
"--n-genera",
|
||||
"4",
|
||||
"--reads-per-genus",
|
||||
"8",
|
||||
"--window-samples",
|
||||
"1200",
|
||||
"--shard-windows",
|
||||
"64",
|
||||
"--seed",
|
||||
"0",
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
def test_cli_end_to_end(tmp_path: Path) -> None:
|
||||
store_dir = tmp_path / "cli_store"
|
||||
assert _cli_store_cmd(store_dir) == 0
|
||||
out_dir = tmp_path / "runs"
|
||||
assert (
|
||||
cli_main(
|
||||
[
|
||||
"train",
|
||||
"--store",
|
||||
str(store_dir),
|
||||
"--out-dir",
|
||||
str(out_dir),
|
||||
"--arch",
|
||||
"cnn",
|
||||
"--budget",
|
||||
"100000",
|
||||
"--batch",
|
||||
"8",
|
||||
"--epochs",
|
||||
"3",
|
||||
"--patience",
|
||||
"3",
|
||||
"--workers",
|
||||
"0",
|
||||
"--min-test-windows",
|
||||
"1",
|
||||
"--amp",
|
||||
"off",
|
||||
"--device",
|
||||
"cpu",
|
||||
"--stage",
|
||||
"arch_ladder",
|
||||
]
|
||||
)
|
||||
== 0
|
||||
)
|
||||
run_dir = out_dir / "cnn_0.1M_s4_g4"
|
||||
assert (run_dir / "config.json").is_file()
|
||||
assert (run_dir / "ckpt.pt").is_file()
|
||||
assert cli_main(["validate", "--run-dir", str(run_dir), "--n-boot", "50"]) == 0
|
||||
assert (run_dir / "val_report.json").is_file()
|
||||
assert cli_main(["test", "--run-dir", str(run_dir), "--n-boot", "50"]) == 0
|
||||
assert (run_dir / "report.json").is_file()
|
||||
# a control run for the gate, then the verdict
|
||||
assert (
|
||||
cli_main(
|
||||
[
|
||||
"train",
|
||||
"--store",
|
||||
str(store_dir),
|
||||
"--out-dir",
|
||||
str(out_dir),
|
||||
"--arch",
|
||||
"cnn",
|
||||
"--budget",
|
||||
"100000",
|
||||
"--batch",
|
||||
"8",
|
||||
"--epochs",
|
||||
"3",
|
||||
"--patience",
|
||||
"3",
|
||||
"--workers",
|
||||
"0",
|
||||
"--min-test-windows",
|
||||
"1",
|
||||
"--amp",
|
||||
"off",
|
||||
"--device",
|
||||
"cpu",
|
||||
"--control",
|
||||
]
|
||||
)
|
||||
== 0
|
||||
)
|
||||
control_dir = out_dir / "cnn_0.1M_s4_g4-shuf"
|
||||
assert cli_main(["test", "--run-dir", str(control_dir), "--n-boot", "50"]) == 0
|
||||
assert cli_main(["gate", "--runs", str(run_dir), "--control", str(control_dir)]) == 0
|
||||
assert (out_dir / "gate_verdict.json").is_file()
|
||||
|
||||
|
||||
def test_cli_test_with_trap_store(tmp_path: Path, trap_store: Path) -> None:
|
||||
store_dir = tmp_path / "s"
|
||||
synthetic_store(
|
||||
store_dir, n_genera=4, reads_per_genus=10, window_samples=1200, shard_windows=64
|
||||
)
|
||||
out_dir = tmp_path / "runs"
|
||||
assert cli_main(["train", "--store", str(store_dir), "--out-dir", str(out_dir),
|
||||
"--arch", "cnn", "--budget", "100000", "--batch", "8", "--epochs", "2",
|
||||
"--min-test-windows", "1", "--workers", "0", "--amp", "off",
|
||||
"--device", "cpu"]) == 0
|
||||
run_dir = out_dir / "cnn_0.1M_s4_g4"
|
||||
assert (
|
||||
cli_main(
|
||||
["test", "--run-dir", str(run_dir), "--trap-store", str(trap_store), "--n-boot", "50"]
|
||||
)
|
||||
== 0
|
||||
)
|
||||
report = json.loads((run_dir / "report.json").read_text())
|
||||
assert report["trap"] is not None
|
||||
assert "trap_fpr_at_tau" in report["trap"]
|
||||
assert (run_dir / "trap_report.json").is_file()
|
||||
|
||||
|
||||
def test_build_subcommand_prints_realised(capsys: pytest.CaptureFixture) -> None:
|
||||
assert cli_main(["build", "--arch", "cnn", "--budget", "1000000", "--n-classes", "11"]) == 0
|
||||
payload = json.loads(capsys.readouterr().out)
|
||||
assert payload["params_realised"] <= int(1.1 * 1_000_000)
|
||||
assert payload["params_realised"] == payload["count_params"]
|
||||
assert payload["n_classes"] == 11
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- device/AMP
|
||||
|
||||
|
||||
def test_select_device_and_amp_policies() -> None:
|
||||
assert select_device("cpu").type == "cpu"
|
||||
assert select_device("").type == "cpu"
|
||||
with pytest.raises(ValueError):
|
||||
select_device("cuda")
|
||||
with pytest.raises(ValueError):
|
||||
select_device("not-a-device")
|
||||
device = torch.device("cpu")
|
||||
from custom_models.train import resolve_amp
|
||||
|
||||
assert resolve_amp("off", device) == "fp32"
|
||||
assert resolve_amp("bf16", device) == "fp32"
|
||||
assert resolve_amp("fp32", device) == "fp32"
|
||||
assert resolve_amp("fp16", device) == "fp32"
|
||||
Reference in New Issue
Block a user