oncodsl / engine /objectives.py
govindbalki's picture
Upload folder using huggingface_hub
0fff343 verified
Raw
History Blame Contribute Delete
7.95 kB
"""Objective specs — what "fitness" means for a program.
The engine itself stays name-blind; each objective wraps the math for one
target+metric combination so the GP loop can dispatch generically.
Two concrete objectives:
- ``BinaryAUROCObjective`` — binary y; 5-fold CV AUROC; prefilter ranks
by ``|AUROC − 0.5|``. The MSI-H/MSS path.
- ``CorrelationObjective(direction)`` — continuous y; per-fold linear
fit then signed Spearman on the held-out fold; prefilter ranks by
``|spearman|``. The mutation-burden path (direction="neg" makes
"lower score ↔ higher TMB" the winning shape).
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Literal
import numpy as np
import pandas as pd
from sklearn.linear_model import LogisticRegression
from sklearn.metrics import roc_auc_score
from sklearn.model_selection import KFold, StratifiedKFold
from scipy.stats import spearmanr
Target = Literal["msi", "tmb"]
Metric = Literal["auroc", "correlation"]
Direction = Literal["neg", "pos"]
class Objective:
"""Abstract base. Subclasses implement the four hook methods."""
target: Target
metric: Metric
direction: Direction | None
binary: bool = False # True for classification (stratified split + StratifiedKFold)
# --- API used by the engine -------------------------------------------
def cv_score(
self,
states: np.ndarray,
y: np.ndarray,
*,
n_folds: int,
random_state: int,
) -> float:
"""Higher = better. `states` is (n_samples, n_sets)."""
raise NotImplementedError
def holdout_score(
self,
states_train: np.ndarray,
y_train: np.ndarray,
states_test: np.ndarray,
y_test: np.ndarray,
) -> float:
"""Single-shot final score on the TEST split. Higher = better."""
raise NotImplementedError
def prefilter_score_per_feature(
self,
ranks: pd.DataFrame,
y: np.ndarray,
) -> np.ndarray:
"""Per-feature univariate ranking signal; higher = more discriminative."""
raise NotImplementedError
def permute(self, y: np.ndarray, rng: np.random.Generator) -> np.ndarray:
"""Return a permuted copy of y for the null distribution."""
return rng.permutation(y)
def fitness_label(self) -> str:
raise NotImplementedError
def to_dict(self) -> dict:
d: dict = {"target": self.target, "metric": self.metric}
if self.direction is not None:
d["direction"] = self.direction
return d
# ---------------------------------------------------------------------------
# Binary AUROC objective (MSI-H vs MSS)
# ---------------------------------------------------------------------------
@dataclass
class BinaryAUROCObjective(Objective):
target: Target = "msi"
metric: Metric = "auroc"
direction: Direction | None = None
binary: bool = True
def cv_score(
self,
states: np.ndarray,
y: np.ndarray,
*,
n_folds: int = 5,
random_state: int = 0,
) -> float:
skf = StratifiedKFold(n_splits=n_folds, shuffle=True, random_state=random_state)
aurocs = []
for tr, te in skf.split(states, y):
lr = LogisticRegression(max_iter=1000)
lr.fit(states[tr], y[tr])
proba = lr.predict_proba(states[te])[:, 1]
aurocs.append(roc_auc_score(y[te], proba))
return float(np.mean(aurocs))
def holdout_score(
self,
states_train: np.ndarray,
y_train: np.ndarray,
states_test: np.ndarray,
y_test: np.ndarray,
) -> float:
lr = LogisticRegression(max_iter=1000)
lr.fit(states_train, y_train)
proba = lr.predict_proba(states_test)[:, 1]
return float(roc_auc_score(y_test, proba))
def prefilter_score_per_feature(
self,
ranks: pd.DataFrame,
y: np.ndarray,
) -> np.ndarray:
# Mann-Whitney U / (n_pos * n_neg) per column, then |·−0.5|.
y_bool = np.asarray(y, dtype=bool)
n_pos = int(y_bool.sum())
n_neg = int(len(y_bool) - n_pos)
if n_pos == 0 or n_neg == 0:
return np.full(ranks.shape[1], np.nan)
sum_ranks_pos = ranks.iloc[y_bool].sum(axis=0).values
U = sum_ranks_pos - n_pos * (n_pos + 1) / 2
auroc = U / (n_pos * n_neg)
return np.abs(auroc - 0.5)
def fitness_label(self) -> str:
return "AUROC"
# ---------------------------------------------------------------------------
# Correlation objective (continuous y, signed)
# ---------------------------------------------------------------------------
@dataclass
class CorrelationObjective(Objective):
"""Continuous y; per-fold Spearman on the prediction.
``direction="neg"`` means a negative correlation between the program
score and the target is what we WANT — we flip the sign so the GP can
still maximise. ``"pos"`` is the natural direction.
"""
direction: Direction = "neg"
target: Target = "tmb"
metric: Metric = "correlation"
binary: bool = False
def _sign(self) -> float:
return -1.0 if self.direction == "neg" else 1.0
def cv_score(
self,
states: np.ndarray,
y: np.ndarray,
*,
n_folds: int = 5,
random_state: int = 0,
) -> float:
# Sum of per-set scores into a single per-patient number — preserves
# direction, unlike a LR fit that aligns predictions with y regardless.
kf = KFold(n_splits=n_folds, shuffle=True, random_state=random_state)
scores = []
combined = states.sum(axis=1)
for _tr, te in kf.split(states):
corr, _ = spearmanr(combined[te], y[te])
if np.isnan(corr):
corr = 0.0
scores.append(self._sign() * float(corr))
return float(np.mean(scores))
def holdout_score(
self,
states_train: np.ndarray,
y_train: np.ndarray,
states_test: np.ndarray,
y_test: np.ndarray,
) -> float:
combined = states_test.sum(axis=1)
corr, _ = spearmanr(combined, y_test)
if np.isnan(corr):
corr = 0.0
return self._sign() * float(corr)
def prefilter_score_per_feature(
self,
ranks: pd.DataFrame,
y: np.ndarray,
) -> np.ndarray:
# Spearman of each column vs y: pearson of ranks. Higher |·| = better.
y_ranks = pd.Series(y).rank().values
y_centered = y_ranks - y_ranks.mean()
y_norm = float(np.sqrt((y_centered ** 2).sum()))
if y_norm == 0.0:
return np.zeros(ranks.shape[1])
X = ranks.values
Xc = X - X.mean(axis=0)
x_norms = np.sqrt((Xc ** 2).sum(axis=0))
x_norms[x_norms == 0.0] = 1.0 # avoid division by zero on constant cols
corr = (Xc.T @ y_centered) / (x_norms * y_norm)
return np.abs(corr)
def fitness_label(self) -> str:
return "|spearman|"
# ---------------------------------------------------------------------------
# Build an objective from a spec dict (used by the API).
# ---------------------------------------------------------------------------
def objective_from_spec(spec: dict) -> Objective:
target = spec.get("target")
metric = spec.get("metric")
if target == "msi" and metric == "auroc":
return BinaryAUROCObjective()
if target == "tmb" and metric == "correlation":
direction = spec.get("direction", "neg")
if direction not in ("neg", "pos"):
raise ValueError(f"direction must be 'neg' or 'pos', got {direction!r}")
return CorrelationObjective(direction=direction)
raise ValueError(
f"Unsupported objective spec {spec!r}. "
"Stage 1 supports: msi+auroc, tmb+correlation+(neg|pos)."
)