- src/custom_models torch modules, store, training, metrics, runner and CLI (package-relative imports for pip installability) - test suite (49 tests) + conftest with fast synthetic-store fixtures - uv-managed dev env (pyproject + CPU torch index), hatchling build - README: uv for the package, pip install into conda envs for the workflow
262 lines
12 KiB
Markdown
262 lines
12 KiB
Markdown
# 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 where the run directory is created. |
|
||
| `--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`. |
|
||
| `--control` | off | Marks the run as a shuffled-label control (overrides `--stage`). |
|
||
| `--notes` | — | Free-text note stored verbatim in `config.json`. |
|
||
|
||
### Run directory layout
|
||
|
||
```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)
|
||
```
|
||
|
||
### 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 \
|
||
--arch cnn --budget 1000000 --genera $(head -n $g genera.txt | paste -sd,) \
|
||
--stage arch_ladder --notes "information ceiling"
|
||
done
|
||
|
||
# 4. shuffled-label control
|
||
python -m custom_models train --store data/store_v1 --out-dir runs \
|
||
--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
|
||
```
|
||
|
||
Run a subset via:
|
||
|
||
```bash
|
||
conda run -n e1-torch-cpu python -m custom_models ...
|
||
```
|