Made more compatible with experimental conditions

This commit is contained in:
Tom Kasper
2026-10-03 15:17:03 +01:00
parent 96ae5d78ba
commit 3e8ccdbdad
7 changed files with 188 additions and 48 deletions
+20 -18
View File
@@ -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:
+18 -3
View File
@@ -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)
+10 -1
View File
@@ -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
+8 -4
View File
@@ -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], [])
+4 -2
View File
@@ -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)
+78 -3
View File
@@ -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:<match>`` /
``excluded_genus:<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,10 +457,14 @@ 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:
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:
@@ -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)
+49 -16
View File
@@ -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"]