- 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
12 KiB
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):
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:
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:
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
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
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
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
# 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:
conda run -n e1-torch-cpu python -m custom_models ...