"""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"