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
+571
View File
@@ -0,0 +1,571 @@
"""Generalised train/validate/test CLI for the E1 genus probe.
Subcommands: ``build`` (resolve a budget), ``store`` (build a signal
store from pod5+labels or a synthetic corpus), ``train`` (a full run,
optionally a shuffled-label control), ``validate`` / ``test`` (evaluate a
run directory's checkpoint, ``test`` optionally attaching a trap store),
``gate`` (verdict over run dirs) and ``configs`` (config round-trip).
Every subcommand seeds first (GATE ZERO §2.3); each returns exit code 0
on success and 1 on usage/runtime errors. ``train`` accepts a previous
run's ``config.json`` as a base with any flag overriding it, so variants
(windows ablation, holdout robustness, retraining) stay reproducible.
"""
from __future__ import annotations
import argparse
import json
import sys
from collections.abc import Sequence
from pathlib import Path
from typing import Any
from .config import (
DataConfig,
ExtractConfig,
ModelConfig,
RunConfig,
TrainConfig,
load_config,
save_config,
)
from .runner import evaluate_gate, evaluate_test, evaluate_val, train_run
def _add_model_args(parser: argparse.ArgumentParser, required: bool) -> None:
"""Register the model-request flags on a subparser.
Args:
parser: Target subparser.
required: Whether ``--arch``/``--budget`` must be given (the
``build`` command demands them; ``train`` falls back to
defaults/base-config values).
"""
parser.add_argument("--arch", choices=["cnn", "linatt"], required=required)
parser.add_argument("--budget", type=int, required=required)
parser.add_argument("--stride", type=int, default=None)
parser.add_argument("--d-embed", type=int, default=None)
parser.add_argument("--budget-tol", type=float, default=None)
parser.add_argument("--n-heads", type=int, default=None)
def _add_data_args(parser: argparse.ArgumentParser) -> None:
"""Register the split-configuration flags on a subparser.
All default to ``None`` so only explicitly passed flags override a
base config (tri-state booleans use ``default=None``).
Args:
parser: Target subparser.
"""
parser.add_argument("--genera", default=None)
parser.add_argument("--val-frac", type=float, default=None)
parser.add_argument("--test-frac", type=float, default=None)
parser.add_argument("--holdout-runs", action="store_true", default=None)
parser.add_argument("--min-test-windows", type=int, default=None)
parser.add_argument("--data-seed", type=int, default=None)
def _add_train_args(parser: argparse.ArgumentParser) -> None:
"""Register the training-semantics flags on a subparser.
Args:
parser: Target subparser.
"""
parser.add_argument("--batch", type=int, default=None)
parser.add_argument("--lr", type=float, default=None)
parser.add_argument("--weight-decay", type=float, default=None)
parser.add_argument("--epochs", type=int, default=None)
parser.add_argument("--warmup", type=float, default=None)
parser.add_argument("--smoothing", type=float, default=None)
parser.add_argument("--amp", choices=["off", "bf16", "fp16"], default=None)
parser.add_argument("--grad-clip", type=float, default=None)
parser.add_argument(
"--metric", choices=["val_recall_macro", "val_loss"], default=None
)
parser.add_argument("--patience", type=int, default=None)
parser.add_argument("--workers", type=int, default=None)
parser.add_argument("--device", default=None)
parser.add_argument("--seed", type=int, default=None)
parser.add_argument("--deterministic", action="store_true", default=None)
def _pick(value: Any, base: Any, default: Any) -> Any:
"""Three-way option resolution: flag > base config > default.
Args:
value: The parsed CLI flag (``None`` when not passed).
base: The corresponding value from a base config (``None``
without one).
default: The spec default used last.
Returns:
The first non-``None`` of the three.
"""
if value is not None:
return value
if base is not None:
return base
return default
def _parse_genera(raw: str | None) -> tuple[str, ...]:
"""Parse a comma-separated genus list.
Args:
raw: e.g. ``"Bacillus,Listeria,Staphylococcus"``; ``None`` means
no restriction (use every genus in the store).
Returns:
The trimmed non-empty names as a tuple.
Raises:
ValueError: If the string contains no non-empty names.
"""
if raw is None:
return ()
genera = tuple(g.strip() for g in raw.split(",") if g.strip())
if not genera:
raise ValueError("--genera parsed to an empty list")
return genera
def _run_config_from_args(args: argparse.Namespace) -> RunConfig:
"""Assemble a :class:`RunConfig` from flags plus optional base config.
Resolution order per field: explicit flag, then the base config's
value (when ``--config`` was given), then the spec default.
``--control`` or ``--stage control`` marks the run as the
shuffled-label control.
Args:
args: Parsed ``train`` arguments.
Returns:
The fully resolved run configuration.
Raises:
ValueError: If the base file is not a ``RunConfig`` or neither
flags nor base supply ``--store`` / ``--out-dir``.
"""
base: RunConfig | None = None
if args.config:
base = load_config(args.config)
if not isinstance(base, RunConfig):
raise ValueError(f"{args.config} does not hold a RunConfig")
base_model = base.model if base else None
base_data = base.data if base else None
base_train = base.train if base else None
model = ModelConfig(
arch=_pick(args.arch, base_model.arch if base_model else None, "cnn"),
param_budget=_pick(
args.budget, base_model.param_budget if base_model else None, 1_000_000
),
stride=_pick(args.stride, base_model.stride if base_model else None, 4),
d_embed=_pick(args.d_embed, base_model.d_embed if base_model else None, 256),
budget_tol=_pick(
args.budget_tol, base_model.budget_tol if base_model else None, 0.10
),
n_heads=_pick(args.n_heads, base_model.n_heads if base_model else None, 4),
)
data = DataConfig(
genera=(
_parse_genera(args.genera)
if args.genera is not None
else (base_data.genera if base_data else ())
),
val_fraction=_pick(
args.val_frac, base_data.val_fraction if base_data else None, 0.10
),
test_fraction=_pick(
args.test_frac, base_data.test_fraction if base_data else None, 0.10
),
holdout_runs=_pick(
args.holdout_runs, base_data.holdout_runs if base_data else None, False
),
min_test_windows_per_genus=_pick(
args.min_test_windows,
base_data.min_test_windows_per_genus if base_data else None,
300,
),
seed=_pick(args.data_seed, base_data.seed if base_data else None, 0),
)
stage = base.stage if base else "arch_ladder"
if args.control:
stage = "control"
elif args.stage:
stage = args.stage
device = _pick(args.device, base_train.device if base_train else None, None)
if device in ("auto", ""):
device = None
train = TrainConfig(
batch_size=_pick(args.batch, base_train.batch_size if base_train else None, 256),
lr=_pick(args.lr, base_train.lr if base_train else None, 3e-4),
weight_decay=_pick(
args.weight_decay, base_train.weight_decay if base_train else None, 0.01
),
max_epochs=_pick(
args.epochs, base_train.max_epochs if base_train else None, 40
),
warmup_frac=_pick(
args.warmup, base_train.warmup_frac if base_train else None, 0.05
),
label_smoothing=_pick(
args.smoothing, base_train.label_smoothing if base_train else None, 0.0
),
amp=_pick(args.amp, base_train.amp if base_train else None, "bf16"),
grad_clip=_pick(
args.grad_clip, base_train.grad_clip if base_train else None, 1.0
),
early_stop_metric=_pick(
args.metric, base_train.early_stop_metric if base_train else None, "val_recall_macro"
),
patience=_pick(args.patience, base_train.patience if base_train else None, 5),
num_workers=_pick(
args.workers, base_train.num_workers if base_train else None, 4
),
device=device,
seed=_pick(args.seed, base_train.seed if base_train else None, 0),
deterministic=_pick(
args.deterministic, base_train.deterministic if base_train else None, True
),
)
store_dir = args.store or (base.store_dir if base else None)
if store_dir is None:
raise ValueError("--store (or a base --config with a store_dir) is required")
out_dir = args.out_dir or (base.out_dir if base else None)
if out_dir is None:
raise ValueError("--out-dir (or a base --config with an out_dir) is required")
run_id = args.run_id or (base.run_id if base else "") or ""
return RunConfig(
run_id=run_id,
stage=stage,
store_dir=Path(store_dir),
data=data,
model=model,
train=train,
out_dir=Path(out_dir),
notes=args.notes or (base.notes if base else "") or "",
)
def _cmd_build(args: argparse.Namespace) -> int:
"""Handle ``build``: resolve a budget and print the realised geometry.
Args:
args: Parsed arguments (arch, budget, stride, d-embed,
budget-tol, n-heads, n-classes, seed).
Returns:
Exit code 0 on success.
"""
from .models import build_probe, count_params
from .seed import seed_everything
seed_everything(args.seed)
cfg = ModelConfig(
arch=args.arch,
param_budget=args.budget,
stride=args.stride if args.stride is not None else 4,
d_embed=args.d_embed if args.d_embed is not None else 256,
budget_tol=args.budget_tol if args.budget_tol is not None else 0.10,
n_heads=args.n_heads if args.n_heads is not None else 4,
)
probe, realised = build_probe(cfg, args.n_classes)
payload = {
"arch": realised.arch,
"param_budget": realised.param_budget,
"params_realised": realised.params_realised,
"stride": realised.stride,
"patch_len": min(4, realised.stride),
"d_model": realised.d_model,
"n_layers": realised.n_layers,
"n_heads": realised.n_heads,
"d_embed": realised.d_embed,
"n_classes": args.n_classes,
"count_params": count_params(probe),
}
print(json.dumps(payload, indent=2, sort_keys=True))
return 0
def _cmd_store(args: argparse.Namespace) -> int:
"""Handle ``store``: build a signal store and print its meta.
Args:
args: Parsed arguments; either ``--synthetic`` or both ``--pod5``
and ``--labels`` must be given, plus windowing flags.
Returns:
Exit code 0 on success.
Raises:
ValueError: On an incomplete mode selection (no ``--synthetic``
and no pod5/labels pair).
"""
from .store import extract_store, load_labels, synthetic_store
out = Path(args.out)
if args.synthetic:
meta = synthetic_store(
out_dir=out,
n_genera=args.n_genera,
reads_per_genus=args.reads_per_genus,
window_samples=args.window_samples,
period=args.period,
noise=args.noise,
seed=args.seed,
shard_windows=args.shard_windows,
genus_prefix=args.genus_prefix,
)
elif args.pod5 and args.labels:
cfg = ExtractConfig(
pod5_paths=tuple(Path(p) for p in args.pod5),
window_samples=args.window_samples,
skip_head_samples=args.skip_head,
max_windows_per_read=args.max_windows_per_read,
min_read_samples=args.min_read_samples,
norm=args.norm,
dtype=args.dtype,
shard_windows=args.shard_windows,
seed=args.seed,
)
labels = load_labels(Path(args.labels))
meta = extract_store(cfg, labels, out)
else:
raise ValueError(
"store needs either --synthetic or both --pod5 PATHS... and --labels PARQUET"
)
print(json.dumps(meta, indent=2, sort_keys=True, default=str))
return 0
def _cmd_train(args: argparse.Namespace) -> int:
"""Handle ``train``: run one training run and print its summary.
Args:
args: Parsed arguments (see :func:`_run_config_from_args`).
Returns:
Exit code 0 on success.
"""
cfg = _run_config_from_args(args)
realised, result, _ = train_run(cfg)
tail = result.history.row(result.history.height - 1, named=True)
metric = (
tail["val_recall_macro"]
if realised.train.early_stop_metric == "val_recall_macro"
else tail["val_loss"]
)
run_dir = realised.out_dir / realised.run_id
print(
json.dumps(
{
"run_id": realised.run_id,
"run_dir": str(run_dir),
"arch": realised.model.arch,
"d_model": realised.model.d_model,
"n_layers": realised.model.n_layers,
"params_realised": realised.model.params_realised,
"best_epoch": result.best_epoch,
f"best_{realised.train.early_stop_metric}": metric,
"seconds": round(result.seconds, 1),
},
indent=2,
sort_keys=True,
)
)
return 0
def _cmd_validate(args: argparse.Namespace) -> int:
"""Handle ``validate``: evaluate a run checkpoint on its val split.
Args:
args: Parsed arguments (``--run-dir``, optional ``--store``,
``--n-boot``).
Returns:
Exit code 0 on success.
"""
payload = evaluate_val(
Path(args.run_dir), store_dir=Path(args.store) if args.store else None, n_boot=args.n_boot
)
print(json.dumps(payload, indent=2, sort_keys=True))
return 0
def _cmd_test(args: argparse.Namespace) -> int:
"""Handle ``test``: evaluate a run checkpoint on its test split.
Args:
args: Parsed arguments (``--run-dir``, optional ``--store`` and
``--trap-store``, ``--n-boot``).
Returns:
Exit code 0 on success.
"""
payload = evaluate_test(
Path(args.run_dir),
store_dir=Path(args.store) if args.store else None,
trap_store_dir=Path(args.trap_store) if args.trap_store else None,
n_boot=args.n_boot,
)
print(json.dumps(payload, indent=2, sort_keys=True))
return 0
def _cmd_gate(args: argparse.Namespace) -> int:
"""Handle ``gate``: compute and print the E1-v2 gate verdict.
Args:
args: Parsed arguments (``--runs``, ``--control``,
``--chance-tol``).
Returns:
Exit code 0 on success.
"""
verdict = evaluate_gate([Path(r) for r in args.runs], Path(args.control), args.chance_tol)
print(json.dumps(verdict.__dict__, indent=2, sort_keys=True))
return 0
def _cmd_configs(args: argparse.Namespace) -> int:
"""Handle ``configs``: load a config JSON and re-save it elsewhere.
A round-trip helper proving ``load_config(save_config(c)) == c`` and
letting users materialise a base config from an existing run.
Args:
args: Parsed arguments (``--config`` source, ``--out``
destination).
Returns:
Exit code 0 on success.
"""
cfg = load_config(args.config)
save_config(cfg, Path(args.out))
print(f"wrote {args.out}")
return 0
def build_parser() -> argparse.ArgumentParser:
"""Construct the full CLI parser with all subcommands.
Returns:
The configured ``argparse.ArgumentParser`` (``--help`` lists
every subcommand and flag).
"""
parser = argparse.ArgumentParser(
prog="custom_models",
description="E1 genus probe: torch models and train/validate/test wrapper",
)
sub = parser.add_subparsers(dest="command", required=True)
p_build = sub.add_parser("build", help="resolve a param budget into a realised probe")
_add_model_args(p_build, required=True)
p_build.add_argument("--n-classes", type=int, default=11)
p_build.add_argument("--seed", type=int, default=0)
p_build.set_defaults(func=_cmd_build)
p_store = sub.add_parser("store", help="build a signal store (pod5+labels or synthetic)")
p_store.add_argument("--out", required=True)
p_store.add_argument("--synthetic", action="store_true")
p_store.add_argument("--n-genera", type=int, default=11)
p_store.add_argument("--reads-per-genus", type=int, default=200)
p_store.add_argument("--period", type=int, default=None)
p_store.add_argument("--noise", type=float, default=0.15)
p_store.add_argument("--genus-prefix", default="genus_")
p_store.add_argument("--pod5", nargs="+", default=None)
p_store.add_argument("--labels", default=None)
p_store.add_argument("--window-samples", type=int, default=12_000)
p_store.add_argument("--skip-head", type=int, default=500)
p_store.add_argument("--min-read-samples", type=int, default=12_500)
p_store.add_argument("--max-windows-per-read", type=int, default=1)
p_store.add_argument("--norm", choices=["median_iqr", "median_mad"], default="median_iqr")
p_store.add_argument("--dtype", choices=["float16", "float32"], default="float16")
p_store.add_argument("--shard-windows", type=int, default=8_192)
p_store.add_argument("--seed", type=int, default=0)
p_store.set_defaults(func=_cmd_store)
p_train = sub.add_parser("train", help="train a probe on a signal store")
p_train.add_argument("--store", default=None)
p_train.add_argument("--out-dir", default=None)
p_train.add_argument("--config", default=None)
p_train.add_argument("--run-id", default=None)
p_train.add_argument(
"--stage",
choices=["arch_ladder", "windows_ablation", "robustness", "control", "extra"],
default=None,
)
p_train.add_argument("--control", action="store_true", default=None)
p_train.add_argument("--notes", default=None)
_add_model_args(p_train, required=False)
_add_data_args(p_train)
_add_train_args(p_train)
p_train.set_defaults(func=_cmd_train)
p_validate = sub.add_parser("validate", help="evaluate a run checkpoint on its val split")
p_validate.add_argument("--run-dir", required=True)
p_validate.add_argument("--store", default=None)
p_validate.add_argument("--n-boot", type=int, default=10_000)
p_validate.set_defaults(func=_cmd_validate)
p_test = sub.add_parser("test", help="evaluate a run checkpoint on its test split")
p_test.add_argument("--run-dir", required=True)
p_test.add_argument("--store", default=None)
p_test.add_argument("--trap-store", default=None)
p_test.add_argument("--n-boot", type=int, default=10_000)
p_test.set_defaults(func=_cmd_test)
p_gate = sub.add_parser("gate", help="gate verdict over run dirs against a control run")
p_gate.add_argument("--runs", nargs="+", required=True)
p_gate.add_argument("--control", required=True)
p_gate.add_argument("--chance-tol", type=float, default=1.5)
p_gate.set_defaults(func=_cmd_gate)
p_configs = sub.add_parser("configs", help="round-trip a config JSON (save_config/load_config)")
p_configs.add_argument("--config", required=True)
p_configs.add_argument("--out", required=True)
p_configs.set_defaults(func=_cmd_configs)
return parser
def main(argv: Sequence[str] | None = None) -> int:
"""CLI entry point (also ``python -m custom_models``).
Args:
argv: Argument list without the program name; ``None`` uses
``sys.argv``.
Returns:
0 on success; 1 on usage or runtime errors (message printed to
stderr).
"""
parser = build_parser()
args = parser.parse_args(list(argv) if argv is not None else None)
try:
return int(args.func(args))
except (ValueError, RuntimeError, FileNotFoundError, ImportError) as exc:
print(f"error: {exc}", file=sys.stderr)
return 1
if __name__ == "__main__":
raise SystemExit(main())