Files

450 lines
16 KiB
Python

"""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, filter_excluded, 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_filter_excluded() -> None:
labels = pl.DataFrame(
{
"read_id": ["a", "b", "c", "d", "e"],
"genus": ["shigella", "escherichia", "shigella", "citrobacter", "salmonella"],
"species": ["s"] * 5,
"ref_name": ["NC_004851", "NC_001416", "NC_004851", "NC_002006", "NC_003197"],
"source_run": ["r"] * 5,
}
)
filtered, dropped, counts = filter_excluded(labels, ("NC_001416",), ("escherichia",))
assert filtered.height == 4
assert dropped.height == 1
assert dropped["read_id"][0] == "b"
assert dropped["reason"][0] == "excluded_ref:NC_001416"
assert "escherichia" not in filtered["genus"]
assert "NC_001416" not in filtered["ref_name"]
assert counts["n_dropped_ref_nc_001416"] == 1
assert counts["n_dropped_genus_escherichia"] == 0
assert counts["n_rows_in"] == 5 and counts["n_rows_out"] == 4
# ref match is case-insensitive; genus filter can also catch the capture
filtered2, dropped2, counts2 = filter_excluded(labels, ("nc_001416",), ())
assert filtered2.height == 4 and counts2["n_dropped_ref_nc_001416"] == 1
assert dropped2["reason"][0] == "excluded_ref:nc_001416"
filtered3, dropped3, _ = filter_excluded(labels, (), ("citrobacter",))
assert filtered3.height == 4
assert dropped3["read_id"][0] == "d" and dropped3["reason"][0] == "excluded_genus:citrobacter"
# empty filters: nothing dropped, empty schema-correct report
filtered4, dropped4, counts4 = filter_excluded(labels, (), ())
assert filtered4.height == 5 and dropped4.height == 0
assert counts4["n_rows_out"] == 5 and "reason" in dropped4.columns
def test_config_roundtrip(tmp_path: Path) -> None:
cfg = RunConfig(
run_id="cnn_1.0M_s4_g11",
stage="ladder_point",
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, "ladder_point") == "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="ladder_point",
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 / "run",
)
realised, _result, _splits = train_run(cfg)
run_dir = realised.out_dir
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, 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
run_dir = tmp_path / "run"
assert (
cli_main(
[
"train",
"--store",
str(store_dir),
"--out-dir",
str(run_dir),
"--arch",
"cnn",
"--budget",
"100000",
"--batch",
"8",
"--epochs",
"3",
"--patience",
"3",
"--workers",
"0",
"--min-test-windows",
"1",
"--amp",
"off",
"--device",
"cpu",
"--stage",
"ladder_point",
]
)
== 0
)
run_dir = tmp_path / "run"
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(tmp_path / "control"),
"--arch",
"cnn",
"--budget",
"100000",
"--batch",
"8",
"--epochs",
"3",
"--patience",
"3",
"--workers",
"0",
"--min-test-windows",
"1",
"--amp",
"off",
"--device",
"cpu",
"--control",
]
)
== 0
)
control_dir = tmp_path / "control"
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 (tmp_path / "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 / "run"
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
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"