- 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
407 lines
15 KiB
Python
407 lines
15 KiB
Python
"""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,
|
|
)
|