Initial release: E1 genus probe package moved out of the E01 workflow

- src/custom_models torch modules, store, training, metrics, runner and CLI
  (package-relative imports for pip installability)
- test suite (49 tests) + conftest with fast synthetic-store fixtures
- uv-managed dev env (pyproject + CPU torch index), hatchling build
- README: uv for the package, pip install into conda envs for the workflow
This commit is contained in:
Tom Kasper
2026-09-22 13:22:50 +01:00
commit 88e98effad
18 changed files with 5988 additions and 0 deletions
+406
View File
@@ -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,
)