ReMEDi / src /remedi /models.py
pranamanam's picture
Upload 62 files
3f98d52 verified
Raw History Blame Contribute Delete
5.97 kB
"""Small molecule-dose response heads, grouped ensembles, and cross-fitted residuals."""
from pathlib import Path
import joblib
import numpy as np
from sklearn.linear_model import Ridge
from sklearn.neural_network import MLPRegressor
from sklearn.preprocessing import StandardScaler
from sklearn.model_selection import GroupKFold
from .chemistry import encode, fingerprints, max_tanimoto
from .io import write_json, file_hash
def design_features(chemical, dose, context, mode="ridge"):
"""Build condition rows from molecular features, dose in uM, and control means.
chemical is (N, P), dose is (N,), and context is (N, D). Ridge
features include molecule-specific linear and quadratic log-dose terms.
"""
dose = np.asarray(dose, float)
if (dose <= 0).any() or not np.isfinite(dose).all():
raise ValueError("Concentrations must be positive micromolar values")
logdose = np.log(dose)[:, None]
parts = [chemical, logdose, context]
if mode == "ridge":
parts += [chemical*logdose, chemical*logdose**2]
return np.concatenate(parts, axis=1)
def _fit_head(x, y, weights, kind, alpha, hidden, seed):
scale = StandardScaler().fit(x, sample_weight=weights)
x = scale.transform(x)
if kind == "ridge":
head = Ridge(alpha=alpha).fit(x, y, sample_weight=weights)
elif kind == "mlp":
head = MLPRegressor(hidden_layer_sizes=(hidden, hidden), alpha=alpha,
max_iter=500, early_stopping=False, random_state=seed)
head.fit(x, y, sample_weight=weights)
else: raise ValueError("head must be ridge or mlp")
return scale, head
def _predict_head(head, x):
return np.asarray(head[1].predict(head[0].transform(x)))
def train(data, output, kind="ridge", alpha=10., hidden=128, ensemble=5, folds=3,
encoder="fingerprint", embeddings=None, seed=0):
"""Fit response heads and molecule-held-out residuals.
data.response contains treated-minus-control means, one row per condition.
Return the fitted model dictionary and save it with training metadata.
"""
if ensemble < 1: raise ValueError("Ensemble must contain at least one head")
mask = data.obs.split.to_numpy() == "train"
if not mask.any(): raise ValueError("No training conditions")
train_obs = data.obs.loc[mask].reset_index(drop=True)
chemical = encode(train_obs.smiles, encoder, embeddings)
x = design_features(chemical, train_obs.dose_um, data.control[mask], kind)
y = data.response[mask]
groups = train_obs.molecule_id.to_numpy()
molecules, counts = np.unique(groups, return_counts=True)
if len(molecules) < max(folds, 3): raise ValueError("Too few independent training molecules")
# Give each training molecule equal total weight across its conditions.
base_weights = np.array([1/counts[np.where(molecules == m)[0][0]] for m in groups])
base_weights *= len(groups)/base_weights.sum()
rng = np.random.default_rng(seed)
heads = []
for k in range(ensemble):
# Resample whole molecules so their doses and replicates stay together.
sampled = rng.choice(molecules, len(molecules), replace=True)
multiplicity = {m: np.sum(sampled == m) for m in molecules}
weights = base_weights*np.array([multiplicity[m] for m in groups])
heads.append(_fit_head(x, y, weights, kind, alpha, hidden, seed+k))
residuals = np.empty_like(y)
residual_distance = np.empty(len(y))
row_fp = fingerprints(train_obs.smiles)
# Estimate residuals on compounds excluded from each fitted fold.
for k, (fit, held) in enumerate(GroupKFold(n_splits=folds).split(x, y, groups)):
head = _fit_head(x[fit], y[fit], base_weights[fit], kind, alpha, hidden, seed+100+k)
residuals[held] = y[held]-_predict_head(head, x[held])
residual_distance[held] = 1-max_tanimoto(row_fp[held], np.unique(row_fp[fit], axis=0))
# Set endpoint-coordinate weights using training response variance.
weight = 1/np.maximum(np.var(y, axis=0), np.median(np.var(y, axis=0))*.05+1e-6)
weight /= weight.sum()
unique_smiles = train_obs.drop_duplicates("molecule_id").smiles.tolist()
train_fp = fingerprints(unique_smiles)
# Cross-fitted residual rows and identities remain available for stratification.
model = {"heads": heads, "kind": kind, "alpha": alpha, "hidden": hidden,
"encoder": encoder, "embeddings": str(embeddings) if embeddings else None,
"residuals": residuals, "residual_distance": residual_distance,
"residual_smiles": train_obs.smiles.to_numpy(),
"train_molecules": molecules.tolist(), "train_smiles": unique_smiles,
"train_fingerprints": train_fp, "weight": weight,
"feature_space_id": data.metadata.get("feature_space_id"), "study": data.metadata.get("study"), "seed": seed}
output = Path(output); output.mkdir(parents=True, exist_ok=True)
joblib.dump(model, output/"model.joblib")
write_json(output/"training.json", {"kind": kind, "alpha": alpha, "hidden": hidden,
"encoder": encoder, "ensemble": ensemble, "crossfit_folds": folds,
"training_molecules": len(molecules), "training_conditions": len(y), "seed": seed,
"feature_space_id": model["feature_space_id"], "model_sha256": file_hash(output/"model.joblib")})
return model
def predict(model, smiles, dose, context, embeddings=None):
"""Return response predictions with shape (ensemble heads, actions, features)."""
chemical = encode(smiles, model["encoder"], embeddings or model["embeddings"])
x = design_features(chemical, dose, np.asarray(context, float), model["kind"])
return np.stack([_predict_head(h, x) for h in model["heads"]])
def check_feature_space(model, data):
if model.get("feature_space_id") != data.metadata.get("feature_space_id"):
raise ValueError("Model and response data use different cell-feature spaces")