Initial release: E1 genus probe package moved out of the E01 workflow

- 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
This commit is contained in:
Tom Kasper
2026-09-22 13:22:50 +01:00
commit 88e98effad
18 changed files with 5988 additions and 0 deletions
+261
View File
@@ -0,0 +1,261 @@
# 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 ...
```