Made more compatible with experimental conditions
This commit is contained in:
+49
-16
@@ -24,7 +24,7 @@ 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.store import SignalStore, filter_excluded, synthetic_store
|
||||
from custom_models.train import select_device
|
||||
|
||||
# ---------------------------------------------------------------- store
|
||||
@@ -147,10 +147,43 @@ def test_permute_read_labels_preserves_counts(fast_store: Path) -> None:
|
||||
# ---------------------------------------------------------------- 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="arch_ladder",
|
||||
stage="ladder_point",
|
||||
store_dir=Path("/data/store"),
|
||||
data=DataConfig(genera=("A", "B"), val_fraction=0.2, seed=7),
|
||||
model=ModelConfig(
|
||||
@@ -172,7 +205,7 @@ def test_config_roundtrip(tmp_path: Path) -> None:
|
||||
|
||||
|
||||
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("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"
|
||||
|
||||
@@ -214,7 +247,7 @@ def test_trainability_duty_cycle(tmp_path: Path) -> None:
|
||||
)
|
||||
cfg = RunConfig(
|
||||
run_id="",
|
||||
stage="arch_ladder",
|
||||
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
|
||||
@@ -228,10 +261,10 @@ def test_trainability_duty_cycle(tmp_path: Path) -> None:
|
||||
device="cpu",
|
||||
batch_size=8,
|
||||
),
|
||||
out_dir=tmp_path / "runs",
|
||||
out_dir=tmp_path / "run",
|
||||
)
|
||||
realised, _result, _splits = train_run(cfg)
|
||||
run_dir = realised.out_dir / realised.run_id
|
||||
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
|
||||
@@ -256,7 +289,7 @@ def test_control_run_lands_at_chance(tmp_path: Path) -> None:
|
||||
)
|
||||
realised, _result, _splits = train_run(cfg)
|
||||
assert realised.run_id.endswith("-shuf")
|
||||
payload = evaluate_test(realised.out_dir / realised.run_id, n_boot=0)
|
||||
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
|
||||
|
||||
@@ -288,7 +321,7 @@ def _cli_store_cmd(store_dir: Path) -> int:
|
||||
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"
|
||||
run_dir = tmp_path / "run"
|
||||
assert (
|
||||
cli_main(
|
||||
[
|
||||
@@ -296,7 +329,7 @@ def test_cli_end_to_end(tmp_path: Path) -> None:
|
||||
"--store",
|
||||
str(store_dir),
|
||||
"--out-dir",
|
||||
str(out_dir),
|
||||
str(run_dir),
|
||||
"--arch",
|
||||
"cnn",
|
||||
"--budget",
|
||||
@@ -316,12 +349,12 @@ def test_cli_end_to_end(tmp_path: Path) -> None:
|
||||
"--device",
|
||||
"cpu",
|
||||
"--stage",
|
||||
"arch_ladder",
|
||||
"ladder_point",
|
||||
]
|
||||
)
|
||||
== 0
|
||||
)
|
||||
run_dir = out_dir / "cnn_0.1M_s4_g4"
|
||||
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
|
||||
@@ -336,7 +369,7 @@ def test_cli_end_to_end(tmp_path: Path) -> None:
|
||||
"--store",
|
||||
str(store_dir),
|
||||
"--out-dir",
|
||||
str(out_dir),
|
||||
str(tmp_path / "control"),
|
||||
"--arch",
|
||||
"cnn",
|
||||
"--budget",
|
||||
@@ -360,10 +393,10 @@ def test_cli_end_to_end(tmp_path: Path) -> None:
|
||||
)
|
||||
== 0
|
||||
)
|
||||
control_dir = out_dir / "cnn_0.1M_s4_g4-shuf"
|
||||
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 (out_dir / "gate_verdict.json").is_file()
|
||||
assert (tmp_path / "gate_verdict.json").is_file()
|
||||
|
||||
|
||||
def test_cli_test_with_trap_store(tmp_path: Path, trap_store: Path) -> None:
|
||||
@@ -371,12 +404,12 @@ def test_cli_test_with_trap_store(tmp_path: Path, trap_store: Path) -> None:
|
||||
synthetic_store(
|
||||
store_dir, n_genera=4, reads_per_genus=10, window_samples=1200, shard_windows=64
|
||||
)
|
||||
out_dir = tmp_path / "runs"
|
||||
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 / "cnn_0.1M_s4_g4"
|
||||
run_dir = out_dir
|
||||
assert (
|
||||
cli_main(
|
||||
["test", "--run-dir", str(run_dir), "--trap-store", str(trap_store), "--n-boot", "50"]
|
||||
|
||||
Reference in New Issue
Block a user