From 3e8ccdbdad8eb6edcf8a73ed74dc9396c351d1bf Mon Sep 17 00:00:00 2001 From: Tom Kasper Date: Sat, 3 Oct 2026 15:17:03 +0100 Subject: [PATCH] Made more compatible with experimental conditions --- README.md | 38 +++++++++-------- src/custom_models/cli.py | 21 +++++++-- src/custom_models/config.py | 11 ++++- src/custom_models/metrics.py | 12 ++++-- src/custom_models/runner.py | 6 ++- src/custom_models/store.py | 83 ++++++++++++++++++++++++++++++++++-- tests/test_pipeline.py | 65 +++++++++++++++++++++------- 7 files changed, 188 insertions(+), 48 deletions(-) diff --git a/README.md b/README.md index 673afbe..6345828 100644 --- a/README.md +++ b/README.md @@ -157,25 +157,27 @@ rehydrated, never guessed, when a saved checkpoint is loaded. | Flag | Default | Role | |------|---------|------| -| `--store` / `--out-dir` | required (or via base config) | Where the signal store lives and where the run directory is created. | +| `--store` / `--out-dir` | required (or via base config) | Where the signal store lives and the run directory itself (created if missing; artifacts written directly into it). | | `--config` | — | Base `RunConfig` JSON (a previous run's `config.json`). | | `--run-id` | derived | If empty: `{arch}_{budget_M}M_s{stride}_g{n_genera}` (control adds `-shuf`). | -| `--stage` | `arch_ladder` | Stage tag: `arch_ladder`, `windows_ablation`, `robustness`, `control`, `extra`. | +| `--stage` | `ladder_point` | Stage tag: `ladder_point`, `windows_ablation`, `robustness`, `control`, `extra`. | | `--control` | off | Marks the run as a shuffled-label control (overrides `--stage`). | | `--notes` | — | Free-text note stored verbatim in `config.json`. | ### Run directory layout +`--out-dir` is the run directory; give one per run (e.g. one per +snakemake rule output). + ```text OUT_DIR/ -└── RUN_ID/ - ├── config.json # realised RunConfig (round-trips exactly) - ├── ckpt.pt # best-epoch weights (CPU tensors) - ├── history.parquet # per-epoch: epoch, lr, train/val loss, val acc, recall, seconds - ├── val_report.json # written by `validate` - ├── report.json # written by `test` - ├── trap_report.json # written by `test` (when --trap-store is given) - └── gate_verdict.json # written by `gate` (common parent of the runs) +├── config.json # realised RunConfig (round-trips exactly) +├── ckpt.pt # best-epoch weights (CPU tensors) +├── history.parquet # per-epoch: epoch, lr, train/val loss, val acc, recall, seconds +├── val_report.json # written by `validate` +├── report.json # written by `test` +├── trap_report.json # written by `test` (when --trap-store is given) +└── gate_verdict.json # written by `gate` (common parent of the run dirs) ``` ### Control runs @@ -239,19 +241,19 @@ python -m custom_models build --arch cnn --budget 1000000 --n-classes 6 # 3. tune the class ladder (train each stage on the same store) for g in 3 4 5 6; do - python -m custom_models train --store data/store_v1 --out-dir runs \ - --arch cnn --budget 1000000 --genera $(head -n $g genera.txt | paste -sd,) \ - --stage arch_ladder --notes "information ceiling" + python -m custom_models train --store data/store_v1 \ + --out-dir runs/g$g --arch cnn --budget 1000000 \ + --genera $(head -n $g genera.txt | paste -sd,) \ + --stage ladder_point --notes "information ceiling" done -# 4. shuffled-label control -python -m custom_models train --store data/store_v1 --out-dir runs \ +# 4. shuffled-label control (its own run directory) +python -m custom_models train --store data/store_v1 --out-dir runs/control \ --arch cnn --budget 1000000 --control # 5. evaluate + gate (candidates need a `report.json` from `test`) -python -m custom_models test --run-dir runs/cnn_1.0M_s4_g6 --trap-store data/trap -python -m custom_models gate --runs runs/cnn_1.0M_s4_g6 \ - --control runs/cnn_1.0M_s4_g6-shuf +python -m custom_models test --run-dir runs/g6 --trap-store data/trap +python -m custom_models gate --runs runs/g6 --control runs/control ``` Run a subset via: diff --git a/src/custom_models/cli.py b/src/custom_models/cli.py index 8d846df..1314b49 100644 --- a/src/custom_models/cli.py +++ b/src/custom_models/cli.py @@ -196,7 +196,7 @@ def _run_config_from_args(args: argparse.Namespace) -> RunConfig: ), seed=_pick(args.data_seed, base_data.seed if base_data else None, 0), ) - stage = base.stage if base else "arch_ladder" + stage = base.stage if base else "ladder_point" if args.control: stage = "control" elif args.stage: @@ -343,6 +343,8 @@ def _cmd_store(args: argparse.Namespace) -> int: norm=args.norm, dtype=args.dtype, shard_windows=args.shard_windows, + exclude_refs=tuple(args.exclude_ref or ()), + exclude_genera=tuple(args.exclude_genus or ()), seed=args.seed, ) labels = load_labels(Path(args.labels)) @@ -373,7 +375,7 @@ def _cmd_train(args: argparse.Namespace) -> int: if realised.train.early_stop_metric == "val_recall_macro" else tail["val_loss"] ) - run_dir = realised.out_dir / realised.run_id + run_dir = realised.out_dir print( json.dumps( { @@ -507,6 +509,19 @@ def build_parser() -> argparse.ArgumentParser: p_store.add_argument("--norm", choices=["median_iqr", "median_mad"], default="median_iqr") p_store.add_argument("--dtype", choices=["float16", "float32"], default="float16") p_store.add_argument("--shard-windows", type=int, default=8_192) + p_store.add_argument( + "--exclude-ref", + nargs="+", + default=None, + help="Drop labelled reads whose ref_name contains any of these " + "substrings (e.g. DCS/lambda control contigs)", + ) + p_store.add_argument( + "--exclude-genus", + nargs="+", + default=None, + help="Drop labelled reads whose genus is any of these (e.g. trap captures)", + ) p_store.add_argument("--seed", type=int, default=0) p_store.set_defaults(func=_cmd_store) @@ -517,7 +532,7 @@ def build_parser() -> argparse.ArgumentParser: p_train.add_argument("--run-id", default=None) p_train.add_argument( "--stage", - choices=["arch_ladder", "windows_ablation", "robustness", "control", "extra"], + choices=["ladder_point", "windows_ablation", "robustness", "control", "extra"], default=None, ) p_train.add_argument("--control", action="store_true", default=None) diff --git a/src/custom_models/config.py b/src/custom_models/config.py index a0bb118..ec13533 100644 --- a/src/custom_models/config.py +++ b/src/custom_models/config.py @@ -16,7 +16,7 @@ from pathlib import Path from typing import Any, Literal Arch = Literal["cnn", "linatt"] -Stage = Literal["arch_ladder", "windows_ablation", "robustness", "control", "extra"] +Stage = Literal["ladder_point", "windows_ablation", "robustness", "control", "extra"] AmpPolicy = Literal["off", "bf16", "fp16"] EarlyStopMetric = Literal["val_recall_macro", "val_loss"] NormMethod = Literal["median_iqr", "median_mad"] @@ -105,6 +105,13 @@ class ExtractConfig: norm: Per-window normaliser applied before storing. dtype: On-disk shard dtype; fp16 halves the store at negligible cost. shard_windows: Windows per ``shard_%05d.npy`` file. + exclude_refs: Reference contig substrings whose labelled reads are + dropped (e.g. ``("NC_001416", "dcs")`` for ONT DNA Control + Sequence spike-ins that would otherwise be mislabelled into + Enterobacteriaceae genera). + exclude_genera: Genus labels dropped before extraction (e.g. trap + captures — reads of trap species that the labeller pulled into + the panel but that belong to no class of the probe). seed: Unused by extraction itself (window choice is deterministic); kept so every config in the pipeline is seed-carrying. @@ -118,6 +125,8 @@ class ExtractConfig: norm: NormMethod = "median_iqr" dtype: Literal["float16", "float32"] = "float16" shard_windows: int = 8_192 + exclude_refs: tuple[str, ...] = () + exclude_genera: tuple[str, ...] = () seed: int = 0 diff --git a/src/custom_models/metrics.py b/src/custom_models/metrics.py index c64703f..eb22be2 100644 --- a/src/custom_models/metrics.py +++ b/src/custom_models/metrics.py @@ -426,6 +426,9 @@ def trap_report( """ adjacency = adjacency if adjacency is not None else DEFAULT_TRAP_ADJACENCY + adjacency = { + k.lower(): {g.lower() for g in v} for k, v in adjacency.items() + } trap_logits = np.asarray(trap_logits, dtype=np.float64) n = trap_logits.shape[0] if len(trap_manifest) != n: @@ -442,11 +445,12 @@ def trap_report( read_arr = trap_manifest["read_id"].to_numpy() accepted = trap_msp >= tau family_ids_by_genus: dict[str, list[int]] = {} + genus_index = {g.lower(): i for i, g in enumerate(genera)} for trap_genus in set(genus_arr.tolist()): - family = adjacency.get(trap_genus, set()) - family_ids_by_genus[trap_genus] = [ - genera.index(g) for g in family if g in genera - ] + family = adjacency.get(trap_genus.lower(), set()) + family_ids_by_genus[trap_genus] = sorted( + genus_index[g] for g in family if g in genus_index + ) in_family = np.array( [ trap_pred[i] in family_ids_by_genus.get(genus_arr[i], []) diff --git a/src/custom_models/runner.py b/src/custom_models/runner.py index 596fff4..c0089e2 100644 --- a/src/custom_models/runner.py +++ b/src/custom_models/runner.py @@ -86,7 +86,9 @@ def train_run(cfg: RunConfig) -> tuple[RunConfig, TrainResult, Splits]: Args: cfg: The run configuration. ``cfg.run_id`` may be empty (then - derived from arch/budget/stride/class count/stage). + derived from arch/budget/stride/class count/stage); + ``cfg.out_dir`` *is* the run directory (created as needed, + artifacts written directly into it). Returns: ``(realised_cfg, result, splits)`` — the realised config that was @@ -119,7 +121,7 @@ def train_run(cfg: RunConfig) -> tuple[RunConfig, TrainResult, Splits]: model=realised_model, genera=tuple(splits.genera), ) - run_dir = realised.out_dir / run_id + run_dir = realised.out_dir run_dir.mkdir(parents=True, exist_ok=True) save_config(realised, run_dir / CONFIG_NAME) result = train_probe(probe, store, splits, realised.train, run_dir) diff --git a/src/custom_models/store.py b/src/custom_models/store.py index 3e874b2..770b024 100644 --- a/src/custom_models/store.py +++ b/src/custom_models/store.py @@ -4,7 +4,9 @@ The store contract is the GATE ZERO schema (§4, §5.3 with the §12.1 binding deviations): ``manifest.parquet`` (one row per window, ``window_id`` equal to row order), ``shard_%05d.npy`` files holding stacked fixed-length median-IQR-normalised windows, and ``store_meta.json`` describing provenance -and drop counts. A synthetic duty-cycle corpus builder provides the +and drop counts. When exclusions are requested (control-spike refs, capture +genera), ``excluded_reads.parquet`` lists every dropped labelled read with +the reason that triggered it. A synthetic duty-cycle corpus builder provides the "labels survive normalisation" fixture (GATE ZERO §12.9) and smoke-test data. """ @@ -43,6 +45,7 @@ MANIFEST_COLUMNS: tuple[str, ...] = ( STORE_META = "store_meta.json" MANIFEST_NAME = "manifest.parquet" +EXCLUDED_NAME = "excluded_reads.parquet" def median_iqr_normalise(x: np.ndarray) -> np.ndarray: @@ -225,6 +228,64 @@ def validate_labels(df: pl.DataFrame) -> None: raise ValueError("labels read_id column contains duplicates") +def filter_excluded( + labels: pl.DataFrame, exclude_refs: tuple[str, ...], exclude_genera: tuple[str, ...] +) -> tuple[pl.DataFrame, pl.DataFrame, dict[str, int]]: + """Drop capture/control reads from a labels table before store building. + + Two independent, ordered filters: (1) reads whose ``ref_name`` contains + any ``exclude_refs`` substring (case-insensitive) — the DCS-style + control-spike contigs that would otherwise soak up or shroud real + alignments; (2) reads whose ``genus`` is one of ``exclude_genera`` — + e.g. trap-genus captures that belong to no class of the probe. A read + matching both gets the ref reason (ref filter runs first). + + Args: + labels: Labels table with ``ref_name`` and ``genus`` columns. + exclude_refs: Substrings matching excluded reference contigs. + exclude_genera: Exact genus names excluded from the store. + + Returns: + ``(filtered, dropped, counts)``: the surviving table, the dropped + rows as-is plus a ``reason`` column (``excluded_ref:`` / + ``excluded_genus:``), and drop-count totals per filter. + + """ + counts: dict[str, int] = {} + dropped_frames: list[pl.DataFrame] = [] + working = labels + for ref in exclude_refs: + mask = ( + pl.col("ref_name") + .str.to_lowercase() + .str.contains(ref.lower(), literal=True) + .fill_null(False) + ) + hits = working.filter(mask) + counts[f"n_dropped_ref_{ref.lower()}"] = hits.height + if hits.height: + dropped_frames.append( + hits.with_columns(pl.lit(f"excluded_ref:{ref}").alias("reason")) + ) + working = working.filter(~mask) + for genus in exclude_genera: + mask = (pl.col("genus") == genus).fill_null(False) + hits = working.filter(mask) + counts[f"n_dropped_genus_{genus.lower()}"] = hits.height + if hits.height: + dropped_frames.append( + hits.with_columns(pl.lit(f"excluded_genus:{genus}").alias("reason")) + ) + working = working.filter(~mask) + if dropped_frames: + dropped = pl.concat(dropped_frames).sort("read_id") + else: + dropped = labels.head(0).with_columns(pl.lit("", dtype=pl.String).alias("reason")) + counts["n_rows_in"] = labels.height + counts["n_rows_out"] = working.height + return working, dropped, counts + + class _StoreWriter: """Streaming builder for a signal store directory. @@ -350,7 +411,7 @@ def extract_store(cfg: ExtractConfig, labels: pl.DataFrame, out_dir: Path) -> Pa their first ``max_windows_per_read`` contiguous windows (after skipping ``skip_head_samples``) are normalised and written as shards. Fully deterministic: no RNG anywhere in the window choice. Drop counters - (short/unlabelled) are reported in ``store_meta.json``. + (short/unlabelled/excluded) are reported in ``store_meta.json``. Args: cfg: Extraction parameters (paths, windowing, normaliser, dtype, @@ -373,7 +434,13 @@ def extract_store(cfg: ExtractConfig, labels: pl.DataFrame, out_dir: Path) -> Pa validate_labels(labels) if not cfg.pod5_paths: raise ValueError("ExtractConfig.pod5_paths is empty") - labelled = labels.filter(pl.col("genus").is_not_null() & (pl.col("genus") != "")) + labelled, dropped, exclude_counts = filter_excluded( + labels, cfg.exclude_refs, cfg.exclude_genera + ) + out_dir = Path(out_dir) + dropped.write_parquet(out_dir / EXCLUDED_NAME) + excluded_ids = set(dropped["read_id"].to_list()) + labelled = labelled.filter(pl.col("genus").is_not_null() & (pl.col("genus") != "")) if labelled.height == 0: raise ValueError("no reads with a non-null genus in the labels table") info = { @@ -390,11 +457,15 @@ def extract_store(cfg: ExtractConfig, labels: pl.DataFrame, out_dir: Path) -> Pa writer = _StoreWriter(out_dir, cfg.window_samples, dtype, cfg.shard_windows) n_dropped_short = 0 n_dropped_unlabelled = 0 + n_dropped_excluded = 0 for pod5_path in sorted(Path(p) for p in cfg.pod5_paths): kept: list[RawSignalRecord] = [] for record in read_pod5_records(pod5_path): if record.read_id not in info: - n_dropped_unlabelled += 1 + if record.read_id in excluded_ids: + n_dropped_excluded += 1 + else: + n_dropped_unlabelled += 1 continue if record.num_samples < cfg.min_read_samples: n_dropped_short += 1 @@ -420,9 +491,13 @@ def extract_store(cfg: ExtractConfig, labels: pl.DataFrame, out_dir: Path) -> Pa "dtype": str(cfg.dtype), "norm": cfg.norm, "pod5_paths": [str(p) for p in sorted(Path(p) for p in cfg.pod5_paths)], + "exclude_refs": list(cfg.exclude_refs), + "exclude_genera": list(cfg.exclude_genera), + **exclude_counts, "n_labelled_reads": labelled.height, "n_dropped_short": n_dropped_short, "n_dropped_unlabelled": n_dropped_unlabelled, + "n_dropped_excluded": n_dropped_excluded, } ) return Path(out_dir) diff --git a/tests/test_pipeline.py b/tests/test_pipeline.py index aab96a8..9a41490 100644 --- a/tests/test_pipeline.py +++ b/tests/test_pipeline.py @@ -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"]