File size: 5,967 Bytes
3f98d52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
"""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")