oncodsl / engine_v2 /nodes.py
govindbalki's picture
Upload folder using huggingface_hub
0fff343 verified
Raw
History Blame Contribute Delete
21.2 kB
"""Typed AST nodes for engine_v2 — full DSL grammar.
A program is a tree of ``Node`` instances. Every node knows its return
type (Matrix / Vector / Scalar / Model) and how to ``execute`` against
an ``ExecContext`` that bundles the opaque-ID matrix with the clinical
fields and labels engine_v2 is allowed to see (stage / age / msi / tmb).
Strict airgap: ``FeatureSet`` leaves carry only opaque IDs; gene-name
strings never appear in the tree or in any payload returned here. Only
the named clinical fields and label columns appear, and only by name
(``msi``, ``tmb``, ``stage``, ``age``), never as column dumps.
"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Iterator
import numpy as np
import pandas as pd
from sklearn.linear_model import LogisticRegression
from engine_v2.types import TType
# ---------------------------------------------------------------------------
# Execution context — what every node may read
# ---------------------------------------------------------------------------
@dataclass
class ExecContext:
"""The view of the cohort the engine is allowed to see at execute time.
- ``M`` : full opaque-ID expression matrix (samples × features).
- ``clinical``: DataFrame indexed like ``M`` with at least ``stage``
(categorical string) and ``age`` (float).
- ``labels`` : dict mapping target name → 1-D numpy array aligned with
``M.index``. Convention: ``"msi"`` is binary 0/1 (1 =
MSI-H); ``"tmb"`` is the continuous mutation count.
- ``fit_ctx`` : optional sibling context used by ``FitApply`` so it
fits its model on TRAIN data and applies the frozen
model to this context's inputs — never fits on test
labels. The pipeline sets this on the test ctx so
``evaluate_holdout`` is honest. Default None (legacy
behaviour: fit and apply on the same ctx).
- ``confounders``: the ORDERED list of clinical column names that
``Effect`` regresses out before measuring its
correlation. Default ``("stage","age")`` preserves
the original behaviour (MSI / TMB / HPV runs are
byte-for-byte unchanged). HNSC cohorts can widen
this to e.g. ``("stage","age","sex","race")`` —
columns that aren't in ``clinical`` are skipped.
Names only — gene identities never enter this set.
"""
M: pd.DataFrame
clinical: pd.DataFrame
labels: dict[str, np.ndarray]
fit_ctx: "ExecContext | None" = None
confounders: tuple[str, ...] = ("stage", "age")
# ---------------------------------------------------------------------------
# Base node
# ---------------------------------------------------------------------------
@dataclass
class Node:
"""Abstract base. Concrete subclasses set ``ttype``."""
ttype: TType = field(init=False)
def children(self) -> list["Node"]:
return []
def depth(self) -> int:
ch = self.children()
return 1 + (max((c.depth() for c in ch), default=0))
def node_count(self) -> int:
return 1 + sum(c.node_count() for c in self.children())
def walk(self) -> Iterator["Node"]:
yield self
for c in self.children():
yield from c.walk()
def feature_ids(self) -> list[str]:
# FeatureSet is a leaf payload (not a Node), reached via Select.
out: list[str] = []
for n in self.walk():
if isinstance(n, Select):
out.extend(n.features.ids)
return out
def repr_typed(self) -> str:
raise NotImplementedError
def execute(self, ctx: ExecContext):
raise NotImplementedError
# ---------------------------------------------------------------------------
# Leaf payloads
# ---------------------------------------------------------------------------
@dataclass
class FeatureSet:
"""A non-empty list of opaque gene IDs. Leaf payload of ``Select``."""
ids: list[str]
def __post_init__(self) -> None:
if not self.ids:
raise ValueError("FeatureSet must have at least 1 gene ID")
seen: set[str] = set()
clean: list[str] = []
for g in self.ids:
if g not in seen:
seen.add(g)
clean.append(g)
self.ids = clean
def repr_typed(self) -> str:
return "[" + ",".join(self.ids) + "]"
# ---------------------------------------------------------------------------
# Matrix nodes
# ---------------------------------------------------------------------------
@dataclass
class MatrixTerminal(Node):
"""The full anonymised expression matrix (rows=patients, cols=opaque IDs)."""
def __post_init__(self) -> None:
self.ttype = TType.MATRIX
def repr_typed(self) -> str:
return "M"
def execute(self, ctx: ExecContext) -> pd.DataFrame:
return ctx.M
@dataclass
class Select(Node):
"""``Select(Matrix, FeatureSet) -> Matrix``."""
matrix: Node
features: FeatureSet
def __post_init__(self) -> None:
self.ttype = TType.MATRIX
def children(self) -> list[Node]:
return [self.matrix]
def repr_typed(self) -> str:
return f"Select({self.matrix.repr_typed()},{self.features.repr_typed()})"
def execute(self, ctx: ExecContext) -> pd.DataFrame:
sub = self.matrix.execute(ctx)
keep = [g for g in self.features.ids if g in sub.columns]
if not keep:
# Degenerate Select — preserve type, but score will be flat.
return sub.iloc[:, :0]
return sub.loc[:, keep]
@dataclass
class Search(Node):
"""``Search(Matrix, k) -> Matrix`` — bounded nested search.
A small, capped univariate ranker that picks the top-k columns of the
incoming Matrix by absolute Spearman correlation with the engine's
current target. Counts toward depth/node budgets like any other node;
introduced by the GP only at the configured low rate. See A4.
"""
matrix: Node
k: int
def __post_init__(self) -> None:
self.ttype = TType.MATRIX
def children(self) -> list[Node]:
return [self.matrix]
def repr_typed(self) -> str:
return f"Search({self.matrix.repr_typed()},{self.k})"
def execute(self, ctx: ExecContext) -> pd.DataFrame:
from engine_v2.types import SEARCH_MAX_COLS, SEARCH_MAX_K
sub = self.matrix.execute(ctx)
if sub.shape[1] == 0:
return sub
# Cap aggressively to keep nested cost bounded.
if sub.shape[1] > SEARCH_MAX_COLS:
sub = sub.iloc[:, :SEARCH_MAX_COLS]
k = min(max(1, self.k), SEARCH_MAX_K, sub.shape[1])
# Pick a target signal — prefer msi if available, else tmb.
target = ctx.labels.get("msi")
if target is None:
target = ctx.labels.get("tmb")
if target is None or len(target) == 0:
return sub.iloc[:, :k]
# Rank columns by |spearman with target|, on the matrix the
# Search node received (TRAIN by construction in our pipeline).
try:
from engine.prefilter import precompute_ranks
r = precompute_ranks(sub)
t = pd.Series(target).rank().values
y_centered = t - t.mean()
denom_y = float(np.sqrt((y_centered ** 2).sum())) or 1.0
X = r.values - r.values.mean(axis=0)
denom_x = np.sqrt((X ** 2).sum(axis=0))
denom_x[denom_x == 0.0] = 1.0
corr = (X.T @ y_centered) / (denom_x * denom_y)
scores = pd.Series(np.abs(corr), index=sub.columns)
top = scores.nlargest(k).index.tolist()
except Exception:
top = list(sub.columns[:k])
return sub.loc[:, top]
# ---------------------------------------------------------------------------
# Vector nodes
# ---------------------------------------------------------------------------
@dataclass
class Reduce(Node):
"""``Reduce(Matrix, Agg) -> Vector``."""
matrix: Node
agg: str
def __post_init__(self) -> None:
self.ttype = TType.VECTOR
def children(self) -> list[Node]:
return [self.matrix]
def repr_typed(self) -> str:
return f"Reduce({self.matrix.repr_typed()},{self.agg})"
def execute(self, ctx: ExecContext) -> pd.Series:
sub = self.matrix.execute(ctx)
if sub.shape[1] == 0:
return pd.Series(np.zeros(sub.shape[0]), index=sub.index)
if self.agg == "mean":
return sub.mean(axis=1)
if self.agg == "median":
return sub.median(axis=1)
if self.agg == "max":
return sub.max(axis=1)
if self.agg == "min":
return sub.min(axis=1)
if self.agg == "var":
return sub.var(axis=1)
raise ValueError(f"Reduce: unknown agg={self.agg!r}")
@dataclass
class Combine(Node):
"""``Combine(Vector, Vector, Op) -> Vector``."""
left: Node
right: Node
op: str
def __post_init__(self) -> None:
self.ttype = TType.VECTOR
def children(self) -> list[Node]:
return [self.left, self.right]
def repr_typed(self) -> str:
return f"Combine({self.left.repr_typed()},{self.right.repr_typed()},{self.op})"
def execute(self, ctx: ExecContext) -> pd.Series:
a = self.left.execute(ctx)
b = self.right.execute(ctx)
a, b = a.align(b, join="inner")
if self.op == "add":
return a + b
if self.op == "sub":
return a - b
if self.op == "mul":
return a * b
if self.op == "mean":
return (a + b) / 2.0
if self.op == "protected_div":
denom = b.where(b.abs() > 1e-9, 1e-9)
return a / denom
raise ValueError(f"Combine: unknown op={self.op!r}")
@dataclass
class Split(Node):
"""``Split(Vector, Predicate) -> Vector``.
Partitions the per-patient input into two groups, applies a different
Reduce-style transform per branch, and recombines into one vector
indexed like the input. ONE level of Split only — synthesis never
nests Split-in-Split. We use a minimal closed-form per branch
(mean-centering inside each side) to keep the operator deterministic
and self-contained.
"""
inner: Node # the Vector being split
predicate: str # one of PREDICATE_KINDS
min_subgroup: int = 5 # hard guard
def __post_init__(self) -> None:
self.ttype = TType.VECTOR
def children(self) -> list[Node]:
return [self.inner]
def repr_typed(self) -> str:
return f"Split({self.inner.repr_typed()},{self.predicate})"
def _mask(self, ctx: ExecContext, v: pd.Series) -> pd.Series:
if self.predicate == "score":
return v >= v.median()
if self.predicate == "stage_late":
stage = ctx.clinical.reindex(v.index).get("stage")
if stage is None:
return pd.Series(False, index=v.index)
return stage.astype(str).str.upper().isin({"III", "IV", "STAGE III", "STAGE IV"})
return pd.Series(False, index=v.index)
def execute(self, ctx: ExecContext) -> pd.Series:
v = self.inner.execute(ctx)
if not isinstance(v, pd.Series):
return v
mask = self._mask(ctx, v).fillna(False).astype(bool)
if mask.sum() < self.min_subgroup or (~mask).sum() < self.min_subgroup:
# Subgroup too small — return the input unchanged so the rest
# of the program can still execute. Fitness will weigh in.
return v.astype(float)
out = v.astype(float).copy()
a = out[mask]
b = out[~mask]
# Mean-centre within each side so a downstream Combine sees a
# contrast rather than a level shift.
out.loc[mask] = a - a.mean()
out.loc[~mask] = b - b.mean()
return out
# ---------------------------------------------------------------------------
# Scalar nodes
# ---------------------------------------------------------------------------
def _align_for_assoc(v: pd.Series, y: np.ndarray):
"""Drop NaNs and align lengths."""
arr = np.asarray(y, dtype=float)
if len(arr) != len(v):
# Reindex y to v's index if possible — otherwise trim.
n = min(len(arr), len(v))
arr = arr[:n]
v = v.iloc[:n]
df = pd.DataFrame({"v": v.astype(float).values, "y": arr}).dropna()
return df["v"].values, df["y"].values
def _spearman_corr(a: np.ndarray, b: np.ndarray) -> float:
"""Spearman correlation with NaN → 0 guard."""
if a.size < 3 or b.size < 3:
return 0.0
s = pd.Series(a).rank().values
t = pd.Series(b).rank().values
if s.std() == 0 or t.std() == 0:
return 0.0
return float(np.corrcoef(s, t)[0, 1])
def _pearson_corr(a: np.ndarray, b: np.ndarray) -> float:
if a.size < 3 or b.size < 3 or a.std() == 0 or b.std() == 0:
return 0.0
return float(np.corrcoef(a, b)[0, 1])
@dataclass
class Associate(Node):
"""``Associate(Vector, target, kind) -> Scalar``.
Plain observational correlation. ``target`` names a label column
(``msi`` or ``tmb``); ``kind`` is ``pearson`` or ``spearman``.
"""
inner: Node
target: str
kind: str = "spearman"
def __post_init__(self) -> None:
self.ttype = TType.SCALAR
def children(self) -> list[Node]:
return [self.inner]
def repr_typed(self) -> str:
return f"Associate({self.inner.repr_typed()},{self.target},{self.kind})"
def execute(self, ctx: ExecContext) -> float:
v = self.inner.execute(ctx)
if not isinstance(v, pd.Series):
return 0.0
y = ctx.labels.get(self.target)
if y is None:
return 0.0
a, b = _align_for_assoc(v, y)
if a.size == 0:
return 0.0
if self.kind == "pearson":
return _pearson_corr(a, b)
return _spearman_corr(a, b)
@dataclass
class Effect(Node):
"""``Effect(Vector, target, adjust=[stage, age]) -> Scalar``.
Observational backdoor adjustment — residualise the Vector and the
target on the clinical confounders (one-hot ``stage`` + continuous
``age``), then take the (kind-specified) correlation of the residuals.
Only as good as the measured confounders.
"""
inner: Node
target: str
kind: str = "spearman"
def __post_init__(self) -> None:
self.ttype = TType.SCALAR
def children(self) -> list[Node]:
return [self.inner]
def repr_typed(self) -> str:
return f"Effect({self.inner.repr_typed()},{self.target},{self.kind})"
def execute(self, ctx: ExecContext) -> float:
v = self.inner.execute(ctx)
if not isinstance(v, pd.Series):
return 0.0
y = ctx.labels.get(self.target)
if y is None:
return 0.0
df = ctx.clinical.reindex(v.index)
# Confounder set: read from ctx, default = ("stage","age") so
# legacy MSI / TMB / HPV runs are byte-for-byte unchanged.
# Columns that aren't in clinical are silently skipped.
confounders = tuple(c for c in (ctx.confounders or ()) if c in df.columns)
# Treat age as continuous, everything else as categorical (one-hot
# with drop_first to avoid the dummy-variable trap; dummy_na=False
# so missing values fall via the dropna below).
cols: dict[str, np.ndarray] = {
"v": v.astype(float).values,
"y": np.asarray(y, dtype=float),
}
cat_specs: list[str] = []
for c in confounders:
if c == "age":
cols["age"] = pd.to_numeric(df["age"], errors="coerce").values
elif c == "stage":
# Legacy: cast to str (turns NaN into the string "nan"),
# then one-hot. Preserved to keep MSI / TMB / HPV runs
# byte-for-byte unchanged.
cols["stage"] = df["stage"].astype(str).values
cat_specs.append("stage")
else:
# New confounders (sex, race, is_oropharynx, …):
# preserve NaN so dropna drops rows with missing
# values rather than lumping them into a "nan" bucket.
cols[c] = df[c].astype("object").where(df[c].notna(), other=np.nan).values
cat_specs.append(c)
full = pd.DataFrame(cols).dropna()
if len(full) < 8:
return 0.0
block_arrays: list[np.ndarray] = [np.ones(len(full))]
if "age" in cols:
block_arrays.append(full["age"].astype(float).values.reshape(-1, 1))
for c in cat_specs:
d = pd.get_dummies(
full[c].astype(str), prefix=c, drop_first=True, dummy_na=False,
).astype(float)
if d.shape[1]:
block_arrays.append(d.values)
X = np.column_stack(block_arrays)
try:
beta_v, *_ = np.linalg.lstsq(X, full["v"].values, rcond=None)
beta_y, *_ = np.linalg.lstsq(X, full["y"].values, rcond=None)
except np.linalg.LinAlgError:
return 0.0
rv = full["v"].values - X @ beta_v
ry = full["y"].values - X @ beta_y
if self.kind == "pearson":
return _pearson_corr(rv, ry)
return _spearman_corr(rv, ry)
# ---------------------------------------------------------------------------
# Model nodes
# ---------------------------------------------------------------------------
@dataclass
class FitApply(Node):
"""``Fit(Vector, labels) -> Model`` then ``Apply(Model, Cohort) -> Vector``.
We fuse Fit + Apply into a single node so the grammar exposes a
Vector-typed transform that "trains and predicts in-place." Fitness
sees a normal Vector output and treats it the same as any other.
"""
inner: Node
target: str # "msi" | "tmb"
def __post_init__(self) -> None:
self.ttype = TType.VECTOR
def children(self) -> list[Node]:
return [self.inner]
def repr_typed(self) -> str:
return f"FitApply({self.inner.repr_typed()},{self.target})"
def execute(self, ctx: ExecContext) -> pd.Series:
# Step 1 — score the APPLY side (ctx). This is what the
# fitted model will be applied to and the result returned.
v_apply = self.inner.execute(ctx)
if not isinstance(v_apply, pd.Series):
return pd.Series(np.zeros(ctx.M.shape[0]), index=ctx.M.index)
# Step 2 — fit on TRAIN ctx if one is set, else fit and apply
# on the same ctx (legacy). Train-only fit is what evaluate_
# holdout uses to keep test labels off the fit; the pipeline's
# full-cohort `_make_full_ctx` deliberately has no labels, so
# the early-return below short-circuits to the raw inner.
fit_ctx = ctx.fit_ctx if ctx.fit_ctx is not None else ctx
if fit_ctx is ctx:
v_fit = v_apply
else:
v_fit_raw = self.inner.execute(fit_ctx)
if not isinstance(v_fit_raw, pd.Series):
return v_apply.astype(float)
v_fit = v_fit_raw
y_fit = fit_ctx.labels.get(self.target)
if y_fit is None or len(y_fit) != len(v_fit):
return v_apply.astype(float)
X_fit = v_fit.astype(float).values.reshape(-1, 1)
X_apply = v_apply.astype(float).values.reshape(-1, 1)
y_arr = np.asarray(y_fit)
finite_fit = np.isfinite(X_fit[:, 0]) & np.isfinite(
y_arr.astype(float),
)
if finite_fit.sum() < 8:
return v_apply.astype(float)
# Binary targets (MSI, HPV) → logistic regression on the 1-D
# score, producing per-patient probability. Both objectives
# share the binary path; TMB stays on continuous OLS.
if self.target in ("msi", "hpv"):
y_bin = (y_arr > 0).astype(int)
if len(np.unique(y_bin[finite_fit])) < 2:
return v_apply.astype(float)
try:
lr = LogisticRegression(max_iter=500)
lr.fit(X_fit[finite_fit], y_bin[finite_fit])
proba = lr.predict_proba(X_apply)[:, 1]
return pd.Series(proba, index=v_apply.index)
except Exception:
return v_apply.astype(float)
# Continuous TMB: simple OLS on the single score.
try:
X_fit_aug = np.column_stack([np.ones(len(X_fit)), X_fit[:, 0]])
X_apply_aug = np.column_stack(
[np.ones(len(X_apply)), X_apply[:, 0]],
)
beta, *_ = np.linalg.lstsq(
X_fit_aug[finite_fit],
y_arr.astype(float)[finite_fit],
rcond=None,
)
pred = X_apply_aug @ beta
return pd.Series(pred, index=v_apply.index)
except Exception:
return v_apply.astype(float)