Compare commits
1
Commits
96ae5d78ba
..
main
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3e8ccdbdad |
@@ -157,25 +157,27 @@ rehydrated, never guessed, when a saved checkpoint is loaded.
|
|||||||
|
|
||||||
| Flag | Default | Role |
|
| 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`). |
|
| `--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`). |
|
| `--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`). |
|
| `--control` | off | Marks the run as a shuffled-label control (overrides `--stage`). |
|
||||||
| `--notes` | — | Free-text note stored verbatim in `config.json`. |
|
| `--notes` | — | Free-text note stored verbatim in `config.json`. |
|
||||||
|
|
||||||
### Run directory layout
|
### Run directory layout
|
||||||
|
|
||||||
|
`--out-dir` is the run directory; give one per run (e.g. one per
|
||||||
|
snakemake rule output).
|
||||||
|
|
||||||
```text
|
```text
|
||||||
OUT_DIR/
|
OUT_DIR/
|
||||||
└── RUN_ID/
|
├── config.json # realised RunConfig (round-trips exactly)
|
||||||
├── config.json # realised RunConfig (round-trips exactly)
|
├── ckpt.pt # best-epoch weights (CPU tensors)
|
||||||
├── ckpt.pt # best-epoch weights (CPU tensors)
|
├── history.parquet # per-epoch: epoch, lr, train/val loss, val acc, recall, seconds
|
||||||
├── history.parquet # per-epoch: epoch, lr, train/val loss, val acc, recall, seconds
|
├── val_report.json # written by `validate`
|
||||||
├── val_report.json # written by `validate`
|
├── report.json # written by `test`
|
||||||
├── report.json # written by `test`
|
├── trap_report.json # written by `test` (when --trap-store is given)
|
||||||
├── trap_report.json # written by `test` (when --trap-store is given)
|
└── gate_verdict.json # written by `gate` (common parent of the run dirs)
|
||||||
└── gate_verdict.json # written by `gate` (common parent of the runs)
|
|
||||||
```
|
```
|
||||||
|
|
||||||
### Control runs
|
### 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)
|
# 3. tune the class ladder (train each stage on the same store)
|
||||||
for g in 3 4 5 6; do
|
for g in 3 4 5 6; do
|
||||||
python -m custom_models train --store data/store_v1 --out-dir runs \
|
python -m custom_models train --store data/store_v1 \
|
||||||
--arch cnn --budget 1000000 --genera $(head -n $g genera.txt | paste -sd,) \
|
--out-dir runs/g$g --arch cnn --budget 1000000 \
|
||||||
--stage arch_ladder --notes "information ceiling"
|
--genera $(head -n $g genera.txt | paste -sd,) \
|
||||||
|
--stage ladder_point --notes "information ceiling"
|
||||||
done
|
done
|
||||||
|
|
||||||
# 4. shuffled-label control
|
# 4. shuffled-label control (its own run directory)
|
||||||
python -m custom_models train --store data/store_v1 --out-dir runs \
|
python -m custom_models train --store data/store_v1 --out-dir runs/control \
|
||||||
--arch cnn --budget 1000000 --control
|
--arch cnn --budget 1000000 --control
|
||||||
|
|
||||||
# 5. evaluate + gate (candidates need a `report.json` from `test`)
|
# 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 test --run-dir runs/g6 --trap-store data/trap
|
||||||
python -m custom_models gate --runs runs/cnn_1.0M_s4_g6 \
|
python -m custom_models gate --runs runs/g6 --control runs/control
|
||||||
--control runs/cnn_1.0M_s4_g6-shuf
|
|
||||||
```
|
```
|
||||||
|
|
||||||
Run a subset via:
|
Run a subset via:
|
||||||
|
|||||||
@@ -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),
|
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:
|
if args.control:
|
||||||
stage = "control"
|
stage = "control"
|
||||||
elif args.stage:
|
elif args.stage:
|
||||||
@@ -343,6 +343,8 @@ def _cmd_store(args: argparse.Namespace) -> int:
|
|||||||
norm=args.norm,
|
norm=args.norm,
|
||||||
dtype=args.dtype,
|
dtype=args.dtype,
|
||||||
shard_windows=args.shard_windows,
|
shard_windows=args.shard_windows,
|
||||||
|
exclude_refs=tuple(args.exclude_ref or ()),
|
||||||
|
exclude_genera=tuple(args.exclude_genus or ()),
|
||||||
seed=args.seed,
|
seed=args.seed,
|
||||||
)
|
)
|
||||||
labels = load_labels(Path(args.labels))
|
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"
|
if realised.train.early_stop_metric == "val_recall_macro"
|
||||||
else tail["val_loss"]
|
else tail["val_loss"]
|
||||||
)
|
)
|
||||||
run_dir = realised.out_dir / realised.run_id
|
run_dir = realised.out_dir
|
||||||
print(
|
print(
|
||||||
json.dumps(
|
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("--norm", choices=["median_iqr", "median_mad"], default="median_iqr")
|
||||||
p_store.add_argument("--dtype", choices=["float16", "float32"], default="float16")
|
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("--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.add_argument("--seed", type=int, default=0)
|
||||||
p_store.set_defaults(func=_cmd_store)
|
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("--run-id", default=None)
|
||||||
p_train.add_argument(
|
p_train.add_argument(
|
||||||
"--stage",
|
"--stage",
|
||||||
choices=["arch_ladder", "windows_ablation", "robustness", "control", "extra"],
|
choices=["ladder_point", "windows_ablation", "robustness", "control", "extra"],
|
||||||
default=None,
|
default=None,
|
||||||
)
|
)
|
||||||
p_train.add_argument("--control", action="store_true", default=None)
|
p_train.add_argument("--control", action="store_true", default=None)
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ from pathlib import Path
|
|||||||
from typing import Any, Literal
|
from typing import Any, Literal
|
||||||
|
|
||||||
Arch = Literal["cnn", "linatt"]
|
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"]
|
AmpPolicy = Literal["off", "bf16", "fp16"]
|
||||||
EarlyStopMetric = Literal["val_recall_macro", "val_loss"]
|
EarlyStopMetric = Literal["val_recall_macro", "val_loss"]
|
||||||
NormMethod = Literal["median_iqr", "median_mad"]
|
NormMethod = Literal["median_iqr", "median_mad"]
|
||||||
@@ -105,6 +105,13 @@ class ExtractConfig:
|
|||||||
norm: Per-window normaliser applied before storing.
|
norm: Per-window normaliser applied before storing.
|
||||||
dtype: On-disk shard dtype; fp16 halves the store at negligible cost.
|
dtype: On-disk shard dtype; fp16 halves the store at negligible cost.
|
||||||
shard_windows: Windows per ``shard_%05d.npy`` file.
|
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);
|
seed: Unused by extraction itself (window choice is deterministic);
|
||||||
kept so every config in the pipeline is seed-carrying.
|
kept so every config in the pipeline is seed-carrying.
|
||||||
|
|
||||||
@@ -118,6 +125,8 @@ class ExtractConfig:
|
|||||||
norm: NormMethod = "median_iqr"
|
norm: NormMethod = "median_iqr"
|
||||||
dtype: Literal["float16", "float32"] = "float16"
|
dtype: Literal["float16", "float32"] = "float16"
|
||||||
shard_windows: int = 8_192
|
shard_windows: int = 8_192
|
||||||
|
exclude_refs: tuple[str, ...] = ()
|
||||||
|
exclude_genera: tuple[str, ...] = ()
|
||||||
seed: int = 0
|
seed: int = 0
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -426,6 +426,9 @@ def trap_report(
|
|||||||
|
|
||||||
"""
|
"""
|
||||||
adjacency = adjacency if adjacency is not None else DEFAULT_TRAP_ADJACENCY
|
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)
|
trap_logits = np.asarray(trap_logits, dtype=np.float64)
|
||||||
n = trap_logits.shape[0]
|
n = trap_logits.shape[0]
|
||||||
if len(trap_manifest) != n:
|
if len(trap_manifest) != n:
|
||||||
@@ -442,11 +445,12 @@ def trap_report(
|
|||||||
read_arr = trap_manifest["read_id"].to_numpy()
|
read_arr = trap_manifest["read_id"].to_numpy()
|
||||||
accepted = trap_msp >= tau
|
accepted = trap_msp >= tau
|
||||||
family_ids_by_genus: dict[str, list[int]] = {}
|
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()):
|
for trap_genus in set(genus_arr.tolist()):
|
||||||
family = adjacency.get(trap_genus, set())
|
family = adjacency.get(trap_genus.lower(), set())
|
||||||
family_ids_by_genus[trap_genus] = [
|
family_ids_by_genus[trap_genus] = sorted(
|
||||||
genera.index(g) for g in family if g in genera
|
genus_index[g] for g in family if g in genus_index
|
||||||
]
|
)
|
||||||
in_family = np.array(
|
in_family = np.array(
|
||||||
[
|
[
|
||||||
trap_pred[i] in family_ids_by_genus.get(genus_arr[i], [])
|
trap_pred[i] in family_ids_by_genus.get(genus_arr[i], [])
|
||||||
|
|||||||
@@ -86,7 +86,9 @@ def train_run(cfg: RunConfig) -> tuple[RunConfig, TrainResult, Splits]:
|
|||||||
|
|
||||||
Args:
|
Args:
|
||||||
cfg: The run configuration. ``cfg.run_id`` may be empty (then
|
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:
|
Returns:
|
||||||
``(realised_cfg, result, splits)`` — the realised config that was
|
``(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,
|
model=realised_model,
|
||||||
genera=tuple(splits.genera),
|
genera=tuple(splits.genera),
|
||||||
)
|
)
|
||||||
run_dir = realised.out_dir / run_id
|
run_dir = realised.out_dir
|
||||||
run_dir.mkdir(parents=True, exist_ok=True)
|
run_dir.mkdir(parents=True, exist_ok=True)
|
||||||
save_config(realised, run_dir / CONFIG_NAME)
|
save_config(realised, run_dir / CONFIG_NAME)
|
||||||
result = train_probe(probe, store, splits, realised.train, run_dir)
|
result = train_probe(probe, store, splits, realised.train, run_dir)
|
||||||
|
|||||||
@@ -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
|
deviations): ``manifest.parquet`` (one row per window, ``window_id`` equal to
|
||||||
row order), ``shard_%05d.npy`` files holding stacked fixed-length
|
row order), ``shard_%05d.npy`` files holding stacked fixed-length
|
||||||
median-IQR-normalised windows, and ``store_meta.json`` describing provenance
|
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.
|
"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"
|
STORE_META = "store_meta.json"
|
||||||
MANIFEST_NAME = "manifest.parquet"
|
MANIFEST_NAME = "manifest.parquet"
|
||||||
|
EXCLUDED_NAME = "excluded_reads.parquet"
|
||||||
|
|
||||||
|
|
||||||
def median_iqr_normalise(x: np.ndarray) -> np.ndarray:
|
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")
|
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:
|
class _StoreWriter:
|
||||||
"""Streaming builder for a signal store directory.
|
"""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
|
their first ``max_windows_per_read`` contiguous windows (after skipping
|
||||||
``skip_head_samples``) are normalised and written as shards. Fully
|
``skip_head_samples``) are normalised and written as shards. Fully
|
||||||
deterministic: no RNG anywhere in the window choice. Drop counters
|
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:
|
Args:
|
||||||
cfg: Extraction parameters (paths, windowing, normaliser, dtype,
|
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)
|
validate_labels(labels)
|
||||||
if not cfg.pod5_paths:
|
if not cfg.pod5_paths:
|
||||||
raise ValueError("ExtractConfig.pod5_paths is empty")
|
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:
|
if labelled.height == 0:
|
||||||
raise ValueError("no reads with a non-null genus in the labels table")
|
raise ValueError("no reads with a non-null genus in the labels table")
|
||||||
info = {
|
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)
|
writer = _StoreWriter(out_dir, cfg.window_samples, dtype, cfg.shard_windows)
|
||||||
n_dropped_short = 0
|
n_dropped_short = 0
|
||||||
n_dropped_unlabelled = 0
|
n_dropped_unlabelled = 0
|
||||||
|
n_dropped_excluded = 0
|
||||||
for pod5_path in sorted(Path(p) for p in cfg.pod5_paths):
|
for pod5_path in sorted(Path(p) for p in cfg.pod5_paths):
|
||||||
kept: list[RawSignalRecord] = []
|
kept: list[RawSignalRecord] = []
|
||||||
for record in read_pod5_records(pod5_path):
|
for record in read_pod5_records(pod5_path):
|
||||||
if record.read_id not in info:
|
if record.read_id not in info:
|
||||||
|
if record.read_id in excluded_ids:
|
||||||
|
n_dropped_excluded += 1
|
||||||
|
else:
|
||||||
n_dropped_unlabelled += 1
|
n_dropped_unlabelled += 1
|
||||||
continue
|
continue
|
||||||
if record.num_samples < cfg.min_read_samples:
|
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),
|
"dtype": str(cfg.dtype),
|
||||||
"norm": cfg.norm,
|
"norm": cfg.norm,
|
||||||
"pod5_paths": [str(p) for p in sorted(Path(p) for p in cfg.pod5_paths)],
|
"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_labelled_reads": labelled.height,
|
||||||
"n_dropped_short": n_dropped_short,
|
"n_dropped_short": n_dropped_short,
|
||||||
"n_dropped_unlabelled": n_dropped_unlabelled,
|
"n_dropped_unlabelled": n_dropped_unlabelled,
|
||||||
|
"n_dropped_excluded": n_dropped_excluded,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
return Path(out_dir)
|
return Path(out_dir)
|
||||||
|
|||||||
+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.metrics import choose_tau_msp, closed_set_report, msp
|
||||||
from custom_models.runner import evaluate_test, evaluate_val, train_run
|
from custom_models.runner import evaluate_test, evaluate_val, train_run
|
||||||
from custom_models.seed import seed_everything
|
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
|
from custom_models.train import select_device
|
||||||
|
|
||||||
# ---------------------------------------------------------------- store
|
# ---------------------------------------------------------------- store
|
||||||
@@ -147,10 +147,43 @@ def test_permute_read_labels_preserves_counts(fast_store: Path) -> None:
|
|||||||
# ---------------------------------------------------------------- configs
|
# ---------------------------------------------------------------- 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:
|
def test_config_roundtrip(tmp_path: Path) -> None:
|
||||||
cfg = RunConfig(
|
cfg = RunConfig(
|
||||||
run_id="cnn_1.0M_s4_g11",
|
run_id="cnn_1.0M_s4_g11",
|
||||||
stage="arch_ladder",
|
stage="ladder_point",
|
||||||
store_dir=Path("/data/store"),
|
store_dir=Path("/data/store"),
|
||||||
data=DataConfig(genera=("A", "B"), val_fraction=0.2, seed=7),
|
data=DataConfig(genera=("A", "B"), val_fraction=0.2, seed=7),
|
||||||
model=ModelConfig(
|
model=ModelConfig(
|
||||||
@@ -172,7 +205,7 @@ def test_config_roundtrip(tmp_path: Path) -> None:
|
|||||||
|
|
||||||
|
|
||||||
def test_run_id_format() -> 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("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"
|
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(
|
cfg = RunConfig(
|
||||||
run_id="",
|
run_id="",
|
||||||
stage="arch_ladder",
|
stage="ladder_point",
|
||||||
store_dir=store_dir,
|
store_dir=store_dir,
|
||||||
data=DataConfig(
|
data=DataConfig(
|
||||||
genera=(), val_fraction=0.25, test_fraction=0.25, min_test_windows_per_genus=1, seed=0
|
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",
|
device="cpu",
|
||||||
batch_size=8,
|
batch_size=8,
|
||||||
),
|
),
|
||||||
out_dir=tmp_path / "runs",
|
out_dir=tmp_path / "run",
|
||||||
)
|
)
|
||||||
realised, _result, _splits = train_run(cfg)
|
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()
|
assert (run_dir / "history.parquet").is_file()
|
||||||
val_payload = evaluate_val(run_dir, n_boot=0)
|
val_payload = evaluate_val(run_dir, n_boot=0)
|
||||||
assert val_payload["closed"]["recall_macro"] >= 0.9
|
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)
|
realised, _result, _splits = train_run(cfg)
|
||||||
assert realised.run_id.endswith("-shuf")
|
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"]
|
chance = 1.0 / payload["n_classes"]
|
||||||
assert 0.0 <= payload["closed"]["recall_macro"] <= 1.5 * chance
|
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:
|
def test_cli_end_to_end(tmp_path: Path) -> None:
|
||||||
store_dir = tmp_path / "cli_store"
|
store_dir = tmp_path / "cli_store"
|
||||||
assert _cli_store_cmd(store_dir) == 0
|
assert _cli_store_cmd(store_dir) == 0
|
||||||
out_dir = tmp_path / "runs"
|
run_dir = tmp_path / "run"
|
||||||
assert (
|
assert (
|
||||||
cli_main(
|
cli_main(
|
||||||
[
|
[
|
||||||
@@ -296,7 +329,7 @@ def test_cli_end_to_end(tmp_path: Path) -> None:
|
|||||||
"--store",
|
"--store",
|
||||||
str(store_dir),
|
str(store_dir),
|
||||||
"--out-dir",
|
"--out-dir",
|
||||||
str(out_dir),
|
str(run_dir),
|
||||||
"--arch",
|
"--arch",
|
||||||
"cnn",
|
"cnn",
|
||||||
"--budget",
|
"--budget",
|
||||||
@@ -316,12 +349,12 @@ def test_cli_end_to_end(tmp_path: Path) -> None:
|
|||||||
"--device",
|
"--device",
|
||||||
"cpu",
|
"cpu",
|
||||||
"--stage",
|
"--stage",
|
||||||
"arch_ladder",
|
"ladder_point",
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
== 0
|
== 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 / "config.json").is_file()
|
||||||
assert (run_dir / "ckpt.pt").is_file()
|
assert (run_dir / "ckpt.pt").is_file()
|
||||||
assert cli_main(["validate", "--run-dir", str(run_dir), "--n-boot", "50"]) == 0
|
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",
|
"--store",
|
||||||
str(store_dir),
|
str(store_dir),
|
||||||
"--out-dir",
|
"--out-dir",
|
||||||
str(out_dir),
|
str(tmp_path / "control"),
|
||||||
"--arch",
|
"--arch",
|
||||||
"cnn",
|
"cnn",
|
||||||
"--budget",
|
"--budget",
|
||||||
@@ -360,10 +393,10 @@ def test_cli_end_to_end(tmp_path: Path) -> None:
|
|||||||
)
|
)
|
||||||
== 0
|
== 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(["test", "--run-dir", str(control_dir), "--n-boot", "50"]) == 0
|
||||||
assert cli_main(["gate", "--runs", str(run_dir), "--control", str(control_dir)]) == 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:
|
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(
|
synthetic_store(
|
||||||
store_dir, n_genera=4, reads_per_genus=10, window_samples=1200, shard_windows=64
|
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),
|
assert cli_main(["train", "--store", str(store_dir), "--out-dir", str(out_dir),
|
||||||
"--arch", "cnn", "--budget", "100000", "--batch", "8", "--epochs", "2",
|
"--arch", "cnn", "--budget", "100000", "--batch", "8", "--epochs", "2",
|
||||||
"--min-test-windows", "1", "--workers", "0", "--amp", "off",
|
"--min-test-windows", "1", "--workers", "0", "--amp", "off",
|
||||||
"--device", "cpu"]) == 0
|
"--device", "cpu"]) == 0
|
||||||
run_dir = out_dir / "cnn_0.1M_s4_g4"
|
run_dir = out_dir
|
||||||
assert (
|
assert (
|
||||||
cli_main(
|
cli_main(
|
||||||
["test", "--run-dir", str(run_dir), "--trap-store", str(trap_store), "--n-boot", "50"]
|
["test", "--run-dir", str(run_dir), "--trap-store", str(trap_store), "--n-boot", "50"]
|
||||||
|
|||||||
Reference in New Issue
Block a user