Files
eDNA_Stream_E01_custom_models/README.md
T

264 lines
12 KiB
Markdown
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# custom_models — E1 genus probe CLI
Torch models and a generalised train/validate/test wrapper for the E1
genus-information-ceiling experiment (GATE ZERO / torch model spec).
`custom-models` is its own pip-installable package (this repo, `src/`
layout, built with hatchling).
## Install
Dev/build environment is **uv** (much faster solver + lockfile):
```bash
uv sync --extra dev # venv at .venv with CPU torch from the uv index
uv run pytest # run the test suite
uv run custom-models build --arch cnn --budget 1000000 --n-classes 11
```
Import into the main workflow's **conda** envs — torch/numpy/polars come
from conda there and already satisfy the requirements, so pip adds only
this package:
```bash
conda run -n e1-torch-cpu pip install /path/to/custom_models
conda run -n e1-torch-cpu python -m custom_models <subcommand> ...
# or with pod5 support for the `store` subcommand:
conda run -n e1-torch-cpu pip install "/path/to/custom_models[pod5]"
```
Run as:
```bash
python -m custom_models <subcommand> ...
```
---
## Subcommands
| Subcommand | Purpose |
|-------------|---------|
| `build` | Resolve a nominal parameter budget into a realised probe geometry; prints JSON. |
| `store` | Build a signal store either from pod5 + labels parquet or a synthetic duty-cycle corpus. |
| `train` | One training run on a store; optionally a shuffled-label control. Writes a run directory. |
| `validate` | Evaluate a run's checkpoint on its validation split (bootstrap CIs). |
| `test` | Evaluate a run's checkpoint on its test split; optionally attach a trap (out-of-bank) store. |
| `gate` | E1-v2 gate verdict: compare candidate run dirs against a control run. |
| `configs` | Round-trip helper: load a config JSON and re-save it (`load_config(save_config(c)) == c`). |
Every subcommand seeds before doing anything; all exit 0 on success, 1 on
usage/runtime errors (message on stderr).
---
## Input: the labels parquet
`store` (pod5 mode) and `extract_store` require a labels table (parquet)
with these columns (order-free; `LABEL_COLUMNS` in `store.py`):
| Column | Type | Role |
|--------------|--------|------|
| `read_id` | str | ONT read id — the join key against the pod5 files. **Must be unique** (duplicates raise `ValueError`). |
| `genus` | str | Class label used by the probe. |
| `species` | str | Fine-grained secondary label (carried through to the manifest; unused by training). |
| `ref_name` | str | Reference sequence name the read was aligned to. |
| `identity` | float | Alignment identity (e.g. minimap2 `id%`), for provenance/filtering elsewhere. |
| `aligned_len`| int | Aligned reference length. |
| `query_len` | int | Read length (basis for the coverage-style ratios). |
| `source_run` | str | Sequencing run the read came from — enables the `holdout_runs` split mode. |
Extra columns are ignored. The genus set actually trained on is chosen at
run time with `train --genera A,B,C` (default: every genus in the store).
## Input: the pod5 files
`store --pod5 FILE [FILE ...] --labels LABELS.parquet --out STORE_DIR`.
Windows are cut per read: skip `--skip-head` samples (adapter/mux
artefacts), then take up to `--max-windows-per-read` contiguous fixed
windows of `--window-samples` samples (default 12,000 ≈ 2.4 s at 5 kHz).
Reads shorter than `--min-read-samples` are dropped. Each window is
normalised (`--norm median_iqr`, the amplitude-erasing default) and
written float16 (`--dtype`).
## Output: the signal store layout
```text
STORE_DIR/
├── store_meta.json # extraction stats, counts, dropped-read accounting
├── manifest.parquet # per-window table:
│ # window_id, read_id, genus, species, source_run,
│ # window_index, n_samples
└── shard_%05d.npy # float16 windows, one file per --shard-windows windows
```
Both training paths (real + synthetic) read through `SignalStore`, which
validates the manifest columns on open.
---
## `train` and the option model
`train` accepts a previous run's `config.json` as `--config` base;
explicit flags override it. Resolution order per field:
```
explicit flag > base config value > spec default
```
This makes variants (ablations, holdout robustness, retraining,
relocation of the run dir) reproducible by editing one or two flags.
### Model parameters (`models.py` role)
| Flag | Default | Role |
|------|---------|------|
| `--arch` | `cnn` | Encoder family: `cnn` (patch-embed + residual conv blocks + gated mean pool) or `linatt` (same conv stack tokenised, linear attention blocks, mean pool). |
| `--budget` | 1,000,000 | Nominal trainable-parameter budget. `build_probe` searches layer widths up to `budget * (1 + --budget-tol)` for ceilings well above the arch's table count. `build`, then the basis for the deterministic search over deeper/wider stacks. |
| `--stride` | 4 | Patch stride in signal samples (token pitch), 4 or 8; also fixes the patch length (4 samples / patch). Smaller stride → more tokens → finer temporal resolution, more compute. |
| `--d-embed` | 256 | Pooled embedding width the encoder terminates in. |
| `--budget-tol` | 0.10 | Relative tolerance; realised params must be `<= (1+tol) × budget`. |
| `--n-heads` | 4 | Attention head count for `linatt` (recorded for exact checkpoint rehydration). |
Realised geometry (`d_model`, `n_layers`, `params_realised`) is filled by
`build_probe` and written into the run's `config.json`; it is
rehydrated, never guessed, when a saved checkpoint is loaded.
### Data parameters
| Flag | Default | Role |
|------|---------|------|
| `--genera` | all in store | Comma-separated class whitelist (e.g. `Bacillus,Listeria`). Genera absent from the store drop out with a warning — the run id records the realised count `gN`. |
| `--val-frac` | 0.10 | Fraction of reads (per genus) used for validation (splits are read-level so windows of one read never straddle a split). |
| `--test-frac` | 0.10 | Fraction of reads for test. |
| `--holdout-runs` | off | Split by `source_run` instead of by read (robustness row: unseen instrument/run). |
| `--min-test-windows` | 300 | Per-genus floor for test windows; `make_splits` refuses to build a split that would starve a genus rather than silently evaluating on nothing. |
| `--data-seed` | 0 | Split permutation seed (independent of the training seed). |
### Training parameters
| Flag | Default | Role |
|------|---------|------|
| `--batch` | 256 | Optimiser batch size (read windows). Train loader uses `drop_last=True`. |
| `--lr` | 3e-4 | Peak AdamW learning rate; the warmup-cosine schedule holds returns to 0 by `--epochs`. |
| `--weight-decay` | 0.01 | Decoupled decay; not applied to norm weights or biases. |
| `--epochs` | 40 | Maximum epochs (early stopping may exit sooner). |
| `--warmup` | 0.05 | Fraction of total optimiser steps spent ramping LR 0→peak. |
| `--smoothing` | 0.0 | Cross-entropy label smoothing. |
| `--amp` | `bf16` | Autocast policy: `bf16` where supported, `fp16` (CUDA + GradScaler) or `off` (fp32; also forced on CPU). Falls back silently to fp32 when unsupported. |
| `--grad-clip` | 1.0 | L2 gradient-norm clip. |
| `--metric` | `val_recall_macro` | Early-stop selection metric (the headline gate metric; `val_loss` is the alternative). |
| `--patience` | 5 | Consecutive epochs without improvement before stopping. |
| `--workers` | 4 | DataLoader worker processes (workers are seeded; Linux pins the `fork` start method because 3.14 default `forkserver` is broken here). |
| `--device` | auto | `cuda` → `mps` → `cpu` resolution; or an explicit torch device string. |
| `--seed` | 0 | Training seed; every stochastic step is seeded from it. |
| `--deterministic` | on | Ensure kernels choose inverse-deterministic ops. |
### Run bookkeeping
| Flag | Default | Role |
|------|---------|------|
| `--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` | `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/
├── 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
`train --control` permutes the read-to-genus mapping across the selected
pool (seeded, preserving the genus multiset exactly); every window
inherits its read's new label. Applied *before* splitting, so the
control's splits and labels live consistently in the permuted world and
a trained control should land at chance accuracy (1/n_classes). Control
runs exist so the gate has a leakage-litmus baseline.
---
## `validate` / `test` / `gate`
```bash
python -m custom_models validate --run-dir OUT_DIR/RUN_ID [--store S] [--n-boot 10000]
python -m custom_models test --run-dir OUT_DIR/RUN_ID [--store S] [--trap-store T] [--n-boot]
python -m custom_models gate --runs RUN_A RUN_B ... --control OUT_DIR/CNN...shuf [--chance-tol 1.5]
```
`--store` points at a different signal store than the one recorded in
`config.json` (e.g. evaluating a synthetic-trained probe on real data —
cross-run transfer); by default the recorded store is re-used.
`--trap-store` runs the probe against an out-of-class "trap" set during
`test`: the closed-set open-set probe reports max-softmax-probability
(msp) calibration, threshold `tau`, and trap-vs-known msp separation
(FPR/FNR at the chosen `tau`).
`gate` evaluates the closed- and open-set reports of each candidate run
against the E1-v2 gate rules (chance = 1/n_classes from the control's
config):
- The **control** must itself be at chance (recall ≤ `--chance-tol` ×
chance); a control above that is INVALID (suspected label leakage) and
halts the interpretation.
- **PASS**: test recall-macro ≥ 0.90 with ≤ 5M realised params and trap
FPR at tau ≤ 0.05.
- **MARGINAL-TRAP**: recall ≥ 0.90 at ≤ 5M but failing trap FPR — adopt
per-class/margin thresholding.
- **MARGINAL**: recall in [0.70, 0.90) at ≤ 5M, or 0.90 only at ≤ 20M
with the trap criterion satisfied.
- **FAIL**: anything else. Missing trap evaluation degrades PASS to
MARGINAL (the open-set criterion cannot be confirmed).
The verdict (`gate_verdict.json`) is written to the runs' common parent
directory.
---
## Typical workflow
```bash
# 1. store
python -m custom_models store --pod5 run1.pod5 run2.pod5 \
--labels reads.parquet --out data/store_v1
# 2. build (check geometry before training)
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/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 (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/g6 --trap-store data/trap
python -m custom_models gate --runs runs/g6 --control runs/control
```
Run a subset via:
```bash
conda run -n e1-torch-cpu python -m custom_models ...
```