"""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, out). 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), } text = json.dumps(payload, indent=2, sort_keys=True) if args.out: out = Path(args.out) out.parent.mkdir(parents=True, exist_ok=True) out.write_text(text + "\n", encoding="utf-8") print(f"wrote {out}") else: print(text) 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.add_argument("--out", default=None, help="Write the JSON payload to PATH instead of stdout") 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())