# 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 ... # 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 ... ``` --- ## 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 ... ```