Files

12 KiB
Raw Permalink Blame History

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 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).

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

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/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:

conda run -n e1-torch-cpu python -m custom_models ...