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 |
|
||||
|------|---------|------|
|
||||
| `--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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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], [])
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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