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,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,
|
||||
)
|
||||
Reference in New Issue
Block a user