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:
@@ -0,0 +1,78 @@
|
||||
"""E1 genus probe: torch models, training wrapper and evaluation (v0.1.0).
|
||||
|
||||
Public API re-exports the pieces a run needs: configs (``config``),
|
||||
determinism (``seed``), the signal store (``store``), splits and loaders
|
||||
(``data``), the torch modules and budget resolver (``models``), the
|
||||
training loop and inference (``train``), metrics (``metrics``), run
|
||||
orchestration and the gate (``runner``) and the CLI (``cli``). The
|
||||
normative contracts live in ``planning/design_docs/preliminary_
|
||||
experiments/E1_torch_model_spec.md`` and ``E1_v2_eleven_taxa_redesign.md``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from .config import (
|
||||
DataConfig,
|
||||
ExtractConfig,
|
||||
ModelConfig,
|
||||
RunConfig,
|
||||
TrainConfig,
|
||||
load_config,
|
||||
run_id_for,
|
||||
save_config,
|
||||
)
|
||||
from .data import ProbeDataset, Splits, make_splits, permute_read_labels
|
||||
from .models import (
|
||||
ConvEncoder,
|
||||
GenusProbe,
|
||||
LinearAttentionEncoder,
|
||||
SignalPatchEmbed,
|
||||
build_probe,
|
||||
count_params,
|
||||
)
|
||||
from .seed import make_generator, seed_everything, worker_init_fn
|
||||
from .store import SignalStore, extract_store, load_labels, synthetic_store
|
||||
from .train import (
|
||||
TrainResult,
|
||||
build_optimizer,
|
||||
embed_probe,
|
||||
predict,
|
||||
select_device,
|
||||
train_probe,
|
||||
)
|
||||
|
||||
__version__ = "0.1.0"
|
||||
|
||||
__all__ = [
|
||||
"ConvEncoder",
|
||||
"DataConfig",
|
||||
"ExtractConfig",
|
||||
"GenusProbe",
|
||||
"LinearAttentionEncoder",
|
||||
"ModelConfig",
|
||||
"ProbeDataset",
|
||||
"RunConfig",
|
||||
"SignalPatchEmbed",
|
||||
"SignalStore",
|
||||
"Splits",
|
||||
"TrainConfig",
|
||||
"TrainResult",
|
||||
"build_optimizer",
|
||||
"build_probe",
|
||||
"count_params",
|
||||
"embed_probe",
|
||||
"extract_store",
|
||||
"load_config",
|
||||
"load_labels",
|
||||
"make_generator",
|
||||
"make_splits",
|
||||
"permute_read_labels",
|
||||
"predict",
|
||||
"run_id_for",
|
||||
"save_config",
|
||||
"seed_everything",
|
||||
"select_device",
|
||||
"synthetic_store",
|
||||
"train_probe",
|
||||
"worker_init_fn",
|
||||
]
|
||||
@@ -0,0 +1,8 @@
|
||||
"""Module entry point: ``python -m custom_models <subcommand>``."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from .cli import main
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -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())
|
||||
@@ -0,0 +1,273 @@
|
||||
"""Frozen configuration dataclasses and JSON (de)serialisation.
|
||||
|
||||
The dataclasses mirror the GATE ZERO config contract (``E1_genus_information_
|
||||
ceiling.md`` §5.1) with the E1-v2 deltas (fixed 11-taxa membership, no
|
||||
class-ladder machinery) and the binding deviations recorded in its §12:
|
||||
``ModelConfig`` carries the realised fields ``d_model`` / ``n_layers`` /
|
||||
``params_realised`` that ``build_probe`` fills in, and ``RunConfig`` records
|
||||
the realised genera so a run directory round-trips deterministically.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass, is_dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any, Literal
|
||||
|
||||
Arch = Literal["cnn", "linatt"]
|
||||
Stage = Literal["arch_ladder", "windows_ablation", "robustness", "control", "extra"]
|
||||
AmpPolicy = Literal["off", "bf16", "fp16"]
|
||||
EarlyStopMetric = Literal["val_recall_macro", "val_loss"]
|
||||
NormMethod = Literal["median_iqr", "median_mad"]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class DataConfig:
|
||||
"""Split-time data selection and read-level split parameters."""
|
||||
|
||||
genera: tuple[str, ...] = ()
|
||||
val_fraction: float = 0.10
|
||||
test_fraction: float = 0.10
|
||||
holdout_runs: bool = False
|
||||
min_test_windows_per_genus: int = 300
|
||||
seed: int = 0
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ModelConfig:
|
||||
"""Probe architecture request; realised geometry is filled by ``build_probe``.
|
||||
|
||||
Attributes:
|
||||
arch: Encoder family, ``"cnn"`` or ``"linatt"``.
|
||||
param_budget: Nominal trainable-parameter budget in params (e.g.
|
||||
``1_000_000``). Only the nominal budget enters the run id; the
|
||||
gate arithmetic always uses ``params_realised``.
|
||||
stride: Patch stride in signal samples, 4 or 8 (token pitch).
|
||||
d_embed: Pooled embedding width the encoder terminates in.
|
||||
budget_tol: Relative tolerance around the budget; the realised count
|
||||
must satisfy ``count_params <= (1 + budget_tol) * param_budget``.
|
||||
n_heads: Attention head count for the ``"linatt"`` arch (recorded so
|
||||
checkpoint rehydration rebuilds the exact geometry).
|
||||
d_model: realised token width (filled by ``build_probe``, ``None``
|
||||
on a request config).
|
||||
n_layers: realised block count (filled by ``build_probe``).
|
||||
params_realised: realised trainable parameter count (filled by
|
||||
``build_probe``; the only source of truth for reporting).
|
||||
|
||||
"""
|
||||
|
||||
arch: Arch
|
||||
param_budget: int
|
||||
stride: int
|
||||
d_embed: int = 256
|
||||
budget_tol: float = 0.10
|
||||
n_heads: int = 4
|
||||
d_model: int | None = None
|
||||
n_layers: int | None = None
|
||||
params_realised: int | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class TrainConfig:
|
||||
"""Training semantics (GATE ZERO §5.1; restated in the torch model spec §5)."""
|
||||
|
||||
batch_size: int = 256
|
||||
lr: float = 3e-4
|
||||
weight_decay: float = 0.01
|
||||
max_epochs: int = 40
|
||||
warmup_frac: float = 0.05
|
||||
label_smoothing: float = 0.0
|
||||
amp: AmpPolicy = "bf16"
|
||||
grad_clip: float = 1.0
|
||||
early_stop_metric: EarlyStopMetric = "val_recall_macro"
|
||||
patience: int = 5
|
||||
num_workers: int = 4
|
||||
device: str | None = None
|
||||
seed: int = 0
|
||||
deterministic: bool = True
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ExtractConfig:
|
||||
"""Pod5-to-store extraction parameters for ``extract_store``.
|
||||
|
||||
Attributes:
|
||||
pod5_paths: Pod5 files to scan, processed in sorted path order.
|
||||
window_samples: Fixed window length in samples (12,000 ≈ 2.4 s at
|
||||
5 kHz ≈ ~1 kb) — the model-input contract.
|
||||
skip_head_samples: Head samples skipped before the first window
|
||||
(adapter / mux artefacts).
|
||||
max_windows_per_read: Windows taken per read; 1 is the headline, 4 is
|
||||
the Stage-A' ablation row. Windows are contiguous from the first.
|
||||
min_read_samples: Reads shorter than this are dropped and counted in
|
||||
the store meta (must fit skip-head plus one full window).
|
||||
norm: Per-window normaliser applied before storing.
|
||||
dtype: On-disk shard dtype; fp16 halves the store at negligible cost.
|
||||
shard_windows: Windows per ``shard_%05d.npy`` file.
|
||||
seed: Unused by extraction itself (window choice is deterministic);
|
||||
kept so every config in the pipeline is seed-carrying.
|
||||
|
||||
"""
|
||||
|
||||
pod5_paths: tuple[Path, ...]
|
||||
window_samples: int = 12_000
|
||||
skip_head_samples: int = 500
|
||||
max_windows_per_read: int = 1
|
||||
min_read_samples: int = 12_500
|
||||
norm: NormMethod = "median_iqr"
|
||||
dtype: Literal["float16", "float32"] = "float16"
|
||||
shard_windows: int = 8_192
|
||||
seed: int = 0
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RunConfig:
|
||||
"""One experiment run: a store, data/model/train sub-configs and outputs.
|
||||
|
||||
``train_run`` fills ``run_id`` (if empty) and ``genera`` (from the realised
|
||||
splits) before the realised config is serialised into the run directory.
|
||||
"""
|
||||
|
||||
run_id: str
|
||||
stage: Stage
|
||||
store_dir: Path
|
||||
data: DataConfig
|
||||
model: ModelConfig
|
||||
train: TrainConfig
|
||||
out_dir: Path
|
||||
genera: tuple[str, ...] = ()
|
||||
notes: str = ""
|
||||
|
||||
|
||||
def run_id_for(arch: str, param_budget: int, stride: int, n_genera: int, stage: str) -> str:
|
||||
"""Build the deterministic, filesystem-safe run identifier.
|
||||
|
||||
Format: ``{arch}_{budget in M:.1f}M_s{stride}_g{n_genera}``, e.g.
|
||||
``cnn_1.0M_s4_g11``; the control stage appends ``-shuf``. Only the
|
||||
nominal budget and class count enter the id — realised geometry never
|
||||
does — so ids are unique within the run matrix and stable across
|
||||
rehydration.
|
||||
|
||||
Args:
|
||||
arch: Encoder family name (``"cnn"`` or ``"linatt"``).
|
||||
param_budget: Nominal parameter budget (divided by 1e6 for the tag).
|
||||
stride: Token stride in samples.
|
||||
n_genera: Number of selected classes (the ``g`` component).
|
||||
stage: Run stage; ``"control"`` triggers the ``-shuf`` suffix.
|
||||
|
||||
Returns:
|
||||
The run id string.
|
||||
|
||||
"""
|
||||
run_id = f"{arch}_{param_budget / 1_000_000:.1f}M_s{stride}_g{n_genera}"
|
||||
if stage == "control":
|
||||
return f"{run_id}-shuf"
|
||||
return run_id
|
||||
|
||||
|
||||
_TYPE_KEY = "__type__"
|
||||
_DATACLASSES: dict[str, type] = {
|
||||
cls.__name__: cls
|
||||
for cls in (DataConfig, ModelConfig, TrainConfig, ExtractConfig, RunConfig)
|
||||
}
|
||||
|
||||
|
||||
def _encode(obj: Any) -> Any:
|
||||
"""Recursively convert a config object into a JSON-friendly payload.
|
||||
|
||||
``Path`` becomes ``{"__type__": "path", "value": str}``, tuples become
|
||||
tagged sequences (so they round-trip as tuples), and dataclass instances
|
||||
become dicts tagged with their class name. All other values pass
|
||||
through unchanged.
|
||||
|
||||
Args:
|
||||
obj: A config dataclass, ``Path``, tuple/list, or plain JSON value.
|
||||
|
||||
Returns:
|
||||
The JSON-serialisable encoding of ``obj``.
|
||||
|
||||
"""
|
||||
if isinstance(obj, Path):
|
||||
return {_TYPE_KEY: "path", "value": str(obj)}
|
||||
if isinstance(obj, (tuple, list)):
|
||||
return {_TYPE_KEY: "seq", "value": [_encode(v) for v in obj]}
|
||||
if is_dataclass(obj) and not isinstance(obj, type):
|
||||
payload: dict[str, Any] = {_TYPE_KEY: type(obj).__name__}
|
||||
for field in obj.__dataclass_fields__:
|
||||
payload[field] = _encode(getattr(obj, field))
|
||||
return payload
|
||||
return obj
|
||||
|
||||
|
||||
def _decode(obj: Any) -> Any:
|
||||
"""Invert :func:`_encode` back into config objects.
|
||||
|
||||
Tagged dicts rebuild ``Path`` values, tuples and registered dataclasses;
|
||||
anything else is returned as-is. Unknown ``__type__`` tags fall through
|
||||
as plain dicts so foreign JSON never crashes decoding.
|
||||
|
||||
Args:
|
||||
obj: A payload produced by :func:`_encode` (or nested JSON value).
|
||||
|
||||
Returns:
|
||||
The reconstructed config object.
|
||||
|
||||
"""
|
||||
if isinstance(obj, dict):
|
||||
kind = obj.get(_TYPE_KEY)
|
||||
if kind == "path":
|
||||
return Path(obj["value"])
|
||||
if kind == "seq":
|
||||
return tuple(_decode(v) for v in obj["value"])
|
||||
if kind in _DATACLASSES:
|
||||
cls = _DATACLASSES[kind]
|
||||
kwargs = {}
|
||||
for name, field in cls.__dataclass_fields__.items():
|
||||
if name in obj:
|
||||
kwargs[name] = _decode(obj[name])
|
||||
elif field.default is not None:
|
||||
pass
|
||||
return cls(**kwargs)
|
||||
return obj
|
||||
|
||||
|
||||
def save_config(cfg: Any, path: Path) -> None:
|
||||
"""Serialise a config dataclass to pretty JSON.
|
||||
|
||||
Parent directories are created as needed. The payload uses the
|
||||
``"__type__"`` discriminator scheme (GATE ZERO §12.7) so
|
||||
``load_config(save_config(c)) == c`` holds exactly.
|
||||
|
||||
Args:
|
||||
cfg: A config dataclass (or nested structure of them).
|
||||
path: Destination JSON file path.
|
||||
|
||||
"""
|
||||
path = Path(path)
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_text(json.dumps(_encode(cfg), indent=2, sort_keys=True) + "\n")
|
||||
|
||||
|
||||
def load_config(path: Path) -> Any:
|
||||
"""Load a config dataclass written by :func:`save_config`.
|
||||
|
||||
Args:
|
||||
path: JSON config file path.
|
||||
|
||||
Returns:
|
||||
The reconstructed config object (type depends on the file).
|
||||
|
||||
Raises:
|
||||
FileNotFoundError: If ``path`` does not exist.
|
||||
ValueError: If the file is not valid JSON or does not decode into
|
||||
the recorded dataclass schema.
|
||||
|
||||
"""
|
||||
path = Path(path)
|
||||
if not path.is_file():
|
||||
raise FileNotFoundError(f"config file not found: {path}")
|
||||
try:
|
||||
return _decode(json.loads(path.read_text()))
|
||||
except (KeyError, TypeError, json.JSONDecodeError) as exc:
|
||||
raise ValueError(f"malformed config file: {path}: {exc}") from exc
|
||||
@@ -0,0 +1,406 @@
|
||||
"""Splits, label permutation for control runs, datasets and data loaders.
|
||||
|
||||
Read-level stratified splitting (all windows of a read share a split),
|
||||
the v2 near-clade safety minimum on per-genus test windows, whole-run
|
||||
holdout mode, the seeded shuffled-label permutation used by control runs,
|
||||
and the torch ``Dataset``/``DataLoader`` plumbing with device-appropriate
|
||||
worker/pinning behaviour.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from dataclasses import dataclass
|
||||
|
||||
import numpy as np
|
||||
import polars as pl
|
||||
import torch
|
||||
from torch.utils.data import DataLoader
|
||||
|
||||
from .config import DataConfig, TrainConfig
|
||||
from .seed import make_generator, worker_init_fn
|
||||
from .store import SignalStore
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Splits:
|
||||
"""Window-index splits plus the full-store label vector.
|
||||
|
||||
Attributes:
|
||||
train: Sorted window ids assigned to training.
|
||||
val: Sorted window ids assigned to validation.
|
||||
test: Sorted window ids assigned to testing.
|
||||
genera: Sorted selected genera; class id = position in this list
|
||||
everywhere downstream.
|
||||
labels: int32 class id for every store window (length = store
|
||||
size), ``-1`` for windows outside the selected genera.
|
||||
|
||||
"""
|
||||
|
||||
train: np.ndarray
|
||||
val: np.ndarray
|
||||
test: np.ndarray
|
||||
genera: list[str]
|
||||
labels: np.ndarray
|
||||
|
||||
|
||||
def make_splits(manifest: pl.DataFrame, cfg: DataConfig) -> Splits:
|
||||
"""Compute deterministic read-level stratified splits from a manifest.
|
||||
|
||||
Selection: ``cfg.genera`` if given (sorted, must all exist in the
|
||||
store — the fixed owner list in E1 v2), else every genus present.
|
||||
Splitting: per genus, reads are shuffled with the seeded generator and
|
||||
assigned by read counts (test first, then val, remainder train), so
|
||||
all windows of a read land in one split. ``holdout_runs`` instead
|
||||
assigns whole ``source_run`` groups to test (seeded shuffle until the
|
||||
test fraction is met) and splits the remaining reads into
|
||||
train/val, requiring every genus to retain both train and test
|
||||
presence. The v2 near-clade guard enforces
|
||||
``min_test_windows_per_genus`` test windows per genus.
|
||||
|
||||
Args:
|
||||
manifest: Store manifest (one row per window; needs ``window_id``,
|
||||
``read_id``, ``genus``, ``source_run``).
|
||||
cfg: Split parameters (genera, fractions, holdout mode, minimum,
|
||||
seed).
|
||||
|
||||
Returns:
|
||||
The splits (see :class:`Splits`).
|
||||
|
||||
Raises:
|
||||
ValueError: On missing manifest columns, an empty manifest,
|
||||
requested genera absent from the store, an empty selection or
|
||||
train split, holdout presence failure, or a per-genus test
|
||||
count below ``min_test_windows_per_genus`` (v2: raise instead
|
||||
of gating on a noisy plan).
|
||||
|
||||
"""
|
||||
for column in ("window_id", "read_id", "genus", "source_run"):
|
||||
if column not in manifest.columns:
|
||||
raise ValueError(f"manifest missing column: {column}")
|
||||
n = manifest.height
|
||||
if n == 0:
|
||||
raise ValueError("manifest is empty")
|
||||
genus_arr = manifest["genus"].to_numpy()
|
||||
read_arr = manifest["read_id"].to_numpy()
|
||||
wid_arr = manifest["window_id"].to_numpy()
|
||||
run_arr = manifest["source_run"].to_numpy()
|
||||
all_genera = sorted(set(genus_arr.tolist()))
|
||||
selected = sorted(set(cfg.genera)) if cfg.genera else list(all_genera)
|
||||
if not selected:
|
||||
raise ValueError("no genera available for splitting")
|
||||
missing = [g for g in selected if g not in set(all_genera)]
|
||||
if missing:
|
||||
raise ValueError(f"genera absent from the store: {missing}")
|
||||
lookup = {g: i for i, g in enumerate(selected)}
|
||||
n_classes = len(selected)
|
||||
labels = np.full(n, -1, dtype=np.int32)
|
||||
for genus, class_id in lookup.items():
|
||||
labels[genus_arr == genus] = class_id
|
||||
selected_mask = labels >= 0
|
||||
if not selected_mask.any():
|
||||
raise ValueError("no windows belong to the selected genera")
|
||||
|
||||
rng = make_generator(cfg.seed)
|
||||
windows_by_read: dict[str, list[int]] = {}
|
||||
genus_by_read: dict[str, str] = {}
|
||||
run_by_read: dict[str, str] = {}
|
||||
for wid, read_id, genus, source_run in zip(
|
||||
wid_arr[selected_mask], read_arr[selected_mask], genus_arr[selected_mask],
|
||||
run_arr[selected_mask]
|
||||
):
|
||||
windows_by_read.setdefault(read_id, []).append(int(wid))
|
||||
genus_by_read[read_id] = genus
|
||||
run_by_read[read_id] = source_run
|
||||
|
||||
reads_by_genus: dict[str, list[str]] = {g: [] for g in selected}
|
||||
for read_id, genus in genus_by_read.items():
|
||||
reads_by_genus[genus].append(read_id)
|
||||
for reads in reads_by_genus.values():
|
||||
reads.sort()
|
||||
|
||||
train_ids: list[int] = []
|
||||
val_ids: list[int] = []
|
||||
test_ids: list[int] = []
|
||||
|
||||
if cfg.holdout_runs:
|
||||
runs = sorted(set(run_by_read.values()))
|
||||
windows_per_run: dict[str, int] = {}
|
||||
for read_id, source_run in run_by_read.items():
|
||||
windows_per_run[source_run] = windows_per_run.get(source_run, 0) + len(
|
||||
windows_by_read[read_id]
|
||||
)
|
||||
order = rng.permutation(len(runs))
|
||||
shuffled_runs = [runs[i] for i in order]
|
||||
target_test = round(cfg.test_fraction * selected_mask.sum())
|
||||
chosen: set[str] = set()
|
||||
accrued = 0
|
||||
for source_run in shuffled_runs:
|
||||
if accrued >= target_test:
|
||||
break
|
||||
chosen.add(source_run)
|
||||
accrued += windows_per_run[source_run]
|
||||
for read_id, source_run in run_by_read.items():
|
||||
if source_run in chosen:
|
||||
test_ids.extend(windows_by_read[read_id])
|
||||
for genus in selected:
|
||||
reads = [r for r in reads_by_genus[genus] if run_by_read[r] not in chosen]
|
||||
order = rng.permutation(len(reads))
|
||||
shuffled = [reads[i] for i in order]
|
||||
n_val = round(cfg.val_fraction * len(shuffled))
|
||||
val_reads = shuffled[:n_val]
|
||||
train_reads = shuffled[n_val:]
|
||||
for r in val_reads:
|
||||
val_ids.extend(windows_by_read[r])
|
||||
for r in train_reads:
|
||||
train_ids.extend(windows_by_read[r])
|
||||
for genus in selected:
|
||||
has_train = any(run_by_read[r] not in chosen for r in reads_by_genus[genus])
|
||||
has_test = any(run_by_read[r] in chosen for r in reads_by_genus[genus])
|
||||
if not (has_train and has_test):
|
||||
raise ValueError(
|
||||
f"holdout_runs split leaves genus {genus!r} without train and test presence"
|
||||
)
|
||||
else:
|
||||
for genus in selected:
|
||||
reads = reads_by_genus[genus]
|
||||
order = rng.permutation(len(reads))
|
||||
shuffled = [reads[i] for i in order]
|
||||
n_test = round(cfg.test_fraction * len(shuffled))
|
||||
n_val = round(cfg.val_fraction * len(shuffled))
|
||||
test_reads = shuffled[:n_test]
|
||||
val_reads = shuffled[n_test : n_test + n_val]
|
||||
train_reads = shuffled[n_test + n_val :]
|
||||
for r in test_reads:
|
||||
test_ids.extend(windows_by_read[r])
|
||||
for r in val_reads:
|
||||
val_ids.extend(windows_by_read[r])
|
||||
for r in train_reads:
|
||||
train_ids.extend(windows_by_read[r])
|
||||
|
||||
if not train_ids:
|
||||
raise ValueError("train split is empty after splitting")
|
||||
train = np.array(sorted(train_ids), dtype=np.int64)
|
||||
val = np.array(sorted(val_ids), dtype=np.int64)
|
||||
test = np.array(sorted(test_ids), dtype=np.int64)
|
||||
if cfg.min_test_windows_per_genus > 0:
|
||||
if len(test) == 0:
|
||||
raise ValueError(
|
||||
"test split is empty; per-genus minimum "
|
||||
f"min_test_windows_per_genus={cfg.min_test_windows_per_genus} cannot be met"
|
||||
)
|
||||
counts = np.bincount(labels[test], minlength=n_classes)
|
||||
short = {
|
||||
selected[i]: int(c)
|
||||
for i, c in enumerate(counts)
|
||||
if 0 < c < cfg.min_test_windows_per_genus
|
||||
}
|
||||
absent = [g for g in selected if not np.any(labels[test] == lookup[g])]
|
||||
if short or absent:
|
||||
raise ValueError(
|
||||
"per-genus test windows below the near-clade safety minimum "
|
||||
f"(min_test_windows_per_genus={cfg.min_test_windows_per_genus}): "
|
||||
f"short={short} absent={absent}"
|
||||
)
|
||||
return Splits(train=train, val=val, test=test, genera=selected, labels=labels)
|
||||
|
||||
|
||||
def permute_read_labels(
|
||||
manifest: pl.DataFrame, genera: tuple[str, ...], seed: int
|
||||
) -> pl.DataFrame:
|
||||
"""Shuffle genus labels across reads for the Stage-C control run.
|
||||
|
||||
The read-to-genus mapping of the selected pool is permuted (seeded,
|
||||
preserving the genus multiset exactly); every window of a read
|
||||
inherits its read's new label. Applied *before* splitting so the
|
||||
control run's splits and labels live in the same permuted world as
|
||||
its training (evaluation re-permutes identically).
|
||||
|
||||
Args:
|
||||
manifest: Store manifest to relabel.
|
||||
genera: Restriction of the permutation pool; empty means all
|
||||
genera present.
|
||||
seed: Seed for the permutation.
|
||||
|
||||
Returns:
|
||||
A new manifest with the ``genus`` column permuted.
|
||||
|
||||
Raises:
|
||||
ValueError: If the pool holds fewer than two reads.
|
||||
|
||||
"""
|
||||
genus_arr = manifest["genus"].to_numpy()
|
||||
read_arr = manifest["read_id"].to_numpy()
|
||||
genus_by_read: dict[str, str] = {}
|
||||
for read_id, genus in zip(read_arr, genus_arr):
|
||||
genus_by_read.setdefault(read_id, genus)
|
||||
pool = sorted(genus_by_read)
|
||||
if genera:
|
||||
pool = [r for r in pool if genus_by_read[r] in set(genera)]
|
||||
if len(pool) < 2:
|
||||
raise ValueError("control run needs at least two reads in the selected pool")
|
||||
values = [genus_by_read[r] for r in pool]
|
||||
rng = make_generator(seed)
|
||||
order = rng.permutation(len(values))
|
||||
permuted = {read_id: values[i] for read_id, i in zip(pool, order)}
|
||||
new_genus = [permuted[read_id] for read_id in read_arr]
|
||||
return manifest.with_columns(pl.Series("genus", new_genus, dtype=pl.String))
|
||||
|
||||
|
||||
class ProbeDataset(torch.utils.data.Dataset):
|
||||
"""Window dataset over a subset of store window ids."""
|
||||
|
||||
def __init__(self, store: SignalStore, indices: np.ndarray, labels: np.ndarray) -> None:
|
||||
"""Bind a store, window ids and the full-store label vector.
|
||||
|
||||
Args:
|
||||
store: Open signal store serving the windows.
|
||||
indices: Window ids to serve (a split, or all ids).
|
||||
labels: Full-store label vector indexed by window id
|
||||
(from :class:`Splits`).
|
||||
|
||||
"""
|
||||
self.store = store
|
||||
self.indices = np.asarray(indices, dtype=np.int64)
|
||||
self.labels = np.asarray(labels, dtype=np.int64)
|
||||
|
||||
def __len__(self) -> int:
|
||||
"""Return the number of windows in this dataset."""
|
||||
return len(self.indices)
|
||||
|
||||
def __getitem__(self, index: int) -> tuple[np.ndarray, int]:
|
||||
"""Fetch one window and its class id.
|
||||
|
||||
Args:
|
||||
index: Dataset row index (not the window id).
|
||||
|
||||
Returns:
|
||||
``(float32 window of shape (L,), class id)``.
|
||||
|
||||
"""
|
||||
window_id = int(self.indices[index])
|
||||
return self.store.get(window_id), int(self.labels[window_id])
|
||||
|
||||
|
||||
def collate_windows(
|
||||
batch: list[tuple[np.ndarray, int]],
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Stack dataset samples into model-ready tensors.
|
||||
|
||||
Args:
|
||||
batch: List of ``(window, class_id)`` pairs from
|
||||
:class:`ProbeDataset`.
|
||||
|
||||
Returns:
|
||||
``(windows (B, L) float32, labels (B,) int64)`` — the probe's
|
||||
input contract.
|
||||
|
||||
"""
|
||||
windows = torch.from_numpy(np.stack([np.asarray(b[0], dtype=np.float32) for b in batch]))
|
||||
labels = torch.tensor([int(b[1]) for b in batch], dtype=torch.int64)
|
||||
return windows, labels
|
||||
|
||||
|
||||
def _loader_mp_context():
|
||||
"""Pick a DataLoader multiprocessing start context for this platform.
|
||||
|
||||
Linux pins ``"fork"`` (torch's historical default) because conda's
|
||||
Python 3.14 builds ship a broken ``forkserver``; other platforms keep
|
||||
the interpreter default (spawn on macOS/Windows).
|
||||
|
||||
Returns:
|
||||
A start-method name or ``None`` for the platform default.
|
||||
|
||||
"""
|
||||
if sys.platform == "linux":
|
||||
return "fork"
|
||||
return None
|
||||
|
||||
|
||||
def make_loaders(
|
||||
store: SignalStore, splits: Splits, cfg: TrainConfig
|
||||
) -> tuple[DataLoader, DataLoader, DataLoader]:
|
||||
"""Build the train/val/test loaders for one run.
|
||||
|
||||
Train: shuffled over a generator seeded from ``cfg.seed`` with
|
||||
``drop_last=True`` (stable step counts) and per-worker seeding.
|
||||
Val/test: unshuffled and complete. ``pin_memory`` only when the
|
||||
resolved device is CUDA; ``persistent_workers`` only when
|
||||
``num_workers > 0``.
|
||||
|
||||
Args:
|
||||
store: The signal store both splits read from.
|
||||
splits: Split window ids and labels.
|
||||
cfg: Training config (batch size, workers, seed, device).
|
||||
|
||||
Returns:
|
||||
``(train_loader, val_loader, test_loader)``.
|
||||
|
||||
"""
|
||||
from .train import select_device
|
||||
|
||||
device = select_device(cfg.device)
|
||||
pin = device.type == "cuda"
|
||||
generator = torch.Generator()
|
||||
generator.manual_seed(cfg.seed)
|
||||
common: dict = {
|
||||
"batch_size": cfg.batch_size,
|
||||
"collate_fn": collate_windows,
|
||||
"num_workers": cfg.num_workers,
|
||||
"pin_memory": pin,
|
||||
"persistent_workers": cfg.num_workers > 0,
|
||||
"worker_init_fn": worker_init_fn,
|
||||
"multiprocessing_context": _loader_mp_context() if cfg.num_workers > 0 else None,
|
||||
}
|
||||
train_loader = DataLoader(
|
||||
ProbeDataset(store, splits.train, splits.labels),
|
||||
shuffle=True,
|
||||
generator=generator,
|
||||
drop_last=True,
|
||||
**common,
|
||||
)
|
||||
val_loader = DataLoader(
|
||||
ProbeDataset(store, splits.val, splits.labels),
|
||||
shuffle=False,
|
||||
drop_last=False,
|
||||
**common,
|
||||
)
|
||||
test_loader = DataLoader(
|
||||
ProbeDataset(store, splits.test, splits.labels),
|
||||
shuffle=False,
|
||||
drop_last=False,
|
||||
**common,
|
||||
)
|
||||
return train_loader, val_loader, test_loader
|
||||
|
||||
|
||||
def make_full_loader(store: SignalStore, cfg: TrainConfig) -> DataLoader:
|
||||
"""Build an unshuffled loader over *every* window of a store.
|
||||
|
||||
Used for trap (TP2) evaluation where the store's genera are not
|
||||
classes of the probe; labels are placeholder ``-1`` values.
|
||||
|
||||
Args:
|
||||
store: Store to evaluate end-to-end.
|
||||
cfg: Training config supplying batching/worker/device behaviour.
|
||||
|
||||
Returns:
|
||||
A DataLoader yielding ``(windows, -1 labels)`` in window order.
|
||||
|
||||
"""
|
||||
from .train import select_device
|
||||
|
||||
device = select_device(cfg.device)
|
||||
labels = np.full(len(store), -1, dtype=np.int64)
|
||||
return DataLoader(
|
||||
ProbeDataset(store, np.arange(len(store), dtype=np.int64), labels),
|
||||
batch_size=cfg.batch_size,
|
||||
shuffle=False,
|
||||
drop_last=False,
|
||||
collate_fn=collate_windows,
|
||||
num_workers=cfg.num_workers,
|
||||
pin_memory=device.type == "cuda",
|
||||
persistent_workers=cfg.num_workers > 0,
|
||||
worker_init_fn=worker_init_fn,
|
||||
multiprocessing_context=_loader_mp_context() if cfg.num_workers > 0 else None,
|
||||
)
|
||||
@@ -0,0 +1,494 @@
|
||||
"""Hand-rolled metrics on CPU numpy (no scikit-learn, no device-dependent math).
|
||||
|
||||
Closed-set reporting with seeded bootstrap CIs, the max-softmax-probability
|
||||
(MSP) open-set diagnostic with val-calibrated tau, and the E1-v2 trap
|
||||
report (false-accept rates at tau, family-collapsed impostor rates, MSP /
|
||||
margin histograms, top-3 retrieval and embedding centroids for the E3
|
||||
hook). Conventions: per-genus recall is NaN for genera absent from the
|
||||
evaluated split; precision and F1 use the zero-division=0 convention so
|
||||
reports stay JSON-clean.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import polars as pl
|
||||
|
||||
DEFAULT_TRAP_ADJACENCY: dict[str, set[str]] = {
|
||||
"Paenibacillus": {"Bacillus"},
|
||||
"Escherichia": {
|
||||
"Cronobacter",
|
||||
"Citrobacter",
|
||||
"Enterobacter",
|
||||
"Klebsiella",
|
||||
"Salmonella",
|
||||
"Shigella",
|
||||
},
|
||||
"Aeromonas": {"Vibrio"},
|
||||
}
|
||||
|
||||
|
||||
def _softmax(logits: np.ndarray) -> np.ndarray:
|
||||
"""Row-wise softmax, numerically stable (max-shifted).
|
||||
|
||||
Args:
|
||||
logits: ``(N, C)`` score matrix.
|
||||
|
||||
Returns:
|
||||
``(N, C)`` probability matrix (rows sum to 1).
|
||||
|
||||
"""
|
||||
z = logits - logits.max(axis=1, keepdims=True)
|
||||
e = np.exp(z)
|
||||
return e / e.sum(axis=1, keepdims=True)
|
||||
|
||||
|
||||
def msp(logits: np.ndarray) -> np.ndarray:
|
||||
"""Max-softmax-probability per row — the open-set confidence score.
|
||||
|
||||
Args:
|
||||
logits: Logit matrix of shape ``(N, C)`` with ``C >= 2``.
|
||||
|
||||
Returns:
|
||||
``(N,)`` float64 MSP values in ``(0, 1]``.
|
||||
|
||||
Raises:
|
||||
ValueError: On non-2-D input or fewer than two classes.
|
||||
|
||||
"""
|
||||
logits = np.asarray(logits, dtype=np.float64)
|
||||
if logits.ndim != 2 or logits.shape[1] < 2:
|
||||
raise ValueError("msp expects logits of shape (N, C) with C >= 2")
|
||||
return _softmax(logits).max(axis=1)
|
||||
|
||||
|
||||
def _confusion(y_true: np.ndarray, y_pred: np.ndarray, n_classes: int) -> np.ndarray:
|
||||
"""Build a ``(C, C)`` count matrix (rows = true, cols = predicted).
|
||||
|
||||
Args:
|
||||
y_true: ``(N,)`` int true class ids.
|
||||
y_pred: ``(N,)`` int predicted class ids.
|
||||
n_classes: Class count C (matrix is always C x C even for absent
|
||||
classes).
|
||||
|
||||
Returns:
|
||||
int64 confusion-count matrix.
|
||||
|
||||
"""
|
||||
codes = y_true.astype(np.int64) * n_classes + y_pred.astype(np.int64)
|
||||
return np.bincount(codes, minlength=n_classes * n_classes).reshape(n_classes, n_classes)
|
||||
|
||||
|
||||
def macro_recall(y_true: np.ndarray, y_pred: np.ndarray, n_classes: int) -> float:
|
||||
"""Macro genus recall — the headline gate metric.
|
||||
|
||||
Mean over genera of ``correct / windows of that genus``; genera absent
|
||||
from ``y_true`` are skipped (NaN-safe). This is the cheap per-epoch
|
||||
variant used during training.
|
||||
|
||||
Args:
|
||||
y_true: ``(N,)`` int true class ids.
|
||||
y_pred: ``(N,)`` int predicted class ids.
|
||||
n_classes: Class count C.
|
||||
|
||||
Returns:
|
||||
Macro recall in ``[0, 1]``; NaN for an empty input pair.
|
||||
|
||||
"""
|
||||
y_true = np.asarray(y_true, dtype=np.int64)
|
||||
y_pred = np.asarray(y_pred, dtype=np.int64)
|
||||
if y_true.size == 0:
|
||||
return float("nan")
|
||||
cm = _confusion(y_true, y_pred, n_classes)
|
||||
rows = cm.sum(axis=1)
|
||||
with np.errstate(invalid="ignore", divide="ignore"):
|
||||
recall = np.where(rows > 0, np.diag(cm) / rows, np.nan)
|
||||
return float(np.nanmean(recall))
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ClosedSetReport:
|
||||
"""Closed-set evaluation summary.
|
||||
|
||||
Attributes:
|
||||
recall_macro: Mean per-genus recall (headline gate metric).
|
||||
recall_micro: Top-1 accuracy (== micro recall).
|
||||
precision_macro: Mean per-class precision (zero-division = 0).
|
||||
f1_macro: Mean per-class F1 (zero-division = 0).
|
||||
per_genus: DataFrame of genus, n_test, recall, precision, f1 and a
|
||||
``trap`` boolean (False for closed-set rows; v2 schema).
|
||||
confusion: ``(C, C)`` count matrix, rows = true, cols = pred,
|
||||
ordered like ``genera``.
|
||||
ci: ``"recall_macro"`` / ``"recall_micro"`` -> 95% bootstrap CI
|
||||
(empty when ``n_boot == 0``).
|
||||
|
||||
"""
|
||||
|
||||
recall_macro: float
|
||||
recall_micro: float
|
||||
precision_macro: float
|
||||
f1_macro: float
|
||||
per_genus: pl.DataFrame
|
||||
confusion: np.ndarray
|
||||
ci: dict[str, tuple[float, float]]
|
||||
|
||||
|
||||
def closed_set_report(
|
||||
y_true: np.ndarray,
|
||||
y_pred: np.ndarray,
|
||||
genera: list[str],
|
||||
n_boot: int = 10_000,
|
||||
seed: int = 0,
|
||||
) -> ClosedSetReport:
|
||||
"""Full closed-set report over predicted labels, with bootstrap CIs.
|
||||
|
||||
The bootstrap resamples evaluation windows with replacement
|
||||
``n_boot`` times (seeded) and recomputes macro/micro recall; the 2.5 /
|
||||
97.5 percentiles give the 95% CIs. E1 fixes ``n_boot=10_000`` and
|
||||
``seed=0`` for cross-run comparability.
|
||||
|
||||
Args:
|
||||
y_true: ``(N,)`` int true class ids in ``[0, C)``.
|
||||
y_pred: ``(N,)`` int predicted class ids in ``[0, C)``.
|
||||
genera: Class names; position = class id (the split's genera
|
||||
list).
|
||||
n_boot: Bootstrap resample count (0 disables CIs).
|
||||
seed: Seed for the bootstrap generator.
|
||||
|
||||
Returns:
|
||||
The :class:`ClosedSetReport`.
|
||||
|
||||
Raises:
|
||||
ValueError: On length mismatch, fewer than 2 classes, an empty
|
||||
evaluation set, or labels outside ``[0, C)``.
|
||||
|
||||
"""
|
||||
y_true = np.asarray(y_true, dtype=np.int64)
|
||||
y_pred = np.asarray(y_pred, dtype=np.int64)
|
||||
if len(y_true) != len(y_pred):
|
||||
raise ValueError("y_true and y_pred must have equal length")
|
||||
n_classes = len(genera)
|
||||
if n_classes < 2:
|
||||
raise ValueError("closed-set report needs at least 2 classes")
|
||||
if y_true.size == 0:
|
||||
raise ValueError("cannot report on an empty evaluation set")
|
||||
lo_ok = y_true.min() >= 0 and y_pred.min() >= 0
|
||||
hi_ok = y_true.max() < n_classes and y_pred.max() < n_classes
|
||||
if not (lo_ok and hi_ok):
|
||||
raise ValueError("labels fall outside [0, n_classes)")
|
||||
confusion = _confusion(y_true, y_pred, n_classes)
|
||||
row_sums = confusion.sum(axis=1)
|
||||
col_sums = confusion.sum(axis=0)
|
||||
diag = np.diag(confusion)
|
||||
with np.errstate(invalid="ignore", divide="ignore"):
|
||||
recall = np.where(row_sums > 0, diag / row_sums, np.nan)
|
||||
precision = np.where(col_sums > 0, diag / col_sums, 0.0)
|
||||
f1 = np.zeros(n_classes, dtype=np.float64)
|
||||
valid = (row_sums > 0) & (np.nan_to_num(recall) + precision > 0)
|
||||
f1[valid] = (
|
||||
2 * recall[valid] * precision[valid] / (recall[valid] + precision[valid])
|
||||
)
|
||||
recall_macro = float(np.nanmean(recall)) if not np.isnan(recall).all() else 0.0
|
||||
recall_micro = float(diag.sum() / confusion.sum())
|
||||
precision_macro = float(precision.mean())
|
||||
f1_macro = float(f1.mean())
|
||||
per_genus = pl.DataFrame(
|
||||
{
|
||||
"genus": genera,
|
||||
"n_test": row_sums.astype(np.int64),
|
||||
"recall": recall,
|
||||
"precision": precision,
|
||||
"f1": f1,
|
||||
"trap": np.zeros(n_classes, dtype=bool),
|
||||
}
|
||||
)
|
||||
ci: dict[str, tuple[float, float]] = {}
|
||||
if n_boot > 0:
|
||||
rng = np.random.default_rng(seed)
|
||||
n = len(y_true)
|
||||
pair_codes = y_true * n_classes + y_pred
|
||||
macro_samples = np.empty(n_boot, dtype=np.float64)
|
||||
micro_samples = np.empty(n_boot, dtype=np.float64)
|
||||
for b in range(n_boot):
|
||||
idx = rng.integers(0, n, size=n)
|
||||
cm = np.bincount(pair_codes[idx], minlength=n_classes * n_classes).reshape(
|
||||
n_classes, n_classes
|
||||
)
|
||||
rows = cm.sum(axis=1)
|
||||
with np.errstate(invalid="ignore", divide="ignore"):
|
||||
rec = np.where(rows > 0, np.diag(cm) / rows, np.nan)
|
||||
macro_samples[b] = np.nanmean(rec)
|
||||
micro_samples[b] = np.diag(cm).sum() / cm.sum()
|
||||
ci["recall_macro"] = (
|
||||
float(np.quantile(macro_samples, 0.025)),
|
||||
float(np.quantile(macro_samples, 0.975)),
|
||||
)
|
||||
ci["recall_micro"] = (
|
||||
float(np.quantile(micro_samples, 0.025)),
|
||||
float(np.quantile(micro_samples, 0.975)),
|
||||
)
|
||||
return ClosedSetReport(
|
||||
recall_macro=recall_macro,
|
||||
recall_micro=recall_micro,
|
||||
precision_macro=precision_macro,
|
||||
f1_macro=f1_macro,
|
||||
per_genus=per_genus,
|
||||
confusion=confusion,
|
||||
ci=ci,
|
||||
)
|
||||
|
||||
|
||||
def choose_tau_msp(
|
||||
val_correct_msp: np.ndarray, target_known_recall: float = 0.95
|
||||
) -> float:
|
||||
"""Calibrate the open-set MSP threshold tau on correct validation reads.
|
||||
|
||||
``tau`` is the ``(1 - target)`` quantile of the MSP values of
|
||||
*correctly classified* validation windows, so at least ``target`` of
|
||||
correct known reads survive ``msp >= tau``.
|
||||
|
||||
Args:
|
||||
val_correct_msp: MSP values of correct validation windows.
|
||||
target_known_recall: Fraction of correct known reads to keep
|
||||
(default 0.95, GATE ZERO §5.9).
|
||||
|
||||
Returns:
|
||||
The threshold ``tau`` in ``[0, 1)``.
|
||||
|
||||
Raises:
|
||||
ValueError: On empty input or a target outside ``(0, 1)``.
|
||||
|
||||
"""
|
||||
arr = np.asarray(val_correct_msp, dtype=np.float64)
|
||||
if arr.size == 0:
|
||||
raise ValueError("choose_tau_msp needs at least one correctly classified window")
|
||||
if not 0.0 < target_known_recall < 1.0:
|
||||
raise ValueError("target_known_recall must lie in (0, 1)")
|
||||
return float(np.quantile(arr, 1.0 - target_known_recall))
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class OpenSetProbeReport:
|
||||
"""E1's simple open-set diagnostic (full calibration is E3's job).
|
||||
|
||||
Attributes:
|
||||
tau: Val-calibrated MSP threshold.
|
||||
false_unknown_rate: Fraction of ALL test windows with
|
||||
``msp < tau`` (routed to "unknown").
|
||||
known_recall_at_tau: Macro recall restricted to windows with
|
||||
``msp >= tau``.
|
||||
kept_fraction: Fraction of test windows with ``msp >= tau``.
|
||||
|
||||
"""
|
||||
|
||||
tau: float
|
||||
false_unknown_rate: float
|
||||
known_recall_at_tau: float
|
||||
kept_fraction: float
|
||||
|
||||
|
||||
def open_set_probe_report(
|
||||
val_logits: np.ndarray,
|
||||
val_y: np.ndarray,
|
||||
test_logits: np.ndarray,
|
||||
test_y: np.ndarray,
|
||||
target_known_recall: float = 0.95,
|
||||
) -> OpenSetProbeReport:
|
||||
"""Compute the open-set probe report from val and test logits.
|
||||
|
||||
Tau comes from the correct-validation MSP quantile
|
||||
(:func:`choose_tau_msp`); the rates are measured on the test split.
|
||||
|
||||
Args:
|
||||
val_logits: ``(N_val, C)`` validation logits.
|
||||
val_y: ``(N_val,)`` validation class ids.
|
||||
test_logits: ``(N_test, C)`` test logits.
|
||||
test_y: ``(N_test,)`` test class ids.
|
||||
target_known_recall: Known-recall target for tau calibration.
|
||||
|
||||
Returns:
|
||||
The :class:`OpenSetProbeReport`.
|
||||
|
||||
Raises:
|
||||
ValueError: On empty validation/test inputs or a validation pass
|
||||
with zero correct windows (tau undefined).
|
||||
|
||||
"""
|
||||
val_logits = np.asarray(val_logits, dtype=np.float64)
|
||||
val_y = np.asarray(val_y, dtype=np.int64)
|
||||
if val_logits.size == 0:
|
||||
raise ValueError("open-set probe needs a non-empty validation set")
|
||||
val_pred = val_logits.argmax(axis=1)
|
||||
correct = val_pred == val_y
|
||||
if not correct.any():
|
||||
raise ValueError("open-set probe needs at least one correct validation window")
|
||||
tau = choose_tau_msp(msp(val_logits)[correct], target_known_recall)
|
||||
test_logits = np.asarray(test_logits, dtype=np.float64)
|
||||
test_y = np.asarray(test_y, dtype=np.int64)
|
||||
if test_logits.size == 0:
|
||||
raise ValueError("open-set probe needs a non-empty test set")
|
||||
test_msp = msp(test_logits)
|
||||
test_pred = test_logits.argmax(axis=1)
|
||||
kept = test_msp >= tau
|
||||
kept_fraction = float(kept.mean())
|
||||
false_unknown_rate = float((~kept).mean())
|
||||
n_classes = test_logits.shape[1]
|
||||
known_recall_at_tau = (
|
||||
macro_recall(test_y[kept], test_pred[kept], n_classes) if kept.any() else float("nan")
|
||||
)
|
||||
return OpenSetProbeReport(
|
||||
tau=tau,
|
||||
false_unknown_rate=false_unknown_rate,
|
||||
known_recall_at_tau=known_recall_at_tau,
|
||||
kept_fraction=kept_fraction,
|
||||
)
|
||||
|
||||
|
||||
def _quantiles(values: np.ndarray) -> dict[str, float]:
|
||||
"""Summarise a value vector with the 5/25/50/75/95 percentiles.
|
||||
|
||||
Args:
|
||||
values: 1-D array (assumed within a known range).
|
||||
|
||||
Returns:
|
||||
Dict with keys ``p05`` .. ``p95``.
|
||||
|
||||
"""
|
||||
q = np.quantile(values, [0.05, 0.25, 0.5, 0.75, 0.95])
|
||||
return {
|
||||
"p05": float(q[0]),
|
||||
"p25": float(q[1]),
|
||||
"p50": float(q[2]),
|
||||
"p75": float(q[3]),
|
||||
"p95": float(q[4]),
|
||||
}
|
||||
|
||||
|
||||
def _histogram(values: np.ndarray, bins: int = 20) -> dict[str, Any]:
|
||||
"""Histogram of a value vector over the unit interval.
|
||||
|
||||
Args:
|
||||
values: 1-D array in ``[0, 1]`` (MSP values, margins).
|
||||
bins: Bin count.
|
||||
|
||||
Returns:
|
||||
``{"edges": [...], "counts": [...]}`` (JSON-friendly).
|
||||
|
||||
"""
|
||||
counts, edges = np.histogram(values, bins=bins, range=(0.0, 1.0))
|
||||
return {"edges": [float(e) for e in edges], "counts": [int(c) for c in counts]}
|
||||
|
||||
|
||||
def trap_report(
|
||||
trap_manifest: pl.DataFrame,
|
||||
trap_logits: np.ndarray,
|
||||
tau: float,
|
||||
genera: list[str],
|
||||
adjacency: dict[str, set[str]] | None = None,
|
||||
embeddings: np.ndarray | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Build the TP2 trap report for eval-only impostor windows.
|
||||
|
||||
Trap genera are never trained classes; every window is pushed through
|
||||
the frozen probe's 11-way head and scored post-hoc. Per trap genus
|
||||
this reports: window/read counts and provenance runs, the false-accept
|
||||
rate at the val-calibrated ``tau`` (``fpr_at_tau``), the
|
||||
family-collapsed impostor rate (``fpr_family`` — prediction in the
|
||||
biologically adjacent classes), MSP and top1-top2 margin quantiles and
|
||||
unit histograms, per-rank top-3 predicted-class counts, and the mean
|
||||
embedding centroid (the E3 hook). Nothing is calibrated *on* traps —
|
||||
tau always comes from the main run's validation split.
|
||||
|
||||
Args:
|
||||
trap_manifest: The trap store's manifest (row order = logits
|
||||
order; genus column names the trap genus per window).
|
||||
trap_logits: ``(N, C)`` logits from the frozen probe.
|
||||
tau: Open-set threshold from :func:`open_set_probe_report`.
|
||||
genera: The probe's class list (position = class id).
|
||||
adjacency: Map trap genus -> adjacent class-genera set; defaults
|
||||
to the owner-list mapping from the v2 design (Paenibacillus ->
|
||||
Bacillus, Escherichia -> the six Enterobacteriaceae classes,
|
||||
Aeromonas -> Vibrio). Unknown trap genera simply get an empty
|
||||
family set.
|
||||
embeddings: Optional ``(N, d_embed)`` pooled embeddings (from
|
||||
:func:`~custom_models.train.embed_probe`) enabling centroid
|
||||
output.
|
||||
|
||||
Returns:
|
||||
JSON-serialisable dict: overall ``trap_fpr_at_tau`` /
|
||||
``trap_fpr_family`` plus a ``per_trap`` breakdown.
|
||||
|
||||
Raises:
|
||||
ValueError: If manifest rows and logits disagree.
|
||||
|
||||
"""
|
||||
adjacency = adjacency if adjacency is not None else DEFAULT_TRAP_ADJACENCY
|
||||
trap_logits = np.asarray(trap_logits, dtype=np.float64)
|
||||
n = trap_logits.shape[0]
|
||||
if len(trap_manifest) != n:
|
||||
raise ValueError("trap manifest rows and logits disagree")
|
||||
probs = _softmax(trap_logits)
|
||||
trap_msp = probs.max(axis=1)
|
||||
trap_pred = trap_logits.argmax(axis=1)
|
||||
order = np.argsort(-probs, axis=1)
|
||||
top1 = order[:, 0]
|
||||
top2 = order[:, 1]
|
||||
margins = probs[np.arange(n), top1] - probs[np.arange(n), top2]
|
||||
genus_arr = trap_manifest["genus"].to_numpy()
|
||||
run_arr = trap_manifest["source_run"].to_numpy()
|
||||
read_arr = trap_manifest["read_id"].to_numpy()
|
||||
accepted = trap_msp >= tau
|
||||
family_ids_by_genus: dict[str, list[int]] = {}
|
||||
for trap_genus in set(genus_arr.tolist()):
|
||||
family = adjacency.get(trap_genus, set())
|
||||
family_ids_by_genus[trap_genus] = [
|
||||
genera.index(g) for g in family if g in genera
|
||||
]
|
||||
in_family = np.array(
|
||||
[
|
||||
trap_pred[i] in family_ids_by_genus.get(genus_arr[i], [])
|
||||
for i in range(n)
|
||||
],
|
||||
dtype=bool,
|
||||
)
|
||||
per_trap: dict[str, Any] = {}
|
||||
for trap_genus in sorted(set(genus_arr.tolist())):
|
||||
idx = np.flatnonzero(genus_arr == trap_genus)
|
||||
family_ids = family_ids_by_genus[trap_genus]
|
||||
pred_g = trap_pred[idx]
|
||||
fpr_family = float(np.isin(pred_g, sorted(family_ids)).mean()) if len(idx) else 0.0
|
||||
ranks = {"rank1": {}, "rank2": {}, "rank3": {}}
|
||||
for rank, col in ((0, "rank1"), (1, "rank2"), (2, "rank3")):
|
||||
counts: dict[str, int] = {}
|
||||
for c in order[idx, rank]:
|
||||
name = genera[int(c)]
|
||||
counts[name] = counts.get(name, 0) + 1
|
||||
ranks[col] = dict(sorted(counts.items(), key=lambda kv: (-kv[1], kv[0])))
|
||||
entry: dict[str, Any] = {
|
||||
"n_eval_windows": len(idx),
|
||||
"source_runs": sorted(set(run_arr[idx].tolist())),
|
||||
"n_reads": len(set(read_arr[idx].tolist())),
|
||||
"fpr_at_tau": float(accepted[idx].mean()),
|
||||
"fpr_family": fpr_family,
|
||||
"adjacent_classes": sorted(family),
|
||||
"msp_quantiles": _quantiles(trap_msp[idx]),
|
||||
"margin_quantiles": _quantiles(margins[idx]),
|
||||
"msp_hist": _histogram(trap_msp[idx]),
|
||||
"margin_hist": _histogram(margins[idx]),
|
||||
"top3_by_rank": ranks,
|
||||
}
|
||||
if embeddings is not None and len(idx):
|
||||
entry["centroid"] = [float(v) for v in embeddings[idx].mean(axis=0)]
|
||||
else:
|
||||
entry["centroid"] = None
|
||||
per_trap[trap_genus] = entry
|
||||
return {
|
||||
"tau": float(tau),
|
||||
"n_eval_windows": int(n),
|
||||
"trap_fpr_at_tau": float(accepted.mean()) if n else 0.0,
|
||||
"trap_fpr_family": float(in_family.mean()) if n else 0.0,
|
||||
"per_trap": per_trap,
|
||||
}
|
||||
@@ -0,0 +1,600 @@
|
||||
"""Torch modules for the E1 supervised genus probe.
|
||||
|
||||
Implements the normative model spec (``E1_torch_model_spec.md`` §3): a
|
||||
``SignalPatchEmbed`` conv patchifier (no positional information), the
|
||||
``ConvEncoder`` (depthwise-separable blocks, dilation capped at 16,
|
||||
zero-initialised attention-weighted pooling) and the
|
||||
``LinearAttentionEncoder`` (pre-norm "fast transformer" blocks with the
|
||||
elu+1 kernel, plain mean pooling), both terminating in a pooled
|
||||
``d_embed`` vector, plus the linear-headed ``GenusProbe`` and the
|
||||
deterministic budget resolver ``build_probe``. No batch-norm anywhere
|
||||
(MPS/AMP-hostile, device-identical semantics required); no softmax or
|
||||
log-softmax inside the model.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import replace
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
|
||||
from .config import ModelConfig
|
||||
|
||||
D_MODEL_LATTICE: tuple[int, ...] = (
|
||||
32,
|
||||
48,
|
||||
64,
|
||||
96,
|
||||
128,
|
||||
192,
|
||||
256,
|
||||
384,
|
||||
512,
|
||||
768,
|
||||
1024,
|
||||
1536,
|
||||
2048,
|
||||
3072,
|
||||
4096,
|
||||
)
|
||||
N_LAYERS_LATTICE: tuple[int, ...] = (2, 3, 4, 6, 8, 12, 16)
|
||||
DILATION_CAP = 16
|
||||
LINEAR_ATTENTION_EPS = 1e-6
|
||||
|
||||
|
||||
def count_params(module: nn.Module) -> int:
|
||||
"""Count trainable parameters of a module (buffers excluded).
|
||||
|
||||
This is the single source of truth for budget checks and reporting —
|
||||
closed-form parameter arithmetic is design intuition only (spec §4).
|
||||
|
||||
Args:
|
||||
module: Any ``torch.nn.Module`` (typically a ``GenusProbe``).
|
||||
|
||||
Returns:
|
||||
Total ``numel`` over parameters with ``requires_grad=True``.
|
||||
|
||||
"""
|
||||
return sum(p.numel() for p in module.parameters() if p.requires_grad)
|
||||
|
||||
|
||||
class SignalPatchEmbed(nn.Module):
|
||||
"""Patchify raw signal into tokens with a strided 1-D convolution.
|
||||
|
||||
``Conv1d(1 -> d_model, kernel=patch_len, stride=stride)`` followed by a
|
||||
transpose to token layout ``(B, T, d_model)``. Bias is kept (it is part
|
||||
of the parameter count). No positional information is added: absolute
|
||||
position within a read is biologically irrelevant for a genus signal,
|
||||
and ~3k-token positional tables would eat a large fraction of the 1M
|
||||
budget point (spec §3.1).
|
||||
"""
|
||||
|
||||
def __init__(self, d_model: int, patch_len: int, stride: int) -> None:
|
||||
"""Create the patch embedding.
|
||||
|
||||
Args:
|
||||
d_model: Token width produced by the convolution.
|
||||
patch_len: Kernel size in signal samples (E1 uses
|
||||
``min(4, stride)``).
|
||||
stride: Token pitch in signal samples (4 or 8).
|
||||
|
||||
Raises:
|
||||
ValueError: If any argument is < 1.
|
||||
|
||||
"""
|
||||
super().__init__()
|
||||
if d_model < 1:
|
||||
raise ValueError("d_model must be >= 1")
|
||||
if patch_len < 1:
|
||||
raise ValueError("patch_len must be >= 1")
|
||||
if stride < 1:
|
||||
raise ValueError("stride must be >= 1")
|
||||
self.patch_len = patch_len
|
||||
self.stride = stride
|
||||
self.proj = nn.Conv1d(1, d_model, kernel_size=patch_len, stride=stride)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Map a raw-signal batch to tokens.
|
||||
|
||||
Args:
|
||||
x: Signal batch of shape ``(B, L)`` float32; ``L`` must be
|
||||
divisible by ``stride`` (guaranteed by the store, validated
|
||||
defensively per the API contract).
|
||||
|
||||
Returns:
|
||||
Tokens of shape ``(B, T, d_model)`` with
|
||||
``T = (L - patch_len) // stride + 1`` (3,000 for L=12,000 and
|
||||
stride 4).
|
||||
|
||||
Raises:
|
||||
ValueError: If ``x`` is not 2-D or ``L % stride != 0``.
|
||||
|
||||
"""
|
||||
if x.dim() != 2:
|
||||
raise ValueError(f"expected signal batch of shape (B, L), got {tuple(x.shape)}")
|
||||
if x.shape[1] % self.stride != 0:
|
||||
raise ValueError(
|
||||
f"window length {x.shape[1]} is not divisible by stride {self.stride}"
|
||||
)
|
||||
tokens = self.proj(x.unsqueeze(1))
|
||||
return tokens.transpose(1, 2)
|
||||
|
||||
|
||||
class ConvBlock(nn.Module):
|
||||
"""One depthwise-separable conv block: depthwise -> pointwise -> LN -> GELU.
|
||||
|
||||
Padding ``dilation * (k - 1) // 2`` preserves the token count so
|
||||
``T`` is invariant through the stack. Post-norm ordering follows the
|
||||
GATE ZERO block sketch (spec §3.2).
|
||||
"""
|
||||
|
||||
def __init__(self, d_model: int, kernel_size: int, dilation: int, dropout: float) -> None:
|
||||
"""Create the block.
|
||||
|
||||
Args:
|
||||
d_model: Token width (channels in/Out).
|
||||
kernel_size: Depthwise kernel size (7 anchors layer 0's
|
||||
receptive field, 5 afterwards).
|
||||
dilation: Depthwise dilation (growth capped at
|
||||
``DILATION_CAP`` by the caller).
|
||||
dropout: Dropout probability applied after activation; 0.0
|
||||
installs an identity (E1 matrices run at 0.0).
|
||||
|
||||
"""
|
||||
super().__init__()
|
||||
padding = dilation * (kernel_size - 1) // 2
|
||||
self.depthwise = nn.Conv1d(
|
||||
d_model,
|
||||
d_model,
|
||||
kernel_size=kernel_size,
|
||||
padding=padding,
|
||||
dilation=dilation,
|
||||
groups=d_model,
|
||||
)
|
||||
self.pointwise = nn.Conv1d(d_model, d_model, kernel_size=1)
|
||||
self.norm = nn.LayerNorm(d_model)
|
||||
self.act = nn.GELU()
|
||||
self.drop = nn.Dropout(dropout) if dropout > 0 else nn.Identity()
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Run the block over a token tensor.
|
||||
|
||||
Args:
|
||||
x: Tokens of shape ``(B, T, d_model)``.
|
||||
|
||||
Returns:
|
||||
Tokens of the same shape (length-preserving padding).
|
||||
|
||||
"""
|
||||
h = x.transpose(1, 2)
|
||||
h = self.pointwise(self.depthwise(h))
|
||||
h = h.transpose(1, 2)
|
||||
return self.drop(self.act(self.norm(h)))
|
||||
|
||||
|
||||
class ConvEncoder(nn.Module):
|
||||
"""Depthwise-separable 1-D CNN encoder with attention-weighted pooling.
|
||||
|
||||
Stack of ``n_layers`` :class:`ConvBlock` instances — kernel 7 in layer
|
||||
0 (receptive-field anchor) then 5, dilation ``min(2**l, 16)`` (GATE
|
||||
ZERO §12.3 cap) — followed by ``weighted_mean_pool``: a zero-init
|
||||
``Linear(d_model -> 1)`` score head (uniform pooling at init, plain
|
||||
mean behaviour) softmaxed over tokens, applied to a
|
||||
``Linear(d_model -> d_embed)`` projection.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
d_model: int,
|
||||
n_layers: int,
|
||||
patch_len: int,
|
||||
stride: int,
|
||||
d_embed: int,
|
||||
dropout: float = 0.0,
|
||||
) -> None:
|
||||
"""Create the encoder.
|
||||
|
||||
Args:
|
||||
d_model: Token width inside the stack.
|
||||
n_layers: Number of conv blocks.
|
||||
patch_len: Patch length passed to the patch embed.
|
||||
stride: Token pitch passed to the patch embed.
|
||||
d_embed: Output pooled-embedding width.
|
||||
dropout: Block dropout (E1 uses 0.0; do not tune inside E1 —
|
||||
that changes the ladder semantics).
|
||||
|
||||
Raises:
|
||||
ValueError: If ``n_layers`` or ``d_embed`` < 1 (patch-embed
|
||||
args are validated in :class:`SignalPatchEmbed`).
|
||||
|
||||
"""
|
||||
super().__init__()
|
||||
if n_layers < 1:
|
||||
raise ValueError("n_layers must be >= 1")
|
||||
if d_embed < 1:
|
||||
raise ValueError("d_embed must be >= 1")
|
||||
self.embed = SignalPatchEmbed(d_model, patch_len, stride)
|
||||
self.blocks = nn.ModuleList(
|
||||
[
|
||||
ConvBlock(
|
||||
d_model,
|
||||
kernel_size=7 if layer == 0 else 5,
|
||||
dilation=min(2**layer, DILATION_CAP),
|
||||
dropout=dropout,
|
||||
)
|
||||
for layer in range(n_layers)
|
||||
]
|
||||
)
|
||||
self.score = nn.Linear(d_model, 1)
|
||||
nn.init.zeros_(self.score.weight)
|
||||
nn.init.zeros_(self.score.bias)
|
||||
self.proj = nn.Linear(d_model, d_embed)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Encode raw signal to a pooled embedding.
|
||||
|
||||
Args:
|
||||
x: Signal batch of shape ``(B, L)`` float32 with ``L``
|
||||
divisible by ``stride``.
|
||||
|
||||
Returns:
|
||||
Pooled embedding of shape ``(B, d_embed)``.
|
||||
|
||||
"""
|
||||
h = self.embed(x)
|
||||
for block in self.blocks:
|
||||
h = block(h)
|
||||
weights = torch.softmax(self.score(h), dim=1)
|
||||
return (weights * self.proj(h)).sum(dim=1)
|
||||
|
||||
|
||||
class LinearAttention(nn.Module):
|
||||
"""Non-causal linear attention with the elu+1 feature map.
|
||||
|
||||
Computes ``out[t] = phi(q_t) @ (sum_s phi(k_s) v_s^T) / (phi(q_t) @
|
||||
sum_s phi(k_s))`` per head via the associative form — ``O(T*d^2)``
|
||||
compute and ``O(d^2)`` memory per layer, no ``T x T`` matrix — which
|
||||
is what keeps ~3k-token sequences trainable inside the 1-3 GPU-h/run
|
||||
budget. ``phi(x) = elu(x) + 1`` is strictly positive, so the
|
||||
denominator is positive; a small floor guards degenerate cases.
|
||||
"""
|
||||
|
||||
def __init__(self, d_model: int, n_heads: int) -> None:
|
||||
"""Create the attention mixer.
|
||||
|
||||
Args:
|
||||
d_model: Token width (queries/keys/values are full-width
|
||||
linears reshaped to heads, so param count is head-count
|
||||
invariant).
|
||||
n_heads: Number of heads (default 4; recorded in the realised
|
||||
``ModelConfig``).
|
||||
|
||||
Raises:
|
||||
ValueError: If ``d_model`` is not divisible by ``n_heads`` or
|
||||
``n_heads`` < 1.
|
||||
|
||||
"""
|
||||
super().__init__()
|
||||
if n_heads < 1:
|
||||
raise ValueError("n_heads must be >= 1")
|
||||
if d_model % n_heads != 0:
|
||||
raise ValueError(f"d_model {d_model} not divisible by n_heads {n_heads}")
|
||||
self.n_heads = n_heads
|
||||
self.head_dim = d_model // n_heads
|
||||
self.q = nn.Linear(d_model, d_model)
|
||||
self.k = nn.Linear(d_model, d_model)
|
||||
self.v = nn.Linear(d_model, d_model)
|
||||
self.out = nn.Linear(d_model, d_model)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Mix tokens with linear attention.
|
||||
|
||||
Args:
|
||||
x: Tokens of shape ``(B, T, d_model)``.
|
||||
|
||||
Returns:
|
||||
Mixed tokens of the same shape.
|
||||
|
||||
"""
|
||||
b, t, d = x.shape
|
||||
q = self.q(x).view(b, t, self.n_heads, self.head_dim).transpose(1, 2)
|
||||
k = self.k(x).view(b, t, self.n_heads, self.head_dim).transpose(1, 2)
|
||||
v = self.v(x).view(b, t, self.n_heads, self.head_dim).transpose(1, 2)
|
||||
phi_q = F.elu(q) + 1.0
|
||||
phi_k = F.elu(k) + 1.0
|
||||
kv = torch.einsum("bhti,bhtj->bhij", phi_k, v)
|
||||
k_sum = phi_k.sum(dim=2)
|
||||
numerator = torch.einsum("bhti,bhij->bhtj", phi_q, kv)
|
||||
denominator = torch.einsum("bhti,bhi->bht", phi_q, k_sum)
|
||||
out = numerator / denominator.clamp_min(LINEAR_ATTENTION_EPS).unsqueeze(-1)
|
||||
out = out.transpose(1, 2).reshape(b, t, d)
|
||||
return self.out(out)
|
||||
|
||||
|
||||
class LinearAttentionBlock(nn.Module):
|
||||
"""Pre-norm transformer block: attention residual, then MLP residual."""
|
||||
|
||||
def __init__(self, d_model: int, n_heads: int, dropout: float) -> None:
|
||||
"""Create the block.
|
||||
|
||||
Args:
|
||||
d_model: Token width.
|
||||
n_heads: Attention head count.
|
||||
dropout: Residual dropout (0.0 in E1 matrices).
|
||||
|
||||
"""
|
||||
super().__init__()
|
||||
self.norm1 = nn.LayerNorm(d_model)
|
||||
self.attn = LinearAttention(d_model, n_heads)
|
||||
self.norm2 = nn.LayerNorm(d_model)
|
||||
self.mlp = nn.Sequential(
|
||||
nn.Linear(d_model, d_model),
|
||||
nn.GELU(),
|
||||
nn.Linear(d_model, d_model),
|
||||
)
|
||||
self.drop1 = nn.Dropout(dropout) if dropout > 0 else nn.Identity()
|
||||
self.drop2 = nn.Dropout(dropout) if dropout > 0 else nn.Identity()
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Run both mixer stages with pre-norm residuals.
|
||||
|
||||
Args:
|
||||
x: Tokens of shape ``(B, T, d_model)``.
|
||||
|
||||
Returns:
|
||||
Tokens of the same shape.
|
||||
|
||||
"""
|
||||
x = x + self.drop1(self.attn(self.norm1(x)))
|
||||
x = x + self.drop2(self.mlp(self.norm2(x)))
|
||||
return x
|
||||
|
||||
|
||||
class LinearAttentionEncoder(nn.Module):
|
||||
"""Patch embed + pre-norm linear-attention blocks + mean pooling.
|
||||
|
||||
Terminates in the same interchangeable pooled representation as
|
||||
:class:`ConvEncoder`: plain mean over tokens followed by a
|
||||
``Linear(d_model -> d_embed)`` projection (no score head — the pooling
|
||||
is deliberately not attention-weighted here, spec §3.3).
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
d_model: int,
|
||||
n_layers: int,
|
||||
patch_len: int,
|
||||
stride: int,
|
||||
d_embed: int,
|
||||
dropout: float = 0.0,
|
||||
n_heads: int = 4,
|
||||
) -> None:
|
||||
"""Create the encoder.
|
||||
|
||||
Args:
|
||||
d_model: Token width inside the stack.
|
||||
n_layers: Number of attention blocks.
|
||||
patch_len: Patch length for the (separately weighted) patch
|
||||
embed.
|
||||
stride: Token pitch for the patch embed.
|
||||
d_embed: Output pooled-embedding width.
|
||||
dropout: Residual dropout (0.0 in E1 matrices).
|
||||
n_heads: Attention head count.
|
||||
|
||||
Raises:
|
||||
ValueError: If ``n_layers`` or ``d_embed`` < 1; head
|
||||
divisibility is enforced in :class:`LinearAttention`.
|
||||
|
||||
"""
|
||||
super().__init__()
|
||||
if n_layers < 1:
|
||||
raise ValueError("n_layers must be >= 1")
|
||||
if d_embed < 1:
|
||||
raise ValueError("d_embed must be >= 1")
|
||||
self.embed = SignalPatchEmbed(d_model, patch_len, stride)
|
||||
self.blocks = nn.ModuleList(
|
||||
[LinearAttentionBlock(d_model, n_heads, dropout) for _ in range(n_layers)]
|
||||
)
|
||||
self.proj = nn.Linear(d_model, d_embed)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Encode raw signal to a pooled embedding.
|
||||
|
||||
Args:
|
||||
x: Signal batch of shape ``(B, L)`` float32 with ``L``
|
||||
divisible by ``stride``.
|
||||
|
||||
Returns:
|
||||
Pooled embedding of shape ``(B, d_embed)``.
|
||||
|
||||
"""
|
||||
h = self.embed(x)
|
||||
for block in self.blocks:
|
||||
h = block(h)
|
||||
return self.proj(h.mean(dim=1))
|
||||
|
||||
|
||||
class GenusProbe(nn.Module):
|
||||
"""The supervised probe: an encoder plus a bare linear classifier.
|
||||
|
||||
The head is exactly ``Linear(d_embed, n_classes)`` — no hidden MLP, no
|
||||
temperature, no bias tricks — because E1 measures the *encoder's*
|
||||
information content; a linear head on a pooled representation is the
|
||||
cleanest statement of "the information is (or is not) in the
|
||||
embedding". Raw logits only; cross-entropy supplies the softmax.
|
||||
"""
|
||||
|
||||
def __init__(self, encoder: nn.Module, d_embed: int, n_classes: int) -> None:
|
||||
"""Assemble the probe.
|
||||
|
||||
Args:
|
||||
encoder: An encoder (``ConvEncoder`` or
|
||||
``LinearAttentionEncoder``) mapping ``(B, L)`` to
|
||||
``(B, d_embed)``.
|
||||
d_embed: Width of the encoder's pooled output; must match or
|
||||
the first forward raises on the head's shape mismatch.
|
||||
n_classes: Number of genera (C = 11 for E1 v2).
|
||||
|
||||
Raises:
|
||||
ValueError: If ``n_classes`` < 2 (empty classification head) or
|
||||
``d_embed`` < 1.
|
||||
|
||||
"""
|
||||
super().__init__()
|
||||
if n_classes < 2:
|
||||
raise ValueError("n_classes must be >= 2")
|
||||
if d_embed < 1:
|
||||
raise ValueError("d_embed must be >= 1")
|
||||
self.encoder = encoder
|
||||
self.head = nn.Linear(d_embed, n_classes)
|
||||
|
||||
def embed(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Return the pooled embedding (E3/E4 reuse path).
|
||||
|
||||
Args:
|
||||
x: Signal batch of shape ``(B, L)`` float32.
|
||||
|
||||
Returns:
|
||||
Embeddings of shape ``(B, d_embed)`` — not L2-normalised;
|
||||
cosine machinery belongs to E3, not the probe.
|
||||
|
||||
"""
|
||||
return self.encoder(x)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""Classify raw signal.
|
||||
|
||||
Args:
|
||||
x: Signal batch of shape ``(B, L)`` float32.
|
||||
|
||||
Returns:
|
||||
Logits of shape ``(B, n_classes)``.
|
||||
|
||||
"""
|
||||
return self.head(self.embed(x))
|
||||
|
||||
|
||||
def _make_encoder(
|
||||
arch: str,
|
||||
d_model: int,
|
||||
n_layers: int,
|
||||
patch_len: int,
|
||||
stride: int,
|
||||
d_embed: int,
|
||||
n_heads: int,
|
||||
) -> nn.Module:
|
||||
"""Instantiate an encoder of the requested family.
|
||||
|
||||
Args:
|
||||
arch: ``"cnn"`` or ``"linatt"``.
|
||||
d_model: Token width.
|
||||
n_layers: Block count.
|
||||
patch_len: Patch length in samples.
|
||||
stride: Token pitch in samples.
|
||||
d_embed: Pooled output width.
|
||||
n_heads: Head count (ignored by ``"cnn"``).
|
||||
|
||||
Returns:
|
||||
A fresh encoder module (default torch init; draws from the global
|
||||
RNG so callers must seed first).
|
||||
|
||||
Raises:
|
||||
ValueError: On an unknown arch.
|
||||
|
||||
"""
|
||||
if arch == "cnn":
|
||||
return ConvEncoder(d_model, n_layers, patch_len, stride, d_embed)
|
||||
if arch == "linatt":
|
||||
return LinearAttentionEncoder(
|
||||
d_model, n_layers, patch_len, stride, d_embed, n_heads=n_heads
|
||||
)
|
||||
raise ValueError(f"unknown arch: {arch!r}")
|
||||
|
||||
|
||||
def build_probe(cfg: ModelConfig, n_classes: int) -> tuple[GenusProbe, ModelConfig]:
|
||||
"""Resolve a parameter budget into a concrete probe (deterministic).
|
||||
|
||||
Searches the ``(d_model, n_layers)`` lattice — ``D_MODEL_LATTICE`` x
|
||||
``N_LAYERS_LATTICE``, pruning by monotonicity (parameter counts grow
|
||||
with both axes) so oversized candidates are never constructed — and
|
||||
among candidates with ``count_params <= (1 + budget_tol) * budget``
|
||||
picks the *largest* (ceiling-cap reading of GATE ZERO §12.3; candidates
|
||||
above the cap are disqualified even if closer in absolute distance).
|
||||
Ties break by ``(n_layers asc, d_model asc)``. The returned realised
|
||||
config is the value serialised into ``config.json``; re-running the
|
||||
search on it reproduces the same geometry (checkpoint rehydration
|
||||
path). Construction draws only from torch's global RNG: callers must
|
||||
call ``seed_everything`` first, and identical (cfg, n_classes, seed)
|
||||
yields bit-identical weights.
|
||||
|
||||
Args:
|
||||
cfg: Model request (arch, nominal budget, stride, d_embed,
|
||||
tolerance, head count); realised fields are overwritten.
|
||||
n_classes: Number of genera (C = 11 for E1 v2).
|
||||
|
||||
Returns:
|
||||
``(probe, realised_cfg)`` where ``realised_cfg`` fills ``d_model``,
|
||||
``n_layers`` and ``params_realised``.
|
||||
|
||||
Raises:
|
||||
ValueError: If ``n_classes`` < 2, ``stride``/``d_embed`` invalid,
|
||||
or no lattice candidate lands within tolerance (the ladder
|
||||
matrix then drops that arch/budget cell — fail fast).
|
||||
|
||||
"""
|
||||
if n_classes < 2:
|
||||
raise ValueError(f"n_classes must be >= 2, got {n_classes}")
|
||||
if cfg.stride < 1:
|
||||
raise ValueError("stride must be >= 1")
|
||||
if cfg.d_embed < 1:
|
||||
raise ValueError("d_embed must be >= 1")
|
||||
patch_len = min(4, cfg.stride)
|
||||
ceiling = int((1.0 + cfg.budget_tol) * cfg.param_budget)
|
||||
best: tuple[int, int, int, GenusProbe] | None = None
|
||||
for d_model in D_MODEL_LATTICE:
|
||||
probe = GenusProbe(
|
||||
_make_encoder(
|
||||
cfg.arch,
|
||||
d_model,
|
||||
N_LAYERS_LATTICE[0],
|
||||
patch_len,
|
||||
cfg.stride,
|
||||
cfg.d_embed,
|
||||
cfg.n_heads,
|
||||
),
|
||||
cfg.d_embed,
|
||||
n_classes,
|
||||
)
|
||||
params = count_params(probe)
|
||||
if params > ceiling:
|
||||
break
|
||||
local: tuple[int, int, int, GenusProbe] = (params, N_LAYERS_LATTICE[0], d_model, probe)
|
||||
for n_layers in N_LAYERS_LATTICE[1:]:
|
||||
probe = GenusProbe(
|
||||
_make_encoder(
|
||||
cfg.arch, d_model, n_layers, patch_len, cfg.stride, cfg.d_embed, cfg.n_heads
|
||||
),
|
||||
cfg.d_embed,
|
||||
n_classes,
|
||||
)
|
||||
params = count_params(probe)
|
||||
if params > ceiling:
|
||||
break
|
||||
local = (params, n_layers, d_model, probe)
|
||||
if best is None or (local[0], -local[1], -local[2]) > (best[0], -best[1], -best[2]):
|
||||
best = local
|
||||
if best is None:
|
||||
raise ValueError(
|
||||
f"no ({cfg.arch}) candidate within {cfg.budget_tol:.0%} of budget {cfg.param_budget}"
|
||||
)
|
||||
params, n_layers, d_model, probe = best
|
||||
realised = replace(
|
||||
cfg,
|
||||
d_model=d_model,
|
||||
n_layers=n_layers,
|
||||
params_realised=params,
|
||||
n_heads=cfg.n_heads,
|
||||
)
|
||||
return probe, realised
|
||||
@@ -0,0 +1,611 @@
|
||||
"""Run orchestration: artifact writing, checkpoint rehydration and the gate.
|
||||
|
||||
``train_run`` executes one run end to end (seed, split, build, train,
|
||||
artifacts); ``evaluate_val`` / ``evaluate_test`` rehydrate a run directory
|
||||
deterministically (same config + same store -> same splits, control runs
|
||||
re-permute identically) and write the evaluation artifact set; and
|
||||
``evaluate_gate`` implements the E1-v2 four-line verdict table over run
|
||||
directories plus the shuffled-label control.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import asdict, dataclass, replace
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import polars as pl
|
||||
import torch
|
||||
|
||||
from .config import RunConfig, load_config, run_id_for, save_config
|
||||
from .data import (
|
||||
Splits,
|
||||
make_full_loader,
|
||||
make_loaders,
|
||||
make_splits,
|
||||
permute_read_labels,
|
||||
)
|
||||
from .metrics import (
|
||||
ClosedSetReport,
|
||||
closed_set_report,
|
||||
msp,
|
||||
open_set_probe_report,
|
||||
trap_report,
|
||||
)
|
||||
from .models import build_probe
|
||||
from .store import SignalStore
|
||||
from .train import (
|
||||
TrainResult,
|
||||
embed_probe,
|
||||
predict,
|
||||
resolve_amp,
|
||||
select_device,
|
||||
train_probe,
|
||||
)
|
||||
|
||||
CONFIG_NAME = "config.json"
|
||||
CKPT_NAME = "ckpt.pt"
|
||||
HISTORY_NAME = "history.parquet"
|
||||
REPORT_NAME = "report.json"
|
||||
VAL_REPORT_NAME = "val_report.json"
|
||||
PER_GENUS_NAME = "per_genus.parquet"
|
||||
CONFUSION_NAME = "confusion.npz"
|
||||
OPEN_SET_NAME = "open_set.json"
|
||||
ERRORS_NAME = "errors.parquet"
|
||||
TRAP_REPORT_NAME = "trap_report.json"
|
||||
GATE_NAME = "gate_verdict.json"
|
||||
|
||||
|
||||
def _write_json(path: Path, payload: Any) -> None:
|
||||
"""Write a JSON artifact atomically enough for run directories.
|
||||
|
||||
Args:
|
||||
path: Destination file (parents created as needed).
|
||||
payload: JSON-serialisable object (dicts/dataclasses already
|
||||
converted by callers).
|
||||
|
||||
"""
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
path.write_text(json.dumps(payload, indent=2, sort_keys=True) + "\n")
|
||||
|
||||
|
||||
def train_run(cfg: RunConfig) -> tuple[RunConfig, TrainResult, Splits]:
|
||||
"""Execute one training run and write its core artifacts.
|
||||
|
||||
Order of operations (all stochastic steps seeded from the config):
|
||||
``seed_everything`` -> open store -> permute labels if the stage is
|
||||
``"control"`` -> deterministic splits -> ``build_probe`` budget
|
||||
resolution (run id defaults to ``run_id_for`` when unset) -> write the
|
||||
realised ``config.json`` -> ``train_probe`` (writes
|
||||
``history.parquet``) -> save ``ckpt.pt`` holding the *best* state.
|
||||
Re-runs of the same config overwrite artifacts (idempotent).
|
||||
|
||||
Args:
|
||||
cfg: The run configuration. ``cfg.run_id`` may be empty (then
|
||||
derived from arch/budget/stride/class count/stage).
|
||||
|
||||
Returns:
|
||||
``(realised_cfg, result, splits)`` — the realised config that was
|
||||
serialised, the training outcome and the splits used.
|
||||
|
||||
Raises:
|
||||
Whatever ``make_splits``, ``build_probe`` and ``train_probe``
|
||||
raise on invalid data or configuration (ValueError /
|
||||
RuntimeError / FileNotFoundError).
|
||||
|
||||
"""
|
||||
from .seed import seed_everything
|
||||
|
||||
seed_everything(cfg.train.seed, cfg.train.deterministic)
|
||||
store = SignalStore(cfg.store_dir)
|
||||
manifest = store.manifest()
|
||||
if cfg.stage == "control":
|
||||
manifest = permute_read_labels(manifest, cfg.data.genera, cfg.data.seed)
|
||||
splits = make_splits(manifest, cfg.data)
|
||||
n_classes = len(splits.genera)
|
||||
probe, realised_model = build_probe(cfg.model, n_classes)
|
||||
run_id = cfg.run_id
|
||||
if not run_id:
|
||||
run_id = run_id_for(
|
||||
cfg.model.arch, cfg.model.param_budget, cfg.model.stride, n_classes, cfg.stage
|
||||
)
|
||||
realised = replace(
|
||||
cfg,
|
||||
run_id=run_id,
|
||||
model=realised_model,
|
||||
genera=tuple(splits.genera),
|
||||
)
|
||||
run_dir = realised.out_dir / run_id
|
||||
run_dir.mkdir(parents=True, exist_ok=True)
|
||||
save_config(realised, run_dir / CONFIG_NAME)
|
||||
result = train_probe(probe, store, splits, realised.train, run_dir)
|
||||
torch.save(
|
||||
{
|
||||
"best_state": result.best_state,
|
||||
"best_epoch": result.best_epoch,
|
||||
"run_id": run_id,
|
||||
"params_realised": realised_model.params_realised,
|
||||
},
|
||||
run_dir / CKPT_NAME,
|
||||
)
|
||||
return realised, result, splits
|
||||
|
||||
|
||||
def load_run(run_dir: Path) -> RunConfig:
|
||||
"""Load a run directory's realised ``RunConfig``.
|
||||
|
||||
Args:
|
||||
run_dir: Directory holding ``config.json`` (as written by
|
||||
:func:`train_run`).
|
||||
|
||||
Returns:
|
||||
The realised run configuration.
|
||||
|
||||
Raises:
|
||||
FileNotFoundError: If the config file is missing.
|
||||
ValueError: If the file does not decode into a ``RunConfig``.
|
||||
|
||||
"""
|
||||
cfg = load_config(Path(run_dir) / CONFIG_NAME)
|
||||
if not isinstance(cfg, RunConfig):
|
||||
raise ValueError(f"{Path(run_dir) / CONFIG_NAME} does not hold a RunConfig")
|
||||
return cfg
|
||||
|
||||
|
||||
def load_probe(run_dir: Path, cfg: RunConfig) -> torch.nn.Module:
|
||||
"""Rebuild a run's probe and load its best checkpoint weights.
|
||||
|
||||
The geometry comes from re-running the deterministic
|
||||
``build_probe`` search on the realised config (a pure function of
|
||||
config + class count, so the shapes always match), then
|
||||
``load_state_dict`` fills the weights.
|
||||
|
||||
Args:
|
||||
run_dir: Run directory holding ``ckpt.pt``.
|
||||
cfg: The run's realised config (model section + genera).
|
||||
|
||||
Returns:
|
||||
The probe (on CPU) with best-epoch weights loaded.
|
||||
|
||||
Raises:
|
||||
FileNotFoundError: If the checkpoint is missing.
|
||||
|
||||
"""
|
||||
run_dir = Path(run_dir)
|
||||
ckpt_path = run_dir / CKPT_NAME
|
||||
if not ckpt_path.is_file():
|
||||
raise FileNotFoundError(f"checkpoint not found: {ckpt_path}")
|
||||
n_classes = len(cfg.genera)
|
||||
probe, _ = build_probe(cfg.model, n_classes)
|
||||
ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=True)
|
||||
probe.load_state_dict(ckpt["best_state"])
|
||||
return probe
|
||||
|
||||
|
||||
def _rehydrate(
|
||||
run_dir: Path, store_dir: Path | None = None
|
||||
) -> tuple[RunConfig, SignalStore, pl.DataFrame, Splits, torch.nn.Module, torch.device, str]:
|
||||
"""Rebuild everything an evaluation needs from a run directory.
|
||||
|
||||
Recomputes the splits from the run's config (deterministic: same
|
||||
store + same seeds -> the training-time splits; control stages
|
||||
re-apply their seeded label permutation first), reloads the probe,
|
||||
and resolves device/AMP. The returned manifest is the *working*
|
||||
manifest (permuted for control runs) so error tables agree with the
|
||||
labels the run was trained on.
|
||||
|
||||
Args:
|
||||
run_dir: Run directory (config + checkpoint).
|
||||
store_dir: Optional store override; defaults to the recorded
|
||||
``store_dir``.
|
||||
|
||||
Returns:
|
||||
``(cfg, store, manifest, splits, probe, device, amp)`` with the
|
||||
probe already moved to ``device``.
|
||||
|
||||
Raises:
|
||||
ValueError: If the recomputed genera disagree with the recorded
|
||||
genera (the store changed since training — evaluation would
|
||||
not be apples-to-apples).
|
||||
|
||||
"""
|
||||
cfg = load_run(run_dir)
|
||||
store = SignalStore(store_dir if store_dir is not None else cfg.store_dir)
|
||||
manifest = store.manifest()
|
||||
if cfg.stage == "control":
|
||||
manifest = permute_read_labels(manifest, cfg.data.genera, cfg.data.seed)
|
||||
splits = make_splits(manifest, cfg.data)
|
||||
if list(cfg.genera) != splits.genera:
|
||||
raise ValueError(
|
||||
"store genera disagree with the recorded run genera; "
|
||||
"the store must be the one the run was trained on"
|
||||
)
|
||||
probe = load_probe(run_dir, cfg)
|
||||
device = select_device(cfg.train.device)
|
||||
amp = resolve_amp(cfg.train.amp, device)
|
||||
probe = probe.to(device)
|
||||
return cfg, store, manifest, splits, probe, device, amp
|
||||
|
||||
|
||||
def _closed_payload(report: ClosedSetReport) -> dict[str, Any]:
|
||||
"""Convert a closed-set report into its JSON payload.
|
||||
|
||||
Args:
|
||||
report: The report to serialise.
|
||||
|
||||
Returns:
|
||||
Dict of the scalar metrics plus CIs as plain lists.
|
||||
|
||||
"""
|
||||
return {
|
||||
"recall_macro": report.recall_macro,
|
||||
"recall_micro": report.recall_micro,
|
||||
"precision_macro": report.precision_macro,
|
||||
"f1_macro": report.f1_macro,
|
||||
"ci": {k: [float(v[0]), float(v[1])] for k, v in report.ci.items()},
|
||||
}
|
||||
|
||||
|
||||
def evaluate_val(
|
||||
run_dir: Path, store_dir: Path | None = None, n_boot: int = 10_000
|
||||
) -> dict[str, Any]:
|
||||
"""Evaluate a run's checkpoint on its validation split.
|
||||
|
||||
Writes ``val_report.json`` (closed-set scalars + bootstrap CIs) into
|
||||
the run directory and returns the same payload.
|
||||
|
||||
Args:
|
||||
run_dir: Run directory (config + checkpoint).
|
||||
store_dir: Optional store override (default: recorded store).
|
||||
n_boot: Bootstrap resample count for CIs.
|
||||
|
||||
Returns:
|
||||
The val-report payload (run id, split, class/param counts,
|
||||
``"closed"`` block).
|
||||
|
||||
Raises:
|
||||
ValueError: If the validation split is empty or the store drifted
|
||||
(via :func:`_rehydrate`).
|
||||
|
||||
"""
|
||||
run_dir = Path(run_dir)
|
||||
cfg, store, _manifest, splits, probe, device, amp = _rehydrate(run_dir, store_dir)
|
||||
if len(splits.val) == 0:
|
||||
raise ValueError("validation split is empty")
|
||||
_, val_loader, _ = make_loaders(store, splits, cfg.train)
|
||||
val_logits, val_y = predict(probe, val_loader, device, amp)
|
||||
val_pred = val_logits.argmax(axis=1)
|
||||
report = closed_set_report(
|
||||
val_y, val_pred, splits.genera, n_boot=n_boot, seed=cfg.data.seed
|
||||
)
|
||||
payload = {
|
||||
"run_id": cfg.run_id,
|
||||
"split": "val",
|
||||
"n_classes": len(splits.genera),
|
||||
"params_realised": cfg.model.params_realised,
|
||||
"n_eval_windows": len(val_y),
|
||||
"closed": _closed_payload(report),
|
||||
}
|
||||
_write_json(run_dir / VAL_REPORT_NAME, payload)
|
||||
return payload
|
||||
|
||||
|
||||
def evaluate_test(
|
||||
run_dir: Path,
|
||||
store_dir: Path | None = None,
|
||||
trap_store_dir: Path | None = None,
|
||||
n_boot: int = 10_000,
|
||||
) -> dict[str, Any]:
|
||||
"""Evaluate a run's checkpoint on its test split (full artifact set).
|
||||
|
||||
Predicts val (for tau calibration) and test; writes
|
||||
``per_genus.parquet``, ``confusion.npz`` (matrix / genera / y_true /
|
||||
y_pred / msp), ``errors.parquet`` (per-error rows), ``open_set.json``,
|
||||
``report.json`` and — when a trap store is given —
|
||||
``trap_report.json`` (TP2: false-accept rates at tau, family-collapsed
|
||||
rates, MSP/margin histograms, top-3 retrieval, centroids).
|
||||
|
||||
Args:
|
||||
run_dir: Run directory (config + checkpoint).
|
||||
store_dir: Optional store override (default: recorded store).
|
||||
trap_store_dir: Optional eval-only trap store (its genera are
|
||||
never classes of the probe).
|
||||
n_boot: Bootstrap resample count for CIs.
|
||||
|
||||
Returns:
|
||||
The report payload: run id, class/param counts, ``"closed"`` and
|
||||
``"open_set"`` blocks, and ``"trap"`` (null when no trap store
|
||||
was supplied).
|
||||
|
||||
Raises:
|
||||
ValueError: On empty val/test splits or store drift (via
|
||||
:func:`_rehydrate`); open-set reporting errors propagate from
|
||||
:func:`~custom_models.metrics.open_set_probe_report` (e.g. no
|
||||
correct validation window — tau undefined).
|
||||
|
||||
"""
|
||||
run_dir = Path(run_dir)
|
||||
cfg, store, manifest, splits, probe, device, amp = _rehydrate(run_dir, store_dir)
|
||||
if len(splits.val) == 0 or len(splits.test) == 0:
|
||||
raise ValueError("val and test splits must both be non-empty")
|
||||
_, val_loader, test_loader = make_loaders(store, splits, cfg.train)
|
||||
val_logits, val_y = predict(probe, val_loader, device, amp)
|
||||
test_logits, test_y = predict(probe, test_loader, device, amp)
|
||||
test_pred = test_logits.argmax(axis=1)
|
||||
report = closed_set_report(
|
||||
test_y, test_pred, splits.genera, n_boot=n_boot, seed=cfg.data.seed
|
||||
)
|
||||
open_set = open_set_probe_report(val_logits, val_y, test_logits, test_y)
|
||||
test_msp = msp(test_logits)
|
||||
run_dir.mkdir(parents=True, exist_ok=True)
|
||||
report.per_genus.write_parquet(run_dir / PER_GENUS_NAME)
|
||||
np.savez_compressed(
|
||||
run_dir / CONFUSION_NAME,
|
||||
matrix=report.confusion,
|
||||
genera=np.array(splits.genera),
|
||||
y_true=test_y,
|
||||
y_pred=test_pred,
|
||||
msp=test_msp,
|
||||
)
|
||||
read_arr = manifest["read_id"].to_numpy()
|
||||
genus_arr = manifest["genus"].to_numpy()
|
||||
wid_arr = manifest["window_id"].to_numpy()
|
||||
read_by_wid = dict(zip(wid_arr.tolist(), read_arr.tolist()))
|
||||
genus_by_wid = dict(zip(wid_arr.tolist(), genus_arr.tolist()))
|
||||
wrong = np.flatnonzero(test_pred != test_y)
|
||||
errors = pl.DataFrame(
|
||||
{
|
||||
"window_id": splits.test[wrong],
|
||||
"read_id": [read_by_wid[int(w)] for w in splits.test[wrong]],
|
||||
"genus": [genus_by_wid[int(w)] for w in splits.test[wrong]],
|
||||
"pred": [splits.genera[int(p)] for p in test_pred[wrong]],
|
||||
"true_id": test_y[wrong],
|
||||
"pred_id": test_pred[wrong],
|
||||
"msp": test_msp[wrong],
|
||||
}
|
||||
)
|
||||
errors.write_parquet(run_dir / ERRORS_NAME)
|
||||
_write_json(run_dir / OPEN_SET_NAME, asdict(open_set))
|
||||
trap_payload: dict[str, Any] | None = None
|
||||
if trap_store_dir is not None:
|
||||
trap_store = SignalStore(trap_store_dir)
|
||||
trap_loader = make_full_loader(trap_store, cfg.train)
|
||||
trap_logits, _ = predict(probe, trap_loader, device, amp)
|
||||
trap_embeddings = embed_probe(probe, trap_loader, device, amp)
|
||||
trap_payload = trap_report(
|
||||
trap_store.manifest(),
|
||||
trap_logits,
|
||||
open_set.tau,
|
||||
splits.genera,
|
||||
embeddings=trap_embeddings,
|
||||
)
|
||||
_write_json(run_dir / TRAP_REPORT_NAME, trap_payload)
|
||||
payload = {
|
||||
"run_id": cfg.run_id,
|
||||
"split": "test",
|
||||
"n_classes": len(splits.genera),
|
||||
"params_realised": cfg.model.params_realised,
|
||||
"n_test_windows": len(test_y),
|
||||
"closed": _closed_payload(report),
|
||||
"open_set": asdict(open_set),
|
||||
"trap": (
|
||||
{
|
||||
"tau": trap_payload["tau"],
|
||||
"trap_fpr_at_tau": trap_payload["trap_fpr_at_tau"],
|
||||
"trap_fpr_family": trap_payload["trap_fpr_family"],
|
||||
"per_trap": trap_payload["per_trap"],
|
||||
}
|
||||
if trap_payload is not None
|
||||
else None
|
||||
),
|
||||
}
|
||||
_write_json(run_dir / REPORT_NAME, payload)
|
||||
return payload
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class GateVerdict:
|
||||
"""E1-v2 gate outcome over a set of runs plus the control run.
|
||||
|
||||
Attributes:
|
||||
verdict: One of ``PASS`` / ``MARGINAL`` / ``MARGINAL-TRAP`` /
|
||||
``FAIL`` / ``INVALID`` (v2 §6 four-line table).
|
||||
best_run: Run id of the run the verdict was decided on.
|
||||
best_recall_macro: That run's macro recall.
|
||||
ci: Its 95% bootstrap CI on macro recall.
|
||||
min_params_achieving_90: Smallest realised parameter count among
|
||||
runs with recall >= 0.90 (None if none reached it).
|
||||
trap_fpr: The deciding run's ``trap_fpr_at_tau`` (None when no
|
||||
trap evaluation was attached — the PASS branch then degrades
|
||||
to MARGINAL with an explanatory rationale).
|
||||
chance: ``1 / n_classes`` from the control run's config.
|
||||
control_recall_macro: The control run's macro recall (must sit at
|
||||
``<= chance_tol * chance`` or the verdict is INVALID).
|
||||
rationale: Human-readable verdict justification quoting the gate
|
||||
table.
|
||||
|
||||
"""
|
||||
|
||||
verdict: str
|
||||
best_run: str
|
||||
best_recall_macro: float
|
||||
ci: tuple[float, float]
|
||||
min_params_achieving_90: int | None
|
||||
trap_fpr: float | None
|
||||
chance: float
|
||||
control_recall_macro: float
|
||||
rationale: str
|
||||
|
||||
|
||||
def _read_run_summary(run_dir: Path) -> dict[str, Any]:
|
||||
"""Summarise one run directory for the gate.
|
||||
|
||||
Args:
|
||||
run_dir: Directory with ``config.json`` and ``report.json``.
|
||||
|
||||
Returns:
|
||||
Dict with ``run_id``, ``params`` (realised, falling back to the
|
||||
nominal budget), ``recall_macro``, ``ci`` and ``trap_fpr``.
|
||||
|
||||
Raises:
|
||||
FileNotFoundError: If either file is missing (the test command
|
||||
must run first).
|
||||
|
||||
"""
|
||||
run_dir = Path(run_dir)
|
||||
cfg = load_run(run_dir)
|
||||
report_path = run_dir / REPORT_NAME
|
||||
if not report_path.is_file():
|
||||
raise FileNotFoundError(f"run {run_dir} has no {REPORT_NAME}; run the test command first")
|
||||
report = json.loads(report_path.read_text())
|
||||
params = cfg.model.params_realised
|
||||
if params is None:
|
||||
params = cfg.model.param_budget
|
||||
trap_fpr = None
|
||||
if isinstance(report.get("trap"), dict):
|
||||
trap_fpr = report["trap"].get("trap_fpr_at_tau")
|
||||
ci = report["closed"].get("ci", {}).get("recall_macro")
|
||||
return {
|
||||
"run_id": cfg.run_id,
|
||||
"dir": str(run_dir),
|
||||
"params": int(params),
|
||||
"recall_macro": float(report["closed"]["recall_macro"]),
|
||||
"ci": (float(ci[0]), float(ci[1])) if ci else (float("nan"), float("nan")),
|
||||
"trap_fpr": trap_fpr,
|
||||
}
|
||||
|
||||
|
||||
def evaluate_gate(
|
||||
runs: Sequence[Path], control: Path, chance_tol: float = 1.5
|
||||
) -> GateVerdict:
|
||||
"""Decide the E1-v2 gate over run directories and write the verdict.
|
||||
|
||||
Logic (v2 §6, chance = 1/n_classes from the control run's config):
|
||||
the control exceeding ``chance_tol * chance`` is INVALID (suspected
|
||||
label leakage — fix before reading anything else). Otherwise, over
|
||||
the runs' *realised* params: PASS needs recall >= 0.90 at <= 5M with
|
||||
``trap_fpr_at_tau <= 0.05``; recall >= 0.90 at <= 5M with failing
|
||||
trap FPR is MARGINAL-TRAP (adopt per-class/margin thresholding);
|
||||
recall in [0.70, 0.90) at <= 5M, or 0.90 only reached at <= 20M with
|
||||
the trap criterion satisfied, is MARGINAL; anything else is FAIL.
|
||||
Missing trap evaluation degrades PASS to MARGINAL (cannot confirm the
|
||||
trap criterion). ``gate_verdict.json`` is written to the runs'
|
||||
common parent directory.
|
||||
|
||||
Args:
|
||||
runs: Run directories (each needs ``config.json`` +
|
||||
``report.json`` from the test command).
|
||||
control: The shuffled-label control run directory.
|
||||
chance_tol: Multiple of chance the control may not exceed.
|
||||
|
||||
Returns:
|
||||
The :class:`GateVerdict` (also written as JSON).
|
||||
|
||||
Raises:
|
||||
FileNotFoundError: If a run directory lacks its report.
|
||||
ValueError: If ``runs`` is empty.
|
||||
|
||||
"""
|
||||
control_dir = Path(control)
|
||||
control_summary = _read_run_summary(control_dir)
|
||||
control_cfg = load_run(control_dir)
|
||||
n_classes = len(control_cfg.genera)
|
||||
chance = 1.0 / n_classes
|
||||
control_recall = control_summary["recall_macro"]
|
||||
if control_recall > chance_tol * chance:
|
||||
verdict = GateVerdict(
|
||||
verdict="INVALID",
|
||||
best_run=control_summary["run_id"],
|
||||
best_recall_macro=control_recall,
|
||||
ci=control_summary["ci"],
|
||||
min_params_achieving_90=None,
|
||||
trap_fpr=None,
|
||||
chance=chance,
|
||||
control_recall_macro=control_recall,
|
||||
rationale=(
|
||||
f"control run recall {control_recall:.3f} exceeds {chance_tol}x chance "
|
||||
f"({chance:.3f}): suspected label leakage; fix before reading the ladder"
|
||||
),
|
||||
)
|
||||
else:
|
||||
entries = [_read_run_summary(Path(r)) for r in runs]
|
||||
if not entries:
|
||||
raise ValueError("no runs passed to the gate")
|
||||
by_5m = [e for e in entries if e["params"] <= 5_000_000]
|
||||
by_20m = [e for e in entries if e["params"] <= 20_000_000]
|
||||
best5 = max(by_5m, key=lambda e: e["recall_macro"]) if by_5m else None
|
||||
best20 = max(by_20m, key=lambda e: e["recall_macro"]) if by_20m else None
|
||||
min_params_90 = min(
|
||||
(e["params"] for e in entries if e["recall_macro"] >= 0.90), default=None
|
||||
)
|
||||
if best5 is not None and best5["recall_macro"] >= 0.90:
|
||||
chosen = best5
|
||||
if chosen["trap_fpr"] is not None and chosen["trap_fpr"] <= 0.05:
|
||||
verdict_str = "PASS"
|
||||
rationale = (
|
||||
f"best <= 5M run {chosen['run_id']} reaches recall_macro "
|
||||
f"{chosen['recall_macro']:.3f} with trap FPR {chosen['trap_fpr']:.3f} <= 0.05"
|
||||
)
|
||||
elif chosen["trap_fpr"] is not None:
|
||||
verdict_str = "MARGINAL-TRAP"
|
||||
rationale = (
|
||||
f"recall_macro {chosen['recall_macro']:.3f} at <= 5M params but trap FPR "
|
||||
f"{chosen['trap_fpr']:.3f} > 0.05: adopt per-class or margin thresholding "
|
||||
"before pod5-first is closed-set-safe"
|
||||
)
|
||||
else:
|
||||
verdict_str = "MARGINAL"
|
||||
rationale = (
|
||||
f"recall_macro {chosen['recall_macro']:.3f} at <= 5M params but no trap "
|
||||
"evaluation attached: cannot confirm the trap criterion"
|
||||
)
|
||||
elif best5 is not None and best5["recall_macro"] >= 0.70:
|
||||
chosen = best5
|
||||
verdict_str = "MARGINAL"
|
||||
rationale = (
|
||||
f"best <= 5M recall_macro {chosen['recall_macro']:.3f} in [0.70, 0.90): "
|
||||
"Stage-1 expectations drop a rank"
|
||||
)
|
||||
elif best20 is not None and best20["recall_macro"] >= 0.90:
|
||||
chosen = best20
|
||||
if chosen["trap_fpr"] is not None and chosen["trap_fpr"] <= 0.05:
|
||||
verdict_str = "MARGINAL"
|
||||
rationale = (
|
||||
f"0.90 only reached at <= 20M params ({chosen['run_id']}, "
|
||||
f"{chosen['params']} params) with trap FPR satisfied"
|
||||
)
|
||||
elif chosen["trap_fpr"] is not None:
|
||||
verdict_str = "MARGINAL-TRAP"
|
||||
rationale = (
|
||||
f"0.90 only at <= 20M params and trap FPR {chosen['trap_fpr']:.3f} > 0.05"
|
||||
)
|
||||
else:
|
||||
verdict_str = "MARGINAL"
|
||||
rationale = "0.90 only at <= 20M params; trap FPR unevaluated"
|
||||
else:
|
||||
best_overall = max(entries, key=lambda e: e["recall_macro"])
|
||||
chosen = best_overall
|
||||
verdict_str = "FAIL"
|
||||
rationale = (
|
||||
f"best recall_macro {chosen['recall_macro']:.3f} < 0.70 at <= 5M params and "
|
||||
"0.90 not reached at <= 20M: pod5-as-primary dead, demote to triage"
|
||||
)
|
||||
verdict = GateVerdict(
|
||||
verdict=verdict_str,
|
||||
best_run=chosen["run_id"],
|
||||
best_recall_macro=chosen["recall_macro"],
|
||||
ci=chosen["ci"],
|
||||
min_params_achieving_90=min_params_90,
|
||||
trap_fpr=chosen["trap_fpr"],
|
||||
chance=chance,
|
||||
control_recall_macro=control_recall,
|
||||
rationale=rationale,
|
||||
)
|
||||
dirs = [Path(r) for r in runs] + [Path(control)]
|
||||
common = os.path.commonpath([str(p.parent) for p in dirs])
|
||||
_write_json(Path(common) / GATE_NAME, asdict(verdict))
|
||||
return verdict
|
||||
@@ -0,0 +1,68 @@
|
||||
"""Determinism helpers: seeding, RNG factories and DataLoader worker init."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import random
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
|
||||
def seed_everything(seed: int, deterministic: bool = True) -> None:
|
||||
"""Seed every RNG the experiment draws from.
|
||||
|
||||
Seeds ``random``, ``numpy`` and ``torch`` (CPU plus all CUDA devices).
|
||||
``CUBLAS_WORKSPACE_CONFIG`` is set (via ``setdefault``) before any CUDA
|
||||
context exists so deterministic cuBLAS is possible on CUDA >= 10.2, and
|
||||
with ``deterministic`` torch is switched to deterministic algorithms in
|
||||
warn-only mode. Idempotent; call once at every CLI entrypoint before
|
||||
anything else (GATE ZERO §2.3 contract).
|
||||
|
||||
Args:
|
||||
seed: Base seed applied to every backend (numpy is folded into its
|
||||
32-bit domain).
|
||||
deterministic: Enable ``torch.use_deterministic_algorithms(True,
|
||||
warn_only=True)``; ops without deterministic implementations warn
|
||||
instead of raising.
|
||||
|
||||
"""
|
||||
os.environ.setdefault("CUBLAS_WORKSPACE_CONFIG", ":4096:8")
|
||||
random.seed(seed)
|
||||
np.random.seed(seed % (2**32))
|
||||
torch.manual_seed(seed)
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
if deterministic:
|
||||
torch.use_deterministic_algorithms(True, warn_only=True)
|
||||
|
||||
|
||||
def make_generator(seed: int) -> np.random.Generator:
|
||||
"""Create a fresh seeded numpy generator.
|
||||
|
||||
Args:
|
||||
seed: Seed for the new generator.
|
||||
|
||||
Returns:
|
||||
A ``numpy.random.default_rng`` generator seeded with ``seed``.
|
||||
|
||||
"""
|
||||
return np.random.default_rng(seed)
|
||||
|
||||
|
||||
def worker_init_fn(worker_id: int) -> None:
|
||||
"""Seed numpy/``random`` inside a DataLoader worker process.
|
||||
|
||||
Derives a per-worker seed from ``torch.initial_seed()`` (which the
|
||||
DataLoader already derived from the loader's own seeded generator), so
|
||||
worker-side randomness is reproducible from the run config. Pass as
|
||||
``DataLoader(worker_init_fn=...)``.
|
||||
|
||||
Args:
|
||||
worker_id: Worker index supplied by the DataLoader (unused beyond
|
||||
the signature contract; the seed comes from torch state).
|
||||
|
||||
"""
|
||||
seed = torch.initial_seed() % (2**32)
|
||||
np.random.seed(seed)
|
||||
random.seed(seed)
|
||||
@@ -0,0 +1,613 @@
|
||||
"""Signal store: pod5 extraction, the on-disk store format and its reader.
|
||||
|
||||
The store contract is the GATE ZERO schema (§4, §5.3 with the §12.1 binding
|
||||
deviations): ``manifest.parquet`` (one row per window, ``window_id`` equal to
|
||||
row order), ``shard_%05d.npy`` files holding stacked fixed-length
|
||||
median-IQR-normalised windows, and ``store_meta.json`` describing provenance
|
||||
and drop counts. A synthetic duty-cycle corpus builder provides the
|
||||
"labels survive normalisation" fixture (GATE ZERO §12.9) and smoke-test data.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections.abc import Iterator
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import polars as pl
|
||||
|
||||
from .config import ExtractConfig
|
||||
|
||||
LABEL_COLUMNS: tuple[str, ...] = (
|
||||
"read_id",
|
||||
"genus",
|
||||
"species",
|
||||
"ref_name",
|
||||
"identity",
|
||||
"aligned_len",
|
||||
"query_len",
|
||||
"source_run",
|
||||
)
|
||||
MANIFEST_COLUMNS: tuple[str, ...] = (
|
||||
"window_id",
|
||||
"read_id",
|
||||
"genus",
|
||||
"species",
|
||||
"source_run",
|
||||
"window_index",
|
||||
"n_samples",
|
||||
)
|
||||
|
||||
STORE_META = "store_meta.json"
|
||||
MANIFEST_NAME = "manifest.parquet"
|
||||
|
||||
|
||||
def median_iqr_normalise(x: np.ndarray) -> np.ndarray:
|
||||
"""Normalise a window by its robust centre and spread: ``(x - med) / IQR``.
|
||||
|
||||
Amplitude and offset are erased (a two-level square wave always maps to
|
||||
the same two normalised levels), which is why genus identity must be
|
||||
encoded as *shape* (duty cycle) in the synthetic corpus.
|
||||
|
||||
Args:
|
||||
x: 1-D signal window, any numeric dtype.
|
||||
|
||||
Returns:
|
||||
float32 normalised window.
|
||||
|
||||
Raises:
|
||||
ValueError: If the interquartile range is zero (dead/constant
|
||||
channel).
|
||||
|
||||
"""
|
||||
x = np.asarray(x, dtype=np.float64)
|
||||
median = np.median(x)
|
||||
p75, p25 = np.percentile(x, [75, 25])
|
||||
iqr = p75 - p25
|
||||
if iqr == 0:
|
||||
raise ValueError("median_iqr normalisation undefined for a constant window")
|
||||
return ((x - median) / iqr).astype(np.float32)
|
||||
|
||||
|
||||
def median_mad_normalise(x: np.ndarray) -> np.ndarray:
|
||||
"""Normalise a window as ``(x - med) / (1.4826 * MAD)``.
|
||||
|
||||
Alternative to :func:`median_iqr_normalise` with a Gaussian-consistent
|
||||
scale estimate.
|
||||
|
||||
Args:
|
||||
x: 1-D signal window, any numeric dtype.
|
||||
|
||||
Returns:
|
||||
float32 normalised window.
|
||||
|
||||
Raises:
|
||||
ValueError: If the median absolute deviation is zero.
|
||||
|
||||
"""
|
||||
x = np.asarray(x, dtype=np.float64)
|
||||
median = np.median(x)
|
||||
mad = np.median(np.abs(x - median))
|
||||
if mad == 0:
|
||||
raise ValueError("median_mad normalisation undefined for a constant window")
|
||||
return ((x - median) / (1.4826 * mad)).astype(np.float32)
|
||||
|
||||
|
||||
_NORMALISERS = {
|
||||
"median_iqr": median_iqr_normalise,
|
||||
"median_mad": median_mad_normalise,
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RawSignalRecord:
|
||||
"""One eagerly materialised pod5 read (pod5 lazy access dies with the reader).
|
||||
|
||||
Attributes:
|
||||
read_id: ONT read id string.
|
||||
signal: Calibrated picoampere trace, float32.
|
||||
num_samples: Length of ``signal`` in samples.
|
||||
sample_rate: Acquisition sample rate (Hz, default 5000 fallback).
|
||||
|
||||
"""
|
||||
|
||||
read_id: str
|
||||
signal: np.ndarray
|
||||
num_samples: int
|
||||
sample_rate: int
|
||||
|
||||
|
||||
def signal_pa(record: Any) -> np.ndarray:
|
||||
"""Return a pod5 record's signal in calibrated picoamps.
|
||||
|
||||
Prefers the reader's own ``signal_pa`` property when it exists and
|
||||
yields a non-empty 1-D array; otherwise applies the record's calibration
|
||||
manually as ``(raw + shift) * scale``.
|
||||
|
||||
Args:
|
||||
record: A ``pod5.ReadRecord`` (duck-typed; only attributes are used).
|
||||
|
||||
Returns:
|
||||
float32 picoampere trace.
|
||||
|
||||
"""
|
||||
if hasattr(record, "signal_pa"):
|
||||
try:
|
||||
values = np.asarray(record.signal_pa, dtype=np.float32)
|
||||
if values.ndim == 1 and values.size > 0:
|
||||
return values
|
||||
except (RuntimeError, ValueError, TypeError, AttributeError):
|
||||
pass
|
||||
shift = float(getattr(record, "shift", 0.0))
|
||||
scale = float(getattr(record, "scale", 1.0))
|
||||
return (np.asarray(record.signal, dtype=np.float32) + shift) * scale
|
||||
|
||||
|
||||
def read_pod5_records(path: Path) -> Iterator[RawSignalRecord]:
|
||||
"""Yield every read of a pod5 file as eager :class:`RawSignalRecord` objects.
|
||||
|
||||
The reader context is closed when the generator is exhausted or dropped,
|
||||
so all calibration/signal access happens inside it.
|
||||
|
||||
Args:
|
||||
path: Pod5 file path.
|
||||
|
||||
Yields:
|
||||
One record per read, in file order.
|
||||
|
||||
Raises:
|
||||
FileNotFoundError: If ``path`` does not exist.
|
||||
ImportError: If the optional ``pod5`` package is not installed.
|
||||
|
||||
"""
|
||||
path = Path(path)
|
||||
if not path.is_file():
|
||||
raise FileNotFoundError(f"pod5 file not found: {path}")
|
||||
try:
|
||||
import pod5
|
||||
except ImportError as exc:
|
||||
raise ImportError(
|
||||
"the pod5 package is required to read pod5 files (pip install pod5)"
|
||||
) from exc
|
||||
with pod5.Reader(str(path)) as reader:
|
||||
for record in reader.reads():
|
||||
signal = signal_pa(record)
|
||||
run_info = getattr(record, "run_info", None)
|
||||
sample_rate = int(getattr(run_info, "sample_rate", 5000))
|
||||
yield RawSignalRecord(
|
||||
read_id=str(record.read_id),
|
||||
signal=signal,
|
||||
num_samples=int(signal.shape[0]),
|
||||
sample_rate=sample_rate,
|
||||
)
|
||||
|
||||
|
||||
def load_labels(path: Path) -> pl.DataFrame:
|
||||
"""Load and validate a labels parquet (schema: ``LABEL_COLUMNS``).
|
||||
|
||||
Args:
|
||||
path: Labels parquet path.
|
||||
|
||||
Returns:
|
||||
The validated polars DataFrame.
|
||||
|
||||
Raises:
|
||||
FileNotFoundError: If ``path`` does not exist.
|
||||
ValueError: On missing columns or duplicate read ids.
|
||||
|
||||
"""
|
||||
path = Path(path)
|
||||
if not path.is_file():
|
||||
raise FileNotFoundError(f"labels file not found: {path}")
|
||||
df = pl.read_parquet(path)
|
||||
validate_labels(df)
|
||||
return df
|
||||
|
||||
|
||||
def validate_labels(df: pl.DataFrame) -> None:
|
||||
"""Check a labels DataFrame against the label-table contract.
|
||||
|
||||
Args:
|
||||
df: Candidate labels table (columns checked order-free).
|
||||
|
||||
Raises:
|
||||
ValueError: If any of ``LABEL_COLUMNS`` is missing or ``read_id``
|
||||
contains duplicates.
|
||||
|
||||
"""
|
||||
missing = sorted(set(LABEL_COLUMNS) - set(df.columns))
|
||||
if missing:
|
||||
raise ValueError(f"labels parquet missing columns: {missing}")
|
||||
if df["read_id"].n_unique() != df.height:
|
||||
raise ValueError("labels read_id column contains duplicates")
|
||||
|
||||
|
||||
class _StoreWriter:
|
||||
"""Streaming builder for a signal store directory.
|
||||
|
||||
Buffers normalised windows, flushes them into fixed-size
|
||||
``shard_%05d.npy`` files, accumulates manifest rows and finally writes
|
||||
``manifest.parquet`` plus ``store_meta.json``.
|
||||
"""
|
||||
|
||||
def __init__(self, out_dir: Path, window_samples: int, dtype: Any, shard_windows: int) -> None:
|
||||
"""Prepare the writer and create ``out_dir``.
|
||||
|
||||
Args:
|
||||
out_dir: Store directory to create/write into.
|
||||
window_samples: Fixed window length for every emitted window.
|
||||
dtype: numpy dtype for the on-disk shards (fp16 or fp32).
|
||||
shard_windows: Windows per shard file.
|
||||
|
||||
Raises:
|
||||
ValueError: If ``window_samples`` or ``shard_windows`` < 1.
|
||||
|
||||
"""
|
||||
if window_samples < 1:
|
||||
raise ValueError("window_samples must be >= 1")
|
||||
if shard_windows < 1:
|
||||
raise ValueError("shard_windows must be >= 1")
|
||||
self.out_dir = Path(out_dir)
|
||||
self.out_dir.mkdir(parents=True, exist_ok=True)
|
||||
self.window_samples = int(window_samples)
|
||||
self.dtype = dtype
|
||||
self.shard_windows = int(shard_windows)
|
||||
self._buffer: list[np.ndarray] = []
|
||||
self._rows: list[dict[str, Any]] = []
|
||||
self._shard_idx = 0
|
||||
self._window_id = 0
|
||||
self.n_windows = 0
|
||||
self.n_reads = 0
|
||||
|
||||
def emit(
|
||||
self,
|
||||
window: np.ndarray,
|
||||
read_id: str,
|
||||
genus: str,
|
||||
species: str | None,
|
||||
source_run: str,
|
||||
window_index: int,
|
||||
) -> None:
|
||||
"""Append one normalised window and its manifest row.
|
||||
|
||||
Args:
|
||||
window: Normalised signal window of length ``window_samples``.
|
||||
read_id: Owning read id.
|
||||
genus: Genus label (never null here; labelling pre-filtered).
|
||||
species: Optional species label (may be ``None``).
|
||||
source_run: Provenance run id for the read.
|
||||
window_index: Zero-based window position within the read.
|
||||
|
||||
"""
|
||||
self._buffer.append(np.asarray(window, dtype=np.dtype(self.dtype)))
|
||||
self._rows.append(
|
||||
{
|
||||
"window_id": self._window_id,
|
||||
"read_id": read_id,
|
||||
"genus": genus,
|
||||
"species": species,
|
||||
"source_run": source_run,
|
||||
"window_index": window_index,
|
||||
"n_samples": self.window_samples,
|
||||
}
|
||||
)
|
||||
self._window_id += 1
|
||||
self.n_windows += 1
|
||||
if len(self._buffer) >= self.shard_windows:
|
||||
self._flush()
|
||||
|
||||
def count_read(self) -> None:
|
||||
"""Record that one read contributed at least one window."""
|
||||
self.n_reads += 1
|
||||
|
||||
def _flush(self) -> None:
|
||||
"""Write the buffered windows as the next shard file, if any."""
|
||||
if not self._buffer:
|
||||
return
|
||||
arr = np.stack(self._buffer).astype(self.dtype)
|
||||
np.save(self.out_dir / f"shard_{self._shard_idx:05d}.npy", arr)
|
||||
self._shard_idx += 1
|
||||
self._buffer = []
|
||||
|
||||
def finish(self, meta: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Flush the tail shard and write manifest and store meta.
|
||||
|
||||
Args:
|
||||
meta: Extra provenance fields merged into ``store_meta.json``
|
||||
(window/count fields are added automatically).
|
||||
|
||||
Returns:
|
||||
The full store meta dict that was written.
|
||||
|
||||
Raises:
|
||||
ValueError: If no window was ever emitted.
|
||||
|
||||
"""
|
||||
self._flush()
|
||||
if self.n_windows == 0:
|
||||
raise ValueError("zero windows written to the signal store")
|
||||
pl.DataFrame(self._rows, schema=list(MANIFEST_COLUMNS)).write_parquet(
|
||||
self.out_dir / MANIFEST_NAME
|
||||
)
|
||||
meta = {
|
||||
"n_windows": self.n_windows,
|
||||
"n_reads": self.n_reads,
|
||||
"shard_windows": self.shard_windows,
|
||||
**meta,
|
||||
}
|
||||
(self.out_dir / STORE_META).write_text(json.dumps(meta, indent=2, sort_keys=True) + "\n")
|
||||
return meta
|
||||
|
||||
|
||||
def extract_store(cfg: ExtractConfig, labels: pl.DataFrame, out_dir: Path) -> Path:
|
||||
"""Build a signal store from pod5 files and a labels table.
|
||||
|
||||
For each sorted pod5 path, reads with a non-null genus label and at
|
||||
least ``min_read_samples`` samples are kept, sorted by read id, and
|
||||
their first ``max_windows_per_read`` contiguous windows (after skipping
|
||||
``skip_head_samples``) are normalised and written as shards. Fully
|
||||
deterministic: no RNG anywhere in the window choice. Drop counters
|
||||
(short/unlabelled) are reported in ``store_meta.json``.
|
||||
|
||||
Args:
|
||||
cfg: Extraction parameters (paths, windowing, normaliser, dtype,
|
||||
shard size).
|
||||
labels: Validated labels table (``LABEL_COLUMNS``); reads with null
|
||||
or empty genus are ignored.
|
||||
out_dir: Directory the store is written into (created as needed).
|
||||
|
||||
Returns:
|
||||
``out_dir`` (the store directory).
|
||||
|
||||
Raises:
|
||||
ValueError: On empty ``pod5_paths``, missing label columns,
|
||||
duplicate read ids, zero labelled reads, or zero surviving
|
||||
windows.
|
||||
FileNotFoundError: If a pod5 path is missing.
|
||||
ImportError: If the ``pod5`` package is unavailable.
|
||||
|
||||
"""
|
||||
validate_labels(labels)
|
||||
if not cfg.pod5_paths:
|
||||
raise ValueError("ExtractConfig.pod5_paths is empty")
|
||||
labelled = labels.filter(pl.col("genus").is_not_null() & (pl.col("genus") != ""))
|
||||
if labelled.height == 0:
|
||||
raise ValueError("no reads with a non-null genus in the labels table")
|
||||
info = {
|
||||
row[0]: (row[1], row[2], row[3])
|
||||
for row in zip(
|
||||
labelled["read_id"].to_list(),
|
||||
labelled["genus"].to_list(),
|
||||
labelled["species"].to_list(),
|
||||
labelled["source_run"].to_list(),
|
||||
)
|
||||
}
|
||||
normalise = _NORMALISERS[cfg.norm]
|
||||
dtype = np.dtype(np.float16 if cfg.dtype == "float16" else np.float32)
|
||||
writer = _StoreWriter(out_dir, cfg.window_samples, dtype, cfg.shard_windows)
|
||||
n_dropped_short = 0
|
||||
n_dropped_unlabelled = 0
|
||||
for pod5_path in sorted(Path(p) for p in cfg.pod5_paths):
|
||||
kept: list[RawSignalRecord] = []
|
||||
for record in read_pod5_records(pod5_path):
|
||||
if record.read_id not in info:
|
||||
n_dropped_unlabelled += 1
|
||||
continue
|
||||
if record.num_samples < cfg.min_read_samples:
|
||||
n_dropped_short += 1
|
||||
continue
|
||||
kept.append(record)
|
||||
kept.sort(key=lambda r: r.read_id)
|
||||
for record in kept:
|
||||
genus, species, source_run = info[record.read_id]
|
||||
for w in range(cfg.max_windows_per_read):
|
||||
start = cfg.skip_head_samples + w * cfg.window_samples
|
||||
if start + cfg.window_samples > record.num_samples:
|
||||
break
|
||||
window = normalise(record.signal[start : start + cfg.window_samples])
|
||||
writer.emit(window, record.read_id, genus, species, source_run, w)
|
||||
writer.count_read()
|
||||
writer.finish(
|
||||
{
|
||||
"synthetic": False,
|
||||
"window_samples": cfg.window_samples,
|
||||
"skip_head_samples": cfg.skip_head_samples,
|
||||
"max_windows_per_read": cfg.max_windows_per_read,
|
||||
"min_read_samples": cfg.min_read_samples,
|
||||
"dtype": str(cfg.dtype),
|
||||
"norm": cfg.norm,
|
||||
"pod5_paths": [str(p) for p in sorted(Path(p) for p in cfg.pod5_paths)],
|
||||
"n_labelled_reads": labelled.height,
|
||||
"n_dropped_short": n_dropped_short,
|
||||
"n_dropped_unlabelled": n_dropped_unlabelled,
|
||||
}
|
||||
)
|
||||
return Path(out_dir)
|
||||
|
||||
|
||||
def synthetic_store(
|
||||
out_dir: Path,
|
||||
n_genera: int = 11,
|
||||
reads_per_genus: int = 200,
|
||||
window_samples: int = 12_000,
|
||||
period: int | None = None,
|
||||
noise: float = 0.15,
|
||||
low: float = 1.0,
|
||||
high: float = 3.0,
|
||||
seed: int = 0,
|
||||
shard_windows: int = 8_192,
|
||||
genus_prefix: str = "genus_",
|
||||
) -> Path:
|
||||
"""Build a synthetic store whose classes are square-wave duty cycles.
|
||||
|
||||
Genus ``g`` of ``n_genera`` is encoded with duty cycle
|
||||
``(g + 1) / (n_genera + 1)`` over a square wave of ``period`` samples
|
||||
with a per-read random phase and Gaussian noise. The encoding survives
|
||||
median-IQR normalisation by construction (GATE ZERO §12.9), so a healthy
|
||||
probe trains to near-perfect recall while the shuffled-label control
|
||||
stays at chance. Windows go through the same writer/normalisation path
|
||||
as real extractions.
|
||||
|
||||
Args:
|
||||
out_dir: Store directory to create.
|
||||
n_genera: Number of classes (distinct duties).
|
||||
reads_per_genus: One window per synthetic read.
|
||||
window_samples: Window length (samples).
|
||||
period: Square-wave period; defaults to ``window_samples // 24``.
|
||||
noise: Gaussian noise sigma added to the raw levels.
|
||||
low: Raw low level (arbitrary units; erased by normalisation).
|
||||
high: Raw high level.
|
||||
seed: Seed for phase/noise draws.
|
||||
shard_windows: Windows per shard file.
|
||||
genus_prefix: Prefix for genus names (e.g. ``"trap_"`` for eval-only
|
||||
impostor corpora).
|
||||
|
||||
Returns:
|
||||
``out_dir`` (the store directory).
|
||||
|
||||
Raises:
|
||||
ValueError: On ``n_genera < 2``, ``reads_per_genus < 1`` or a
|
||||
degenerate ``period``.
|
||||
|
||||
"""
|
||||
if n_genera < 2:
|
||||
raise ValueError("synthetic corpus needs n_genera >= 2")
|
||||
if reads_per_genus < 1:
|
||||
raise ValueError("reads_per_genus must be >= 1")
|
||||
if period is None:
|
||||
period = max(4, window_samples // 24)
|
||||
if period < 2:
|
||||
raise ValueError("period must be >= 2")
|
||||
rng = np.random.default_rng(seed)
|
||||
dtype = np.dtype(np.float16)
|
||||
writer = _StoreWriter(out_dir, window_samples, dtype, shard_windows)
|
||||
for g in range(n_genera):
|
||||
genus = f"{genus_prefix}{g:02d}"
|
||||
duty = (g + 1) / (n_genera + 1)
|
||||
for r in range(reads_per_genus):
|
||||
read_id = f"{genus}_read_{r:06d}"
|
||||
phase = float(rng.uniform(0.0, period))
|
||||
offsets = (np.arange(window_samples, dtype=np.float64) + phase) % period
|
||||
signal = np.where(offsets < duty * period, high, low)
|
||||
signal = signal + rng.normal(0.0, noise, window_samples)
|
||||
window = median_iqr_normalise(signal)
|
||||
writer.emit(window, read_id, genus, None, "synthetic", 0)
|
||||
writer.count_read()
|
||||
writer.finish(
|
||||
{
|
||||
"synthetic": True,
|
||||
"window_samples": window_samples,
|
||||
"dtype": "float16",
|
||||
"norm": "median_iqr",
|
||||
"period": period,
|
||||
"noise": noise,
|
||||
"n_genera": n_genera,
|
||||
"reads_per_genus": reads_per_genus,
|
||||
"seed": seed,
|
||||
}
|
||||
)
|
||||
return Path(out_dir)
|
||||
|
||||
|
||||
class SignalStore:
|
||||
"""Random-access reader for a signal store directory.
|
||||
|
||||
Validates the manifest schema and meta on construction, then serves
|
||||
windows lazily from memory-mapped shards. Safe to share across
|
||||
forked DataLoader workers (read-only memmaps).
|
||||
"""
|
||||
|
||||
def __init__(self, store_dir: Path) -> None:
|
||||
"""Open and validate the store.
|
||||
|
||||
Args:
|
||||
store_dir: Directory holding ``manifest.parquet``,
|
||||
``store_meta.json`` and ``shard_%05d.npy`` files.
|
||||
|
||||
Raises:
|
||||
FileNotFoundError: If the manifest or store meta is missing.
|
||||
ValueError: If manifest columns are missing, ``window_id`` is
|
||||
not the row order, or required meta keys are absent.
|
||||
|
||||
"""
|
||||
self.dir = Path(store_dir)
|
||||
manifest_path = self.dir / MANIFEST_NAME
|
||||
if not manifest_path.is_file():
|
||||
raise FileNotFoundError(f"signal store manifest not found: {manifest_path}")
|
||||
self._manifest = pl.read_parquet(manifest_path)
|
||||
missing = sorted(set(MANIFEST_COLUMNS) - set(self._manifest.columns))
|
||||
if missing:
|
||||
raise ValueError(f"store manifest missing columns: {missing}")
|
||||
window_ids = self._manifest["window_id"].to_numpy()
|
||||
if not np.array_equal(window_ids, np.arange(self._manifest.height)):
|
||||
raise ValueError("store manifest window_id must equal the row order")
|
||||
meta_path = self.dir / STORE_META
|
||||
if not meta_path.is_file():
|
||||
raise FileNotFoundError(f"store meta not found: {meta_path}")
|
||||
self._meta = json.loads(meta_path.read_text())
|
||||
for key in ("window_samples", "shard_windows"):
|
||||
if key not in self._meta:
|
||||
raise ValueError(f"store_meta.json missing key: {key}")
|
||||
self._shard_windows = int(self._meta["shard_windows"])
|
||||
self._shards: dict[int, np.ndarray] = {}
|
||||
|
||||
def __len__(self) -> int:
|
||||
"""Return the number of windows in the store."""
|
||||
return self._manifest.height
|
||||
|
||||
def window_len(self) -> int:
|
||||
"""Return the fixed window length in samples."""
|
||||
return int(self._meta["window_samples"])
|
||||
|
||||
def manifest(self) -> pl.DataFrame:
|
||||
"""Return the (cached) manifest DataFrame, one row per window."""
|
||||
return self._manifest
|
||||
|
||||
def meta(self) -> dict[str, Any]:
|
||||
"""Return a copy of ``store_meta.json`` contents."""
|
||||
return dict(self._meta)
|
||||
|
||||
def _shard(self, index: int) -> np.ndarray:
|
||||
"""Return the (memory-mapped) shard array covering ``index``.
|
||||
|
||||
Args:
|
||||
index: Global window id.
|
||||
|
||||
Returns:
|
||||
The shard's 2-D array (opened and cached on first use).
|
||||
|
||||
Raises:
|
||||
FileNotFoundError: If the shard file is absent.
|
||||
|
||||
"""
|
||||
shard = index // self._shard_windows
|
||||
if shard not in self._shards:
|
||||
path = self.dir / f"shard_{shard:05d}.npy"
|
||||
if not path.is_file():
|
||||
raise FileNotFoundError(f"missing signal shard: {path}")
|
||||
self._shards[shard] = np.load(path, mmap_mode="r")
|
||||
return self._shards[shard]
|
||||
|
||||
def get(self, index: int) -> np.ndarray:
|
||||
"""Fetch one window as a detached float32 array.
|
||||
|
||||
Args:
|
||||
index: Global window id (0-based, the manifest row order).
|
||||
|
||||
Returns:
|
||||
float32 array of shape ``(window_len(),)``.
|
||||
|
||||
Raises:
|
||||
IndexError: If ``index`` is out of range for the store.
|
||||
|
||||
"""
|
||||
if index < 0 or index >= len(self):
|
||||
raise IndexError(f"window index out of range: {index}")
|
||||
arr = self._shard(index)
|
||||
row = index % self._shard_windows
|
||||
if row >= arr.shape[0]:
|
||||
raise IndexError(f"window {index} beyond shard rows")
|
||||
return np.asarray(arr[row], dtype=np.float32)
|
||||
@@ -0,0 +1,466 @@
|
||||
"""Training semantics: device/AMP policy, optimizer, loop and inference.
|
||||
|
||||
Encodes the training contract of the torch model spec §5: cross-entropy
|
||||
with optional label smoothing, AdamW with no decay on norm/bias
|
||||
parameters, linear-warmup cosine schedule to zero, L2 gradient clipping,
|
||||
per-epoch validation with early stopping on ``val_recall_macro`` (the
|
||||
headline gate metric) and best-epoch checkpointing to CPU tensors. All
|
||||
reported metrics are computed from CPU numpy logits so MPS/CUDA numeric
|
||||
differences cannot enter the reported numbers.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
import time
|
||||
from contextlib import ContextDecorator, nullcontext
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import polars as pl
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch.utils.data import DataLoader
|
||||
|
||||
from .config import TrainConfig
|
||||
from .data import Splits, make_loaders
|
||||
from .metrics import macro_recall
|
||||
from .models import GenusProbe
|
||||
from .store import SignalStore
|
||||
|
||||
|
||||
def select_device(requested: str | None) -> torch.device:
|
||||
"""Resolve the compute device (cuda -> mps -> cpu, never hard-coded CUDA).
|
||||
|
||||
Args:
|
||||
requested: ``None``/``""``/``"auto"`` picks the first available of
|
||||
cuda, mps, cpu; otherwise a torch device string such as
|
||||
``"cpu"``, ``"cuda"`` or ``"cuda:1"``.
|
||||
|
||||
Returns:
|
||||
The resolved ``torch.device``.
|
||||
|
||||
Raises:
|
||||
ValueError: If ``requested`` is not parseable as a device, names
|
||||
an unsupported type, or names cuda/mps on a machine where it
|
||||
is unavailable.
|
||||
|
||||
"""
|
||||
if requested is None or requested in ("", "auto"):
|
||||
if torch.cuda.is_available():
|
||||
return torch.device("cuda")
|
||||
if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
|
||||
return torch.device("mps")
|
||||
return torch.device("cpu")
|
||||
try:
|
||||
device = torch.device(requested)
|
||||
except RuntimeError as exc:
|
||||
raise ValueError(f"invalid device: {requested!r}") from exc
|
||||
if device.type == "cuda" and not torch.cuda.is_available():
|
||||
raise ValueError("cuda requested but torch.cuda.is_available() is False")
|
||||
if device.type == "mps" and not (
|
||||
hasattr(torch.backends, "mps") and torch.backends.mps.is_available()
|
||||
):
|
||||
raise ValueError("mps requested but not available")
|
||||
if device.type not in ("cuda", "mps", "cpu"):
|
||||
raise ValueError(f"unsupported device type: {device.type}")
|
||||
return device
|
||||
|
||||
|
||||
def resolve_amp(amp: str, device: torch.device) -> str:
|
||||
"""Map a requested AMP policy onto what this device actually supports.
|
||||
|
||||
Policy (spec §5): ``bf16`` where supported (CUDA with bf16 tensors,
|
||||
recent MPS) else silently fp32; ``fp16`` on CUDA only (GradScaler
|
||||
handles it); ``off`` is fp32 everywhere; no autocast on CPU for
|
||||
reproducibility. Idempotent — passing an already-resolved policy
|
||||
returns it unchanged.
|
||||
|
||||
Args:
|
||||
amp: Requested policy: ``"off"``, ``"bf16"`` or ``"fp16"`` (an
|
||||
already-resolved ``"fp32"`` also round-trips).
|
||||
device: Device the policy is resolved against.
|
||||
|
||||
Returns:
|
||||
Effective policy: ``"fp32"``, ``"bf16"`` or ``"fp16"``.
|
||||
|
||||
"""
|
||||
if amp == "off":
|
||||
return "fp32"
|
||||
if device.type == "cuda":
|
||||
if amp == "bf16":
|
||||
return "bf16" if torch.cuda.is_bf16_supported() else "fp32"
|
||||
if amp == "fp16":
|
||||
return "fp16"
|
||||
if device.type == "mps" and amp == "bf16":
|
||||
try:
|
||||
with torch.autocast(device_type="mps", dtype=torch.bfloat16):
|
||||
torch.tensor([1.0]).mul(2.0)
|
||||
return "bf16"
|
||||
except (RuntimeError, TypeError, ValueError):
|
||||
return "fp32"
|
||||
if amp == "fp16" and device.type != "cuda":
|
||||
print(
|
||||
f"warning: amp={amp!r} is CUDA-only, running fp32 on {device.type}"
|
||||
)
|
||||
return "fp32"
|
||||
return "fp32"
|
||||
|
||||
|
||||
def autocast_context(policy: str, device: torch.device) -> ContextDecorator:
|
||||
"""Build the autocast context for an effective AMP policy.
|
||||
|
||||
Args:
|
||||
policy: Effective policy from :func:`resolve_amp`
|
||||
(``"fp32"``/``"bf16"``/``"fp16"``).
|
||||
device: Device the ops run on.
|
||||
|
||||
Returns:
|
||||
A usable context manager — ``torch.autocast`` for bf16/fp16 on
|
||||
accelerators, a no-op context for fp32 and for any policy on CPU.
|
||||
|
||||
"""
|
||||
if policy in ("fp32", "off") or device.type == "cpu":
|
||||
return nullcontext()
|
||||
dtype = torch.bfloat16 if policy == "bf16" else torch.float16
|
||||
return torch.autocast(device_type=device.type, dtype=dtype)
|
||||
|
||||
|
||||
def build_optimizer(
|
||||
probe: GenusProbe, cfg: TrainConfig, steps_per_epoch: int
|
||||
) -> tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR]:
|
||||
"""Create AdamW with the no-decay split and warmup-cosine schedule.
|
||||
|
||||
Parameters with ``ndim <= 1`` or a ``bias`` name (all LayerNorm
|
||||
weights, every bias) go into a zero-decay group; everything else gets
|
||||
``cfg.weight_decay`` decoupled decay. The LambdaLR factor ramps
|
||||
linearly from 0 to 1 over ``warmup_frac`` of the total steps, then
|
||||
cosines from 1 to 0 across the remainder (the schedule therefore hits
|
||||
zero exactly at ``max_epochs``).
|
||||
|
||||
Args:
|
||||
probe: The probe whose parameters are optimised.
|
||||
cfg: Training config (lr, weight decay, epochs, warmup fraction).
|
||||
steps_per_epoch: Optimiser steps per epoch (``len(train_loader)``
|
||||
with ``drop_last=True``).
|
||||
|
||||
Returns:
|
||||
``(optimizer, scheduler)``; the caller steps the scheduler once
|
||||
per optimiser step.
|
||||
|
||||
"""
|
||||
decay: list[torch.nn.Parameter] = []
|
||||
no_decay: list[torch.nn.Parameter] = []
|
||||
for name, param in probe.named_parameters():
|
||||
if not param.requires_grad:
|
||||
continue
|
||||
if param.ndim <= 1 or name.endswith("bias"):
|
||||
no_decay.append(param)
|
||||
else:
|
||||
decay.append(param)
|
||||
groups = [
|
||||
{"params": no_decay, "weight_decay": 0.0},
|
||||
{"params": decay, "weight_decay": cfg.weight_decay},
|
||||
]
|
||||
optimizer = torch.optim.AdamW(groups, lr=cfg.lr)
|
||||
total_steps = max(1, steps_per_epoch * cfg.max_epochs)
|
||||
warmup_steps = max(1, int(cfg.warmup_frac * total_steps))
|
||||
|
||||
def lr_lambda(step: int) -> float:
|
||||
"""Linear-warmup cosine LR factor for LambdaLR.
|
||||
|
||||
Ramps from 0 → 1 over ``warmup_steps``, then cosines from 1 → 0
|
||||
across the remaining steps so the schedule hits zero at
|
||||
``max_epochs``.
|
||||
|
||||
Args:
|
||||
step: Current global optimiser step (0-based).
|
||||
|
||||
Returns:
|
||||
LR multiplier in [0, 1].
|
||||
|
||||
"""
|
||||
if step < warmup_steps:
|
||||
return (step + 1) / warmup_steps
|
||||
progress = (step - warmup_steps) / max(1, total_steps - warmup_steps)
|
||||
return 0.5 * (1.0 + math.cos(math.pi * min(1.0, progress)))
|
||||
|
||||
scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)
|
||||
return optimizer, scheduler
|
||||
|
||||
|
||||
def _mean_cross_entropy(logits: np.ndarray, y: np.ndarray) -> float:
|
||||
"""Mean cross-entropy of CPU logits (plain, no label smoothing).
|
||||
|
||||
Args:
|
||||
logits: ``(N, C)`` score matrix.
|
||||
y: ``(N,)`` int class ids.
|
||||
|
||||
Returns:
|
||||
The mean negative log-likelihood under a numerically stable
|
||||
log-softmax.
|
||||
|
||||
"""
|
||||
z = logits - logits.max(axis=1, keepdims=True)
|
||||
log_probs = z - np.log(np.exp(z).sum(axis=1, keepdims=True))
|
||||
return float(-log_probs[np.arange(len(y)), y].mean())
|
||||
|
||||
|
||||
@dataclass
|
||||
class TrainResult:
|
||||
"""Outcome of one training run.
|
||||
|
||||
Attributes:
|
||||
best_state: Weights of the best epoch (CPU tensor copies) — what
|
||||
``ckpt.pt`` stores; not final-epoch weights.
|
||||
best_epoch: 1-based epoch whose metric won (or -1 if training
|
||||
never improved).
|
||||
history: Per-epoch table: epoch, lr, train_loss, val_loss,
|
||||
val_acc, val_recall_macro, seconds (also written to
|
||||
``history.parquet``).
|
||||
seconds: Total wall time over all epochs run.
|
||||
|
||||
"""
|
||||
|
||||
best_state: dict[str, torch.Tensor]
|
||||
best_epoch: int
|
||||
history: pl.DataFrame
|
||||
seconds: float
|
||||
|
||||
|
||||
def train_probe(
|
||||
probe: GenusProbe,
|
||||
store: SignalStore,
|
||||
splits: Splits,
|
||||
cfg: TrainConfig,
|
||||
out_dir: Any,
|
||||
) -> TrainResult:
|
||||
"""Train a probe with early stopping and write ``history.parquet``.
|
||||
|
||||
Per epoch: full train pass (AMP per the resolved policy, GradScaler
|
||||
for fp16, unscaled L2 grad clipping), then validation via
|
||||
:func:`predict`; the early-stop metric improves or the patience
|
||||
counter advances. The best epoch's weights are copied to CPU and
|
||||
restored into the probe before returning. If the projected wall time
|
||||
over ``max_epochs`` exceeds 3 h, a warning is printed (the GPU-hour
|
||||
budget is a constraint, not a suggestion).
|
||||
|
||||
Args:
|
||||
probe: The probe to train (moved to the resolved device).
|
||||
store: Store the splits read from.
|
||||
splits: Split window ids and labels (val must be non-empty).
|
||||
cfg: Training semantics.
|
||||
out_dir: Directory for ``history.parquet`` (the run directory;
|
||||
created as needed).
|
||||
|
||||
Returns:
|
||||
The :class:`TrainResult` (best weights, history, timing).
|
||||
|
||||
Raises:
|
||||
ValueError: On invalid epochs/metric, a train split smaller than
|
||||
the batch size (``drop_last`` would empty it) or an empty
|
||||
train loader.
|
||||
RuntimeError: If the validation split is empty (nothing to select
|
||||
on).
|
||||
|
||||
"""
|
||||
import pathlib
|
||||
|
||||
out_dir = pathlib.Path(out_dir)
|
||||
if cfg.max_epochs < 1:
|
||||
raise ValueError("max_epochs must be >= 1")
|
||||
if cfg.early_stop_metric not in ("val_recall_macro", "val_loss"):
|
||||
raise ValueError(f"unknown early_stop_metric: {cfg.early_stop_metric}")
|
||||
if len(splits.val) == 0:
|
||||
raise RuntimeError("validation split is empty")
|
||||
if len(splits.train) < cfg.batch_size:
|
||||
raise ValueError(
|
||||
f"train split ({len(splits.train)}) smaller than batch size ({cfg.batch_size})"
|
||||
)
|
||||
device = select_device(cfg.device)
|
||||
amp = resolve_amp(cfg.amp, device)
|
||||
train_loader, val_loader, _ = make_loaders(store, splits, cfg)
|
||||
probe = probe.to(device)
|
||||
steps_per_epoch = len(train_loader)
|
||||
if steps_per_epoch == 0:
|
||||
raise ValueError("train loader is empty")
|
||||
optimizer, scheduler = build_optimizer(probe, cfg, steps_per_epoch)
|
||||
scaler = torch.amp.GradScaler(device.type) if amp == "fp16" else None
|
||||
maximize = cfg.early_stop_metric == "val_recall_macro"
|
||||
best_metric = -math.inf if maximize else math.inf
|
||||
best_epoch = -1
|
||||
best_state: dict[str, torch.Tensor] | None = None
|
||||
bad_epochs = 0
|
||||
rows: list[dict[str, float]] = []
|
||||
n_classes = int(probe.head.out_features)
|
||||
start = time.perf_counter()
|
||||
epochs_run = 0
|
||||
for epoch in range(1, cfg.max_epochs + 1):
|
||||
epoch_start = time.perf_counter()
|
||||
probe.train()
|
||||
running_loss = 0.0
|
||||
n_seen = 0
|
||||
for windows, labels in train_loader:
|
||||
windows = windows.to(device, non_blocking=device.type == "cuda")
|
||||
labels = labels.to(device, non_blocking=device.type == "cuda")
|
||||
optimizer.zero_grad(set_to_none=True)
|
||||
with autocast_context(amp, device):
|
||||
logits = probe(windows)
|
||||
loss = F.cross_entropy(
|
||||
logits, labels, label_smoothing=cfg.label_smoothing
|
||||
)
|
||||
if scaler is not None:
|
||||
scaler.scale(loss).backward()
|
||||
scaler.unscale_(optimizer)
|
||||
if cfg.grad_clip > 0:
|
||||
torch.nn.utils.clip_grad_norm_(
|
||||
(p for g in optimizer.param_groups for p in g["params"]),
|
||||
cfg.grad_clip,
|
||||
)
|
||||
scaler.step(optimizer)
|
||||
scaler.update()
|
||||
else:
|
||||
loss.backward()
|
||||
if cfg.grad_clip > 0:
|
||||
torch.nn.utils.clip_grad_norm_(
|
||||
(p for g in optimizer.param_groups for p in g["params"]),
|
||||
cfg.grad_clip,
|
||||
)
|
||||
optimizer.step()
|
||||
scheduler.step()
|
||||
running_loss += float(loss.item()) * windows.shape[0]
|
||||
n_seen += windows.shape[0]
|
||||
train_loss = running_loss / max(1, n_seen)
|
||||
val_logits, val_y = predict(probe, val_loader, device, amp)
|
||||
val_pred = val_logits.argmax(axis=1)
|
||||
val_loss = _mean_cross_entropy(val_logits, val_y)
|
||||
val_acc = float((val_pred == val_y).mean()) if len(val_y) else float("nan")
|
||||
val_recall = macro_recall(val_y, val_pred, n_classes)
|
||||
seconds = time.perf_counter() - epoch_start
|
||||
metric = val_recall if maximize else val_loss
|
||||
improved = metric > best_metric if maximize else metric < best_metric
|
||||
if improved:
|
||||
best_metric = metric
|
||||
best_epoch = epoch
|
||||
best_state = {
|
||||
key: value.detach().to("cpu").clone()
|
||||
for key, value in probe.state_dict().items()
|
||||
}
|
||||
bad_epochs = 0
|
||||
else:
|
||||
bad_epochs += 1
|
||||
rows.append(
|
||||
{
|
||||
"epoch": epoch,
|
||||
"lr": optimizer.param_groups[0]["lr"],
|
||||
"train_loss": train_loss,
|
||||
"val_loss": val_loss,
|
||||
"val_acc": val_acc,
|
||||
"val_recall_macro": val_recall,
|
||||
"seconds": seconds,
|
||||
}
|
||||
)
|
||||
epochs_run = epoch
|
||||
if bad_epochs >= cfg.patience:
|
||||
break
|
||||
total_seconds = time.perf_counter() - start
|
||||
if best_state is None:
|
||||
raise RuntimeError("training completed without a best epoch")
|
||||
probe.load_state_dict(best_state)
|
||||
history = pl.DataFrame(rows)
|
||||
out_dir.mkdir(parents=True, exist_ok=True)
|
||||
history.write_parquet(out_dir / "history.parquet")
|
||||
projected = (total_seconds / max(1, epochs_run)) * cfg.max_epochs
|
||||
if projected > 3 * 3600:
|
||||
print(
|
||||
f"warning: run projects to {projected / 3600:.1f} h over {cfg.max_epochs} epochs "
|
||||
"(budget guideline is 1-3 h per run)"
|
||||
)
|
||||
return TrainResult(
|
||||
best_state=best_state,
|
||||
best_epoch=best_epoch,
|
||||
history=history,
|
||||
seconds=total_seconds,
|
||||
)
|
||||
|
||||
|
||||
def predict(
|
||||
probe: GenusProbe,
|
||||
loader: DataLoader,
|
||||
device: torch.device,
|
||||
amp: str,
|
||||
) -> tuple[np.ndarray, np.ndarray]:
|
||||
"""Run the probe over a loader and collect logits on CPU.
|
||||
|
||||
The module is switched to eval mode for the pass and its prior mode is
|
||||
restored on exit; gradients are disabled throughout.
|
||||
|
||||
Args:
|
||||
probe: The probe to evaluate (already on ``device``).
|
||||
loader: Any window loader (its labels are returned alongside).
|
||||
device: Compute device.
|
||||
amp: Requested or already-resolved AMP policy (re-resolved
|
||||
idempotently).
|
||||
|
||||
Returns:
|
||||
``(logits (N, C) float32 CPU, labels (N,) int64)`` in loader
|
||||
order; empty loaders yield correctly-shaped empty arrays.
|
||||
|
||||
"""
|
||||
policy = resolve_amp(amp, device)
|
||||
was_training = probe.training
|
||||
probe.eval()
|
||||
logits_chunks: list[np.ndarray] = []
|
||||
label_chunks: list[np.ndarray] = []
|
||||
with torch.no_grad():
|
||||
for windows, labels in loader:
|
||||
windows = windows.to(device, non_blocking=device.type == "cuda")
|
||||
with autocast_context(policy, device):
|
||||
out = probe(windows)
|
||||
logits_chunks.append(out.detach().to("cpu").float().numpy())
|
||||
label_chunks.append(labels.numpy())
|
||||
probe.train(was_training)
|
||||
if not logits_chunks:
|
||||
return (
|
||||
np.zeros((0, int(probe.head.out_features)), dtype=np.float32),
|
||||
np.zeros(0, dtype=np.int64),
|
||||
)
|
||||
return np.concatenate(logits_chunks), np.concatenate(label_chunks)
|
||||
|
||||
|
||||
def embed_probe(
|
||||
probe: GenusProbe,
|
||||
loader: DataLoader,
|
||||
device: torch.device,
|
||||
amp: str,
|
||||
) -> np.ndarray:
|
||||
"""Run the probe's ``embed()`` over a loader and collect embeddings.
|
||||
|
||||
Companion to :func:`predict` for the E3 reuse path (e.g. trap centroid
|
||||
vectors); same eval-mode/grad-free semantics.
|
||||
|
||||
Args:
|
||||
probe: The probe whose encoder embedding is extracted.
|
||||
loader: Any window loader (labels ignored).
|
||||
device: Compute device.
|
||||
amp: Requested or already-resolved AMP policy.
|
||||
|
||||
Returns:
|
||||
Embeddings of shape ``(N, d_embed)`` float32 on CPU.
|
||||
|
||||
"""
|
||||
policy = resolve_amp(amp, device)
|
||||
was_training = probe.training
|
||||
probe.eval()
|
||||
chunks: list[np.ndarray] = []
|
||||
with torch.no_grad():
|
||||
for windows, _ in loader:
|
||||
windows = windows.to(device, non_blocking=device.type == "cuda")
|
||||
with autocast_context(policy, device):
|
||||
out = probe.embed(windows)
|
||||
chunks.append(out.detach().to("cpu").float().numpy())
|
||||
probe.train(was_training)
|
||||
if not chunks:
|
||||
return np.zeros((0, int(probe.head.in_features)), dtype=np.float32)
|
||||
return np.concatenate(chunks)
|
||||
Reference in New Issue
Block a user