File size: 13,660 Bytes
07fcdfe | 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 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 | #!/usr/bin/env python3
"""Drug recommendation evaluation — multiple metrics.
Given the limitations of exact LOCO (dose not encoded in drug_emb),
we evaluate:
1. Drug class retrieval: given a condition, rank drugs by target class match
- Can the model distinguish drugs with same vs different protein targets?
2. Drug embedding nearest-neighbor: use drug embeddings directly
- Do drugs with similar Morgan fingerprints get similar model scores?
3. Gate ablation: measure how drug gate value affects discrimination
- If gate → 0: model ignores drug (CRISPRi-only baseline)
- If gate → 1: model relies entirely on drug
4. Per-drug score consistency: for same drug across doses, score variance
- Low variance = model gives consistent predictions (GOOD)
Usage:
python scripts/evaluate_drug_recommendation.py \
--checkpoint outputs/causal_flow_drug/best_checkpoint.pt \
--preprocessed data/processed/sciplex3_k562_24h.pt \
--gene-map data/processed/sciplex3_k562_24h_gene_map.json \
--smiles data/chembl_smiles.csv \
--target-num-genes 2085 \
--output outputs/causal_flow_drug/eval_results.json
"""
import argparse
import json
import os
import sys
import warnings
from collections import defaultdict
warnings.filterwarnings("ignore")
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "src"))
import numpy as np
import torch
from gidflow.data.sciplex_dataset import Sciplex3Dataset
from gidflow.models import CausalFlowGIDModel
def load_model(checkpoint_path: str, num_genes: int, device: torch.device) -> CausalFlowGIDModel:
ckpt = torch.load(checkpoint_path, map_location=device, weights_only=False)
sd = ckpt["model_state_dict"]
has_drug = "drug_gate" in sd
encoder_output = sd["source_encoder.mlp.6.weight"].shape[0]
encoder_hidden = sd["source_encoder.mlp.0.weight"].shape[0]
gap_output = sd["gap_encoder.mlp.6.weight"].shape[0]
gap_hidden = sd["gap_encoder.mlp.0.weight"].shape[0]
planner_hidden = sd["causal_planner.per_gene_mlp.0.weight"].shape[0]
src_layers = len([k for k in sd if "source_encoder" in k and ".mlp." in k and "weight" in k])
n_layers = (src_layers - 1) // 3 + 1
pert_enc_in = sd["flow_response.pert_encoder.0.weight"].shape[1]
drug_emb_dim = pert_enc_in - num_genes if has_drug else 0
flow_pert_emb_dim = sd["flow_response.pert_encoder.4.weight"].shape[0]
flow_hidden_dim = sd["flow_response.pert_encoder.0.weight"].shape[0]
gene_emb_shape = sd.get("causal_planner.causal_estimator.gene_embedding.weight", torch.zeros(1)).shape
causal_gene_emb_dim = gene_emb_shape[1] if len(gene_emb_shape) > 1 else 128
model = CausalFlowGIDModel(
num_genes=num_genes,
encoder_hidden=encoder_hidden,
encoder_output=encoder_output,
gap_hidden=gap_hidden,
gap_output=gap_output,
causal_gene_emb_dim=causal_gene_emb_dim,
causal_n_heads=4,
causal_n_layers=2,
planner_hidden=planner_hidden,
planner_n_layers=n_layers,
flow_latent_dim=256,
flow_hidden_dim=flow_hidden_dim,
flow_n_layers=3,
flow_time_embed_dim=128,
flow_pert_emb_dim=flow_pert_emb_dim,
n_layers=n_layers,
use_cooccurrence=True,
use_latent=False,
drug_emb_dim=drug_emb_dim,
use_drug_encoder=has_drug,
use_drug_gene_bridge=has_drug,
).to(device)
model.load_state_dict(sd, strict=False)
model.eval()
gate_val = torch.sigmoid(model.drug_gate).item()
print(f" Loaded: num_genes={num_genes}, drug_emb_dim={drug_emb_dim}, drug_gate={gate_val:.4f}")
return model
@torch.no_grad()
def compute_flow_score(model, src, tgt, pert, smiles, device) -> float:
"""Compute -flow_loss for a single condition (higher = better).
Uses model's DrugEncoder to encode SMILES dynamically.
"""
src = src.unsqueeze(0).to(device)
tgt = tgt.unsqueeze(0).to(device)
pert = pert.unsqueeze(0).to(device)
src_mask = torch.ones(src.size(1), device=device).unsqueeze(0)
tgt_mask = torch.ones(tgt.size(1), device=device).unsqueeze(0)
out = model(src, tgt, src_mask, tgt_mask, true_perturbation=pert, drug_smiles=[smiles])
return -out["flow_loss"].item()
def evaluate_drug_class_retrieval(model, conditions, get_X, get_smiles, device, n_test=100):
"""Test: can the model rank same-target-class drugs higher than different-class drugs?
For each test condition:
1. Score against all candidates (same drug, all doses + other drugs)
2. Check if same-target-class drugs rank higher
"""
rng = np.random.default_rng(42)
test_indices = rng.choice(len(conditions), size=min(n_test, len(conditions)), replace=False)
# Group conditions by target category
target_groups = defaultdict(list)
for i, cond in enumerate(conditions):
target_groups[cond.get("target", "")].append(i)
print(f" Target categories: {len(target_groups)}")
for t, idxs in sorted(target_groups.items(), key=lambda x: -len(x[1]))[:5]:
print(f" {t or '(none)'}: {len(idxs)} conditions")
within_class_ranks = []
between_class_ranks = []
for test_idx in test_indices:
test_cond = conditions[test_idx]
test_target = test_cond.get("target", "")
# Get test source cells
ns = min(32, len(test_cond["vehicle_cell_idx"]))
nt = min(32, len(test_cond["drug_cell_idx"]))
src_idx = rng.choice(test_cond["vehicle_cell_idx"], size=ns, replace=False)
tgt_idx = rng.choice(test_cond["drug_cell_idx"], size=nt, replace=False)
src = torch.from_numpy(get_X[src_idx]).float()
tgt = torch.from_numpy(get_X[tgt_idx]).float()
# Score all candidates
scores = []
for i, cond in enumerate(conditions):
smiles = get_smiles(cond["drug_name"])
pert = torch.from_numpy(cond["pert_vec"]).float()
score = compute_flow_score(model, src, tgt, pert, smiles, device)
scores.append((i, score, cond.get("target", "")))
scores.sort(key=lambda x: -x[1]) # descending score
# Find rank of same-target-class drugs
same_class_indices = set(target_groups.get(test_target, []))
same_class_ranks = [rank + 1 for rank, (i, _, _) in enumerate(scores) if i in same_class_indices]
diff_class_ranks = [rank + 1 for rank, (i, _, _) in enumerate(scores) if i not in same_class_indices]
if same_class_ranks:
within_class_ranks.append(np.median(same_class_ranks))
if diff_class_ranks:
between_class_ranks.append(np.median(diff_class_ranks))
result = {
"within_class_median_rank": float(np.median(within_class_ranks)) if within_class_ranks else 0,
"between_class_median_rank": float(np.median(between_class_ranks)) if between_class_ranks else 0,
"n_test": len(test_indices),
"n_target_classes": len(target_groups),
}
print(f" Within-class median rank: {result['within_class_median_rank']:.1f}")
print(f" Between-class median rank: {result['between_class_median_rank']:.1f}")
if result["within_class_median_rank"] > 0:
improvement = result["between_class_median_rank"] - result["within_class_median_rank"]
print(f" Improvement (between - within): {improvement:.1f}")
result["rank_improvement"] = improvement
return result
def evaluate_dose_consistency(model, conditions, get_X, get_smiles, device):
"""Test: for same drug, different doses, are scores consistent?
Low variance = model gives similar predictions for same drug (expected,
since drug_emb doesn't encode dose).
"""
rng = np.random.default_rng(42)
# Group by drug
drug_conditions = defaultdict(list)
for i, cond in enumerate(conditions):
drug_conditions[cond["drug_name"]].append(i)
# For drugs with ≥3 conditions, compute score variance across doses
variances = []
drug_names = []
for drug, indices in drug_conditions.items():
if len(indices) < 3:
continue
# Use first condition as reference source
ref_cond = conditions[indices[0]]
ns = min(32, len(ref_cond["vehicle_cell_idx"]))
nt = min(32, len(ref_cond["drug_cell_idx"]))
src_idx = rng.choice(ref_cond["vehicle_cell_idx"], size=ns, replace=False)
src = torch.from_numpy(get_X[src_idx]).float()
scores = []
for idx in indices:
cond = conditions[idx]
tgt_idx = rng.choice(cond["drug_cell_idx"], size=nt, replace=False)
tgt = torch.from_numpy(get_X[tgt_idx]).float()
smiles = get_smiles(cond["drug_name"])
pert = torch.from_numpy(cond["pert_vec"]).float()
score = compute_flow_score(model, src, tgt, pert, smiles, device)
scores.append(score)
variances.append(np.var(scores))
drug_names.append(drug)
result = {
"n_drugs_tested": len(variances),
"mean_score_variance": float(np.mean(variances)) if variances else 0,
"median_score_variance": float(np.median(variances)) if variances else 0,
}
print(f" Drugs tested (≥3 conditions): {len(variances)}")
print(f" Mean score variance across doses: {result['mean_score_variance']:.6f}")
print(f" → Low variance = model gives consistent predictions for same drug")
return result
def evaluate_zero_vs_nonzero_embeddings(model, conditions, get_X, get_smiles, device):
"""Test: do conditions with zero vs non-zero SMILES get different scores?"""
rng = np.random.default_rng(42)
zero_scores = []
nonzero_scores = []
for i, cond in enumerate(conditions):
smiles = get_smiles(cond["drug_name"])
if not smiles:
group = zero_scores
else:
group = nonzero_scores
ns = min(32, len(cond["vehicle_cell_idx"]))
nt = min(32, len(cond["drug_cell_idx"]))
src_idx = rng.choice(cond["vehicle_cell_idx"], size=ns, replace=False)
tgt_idx = rng.choice(cond["drug_cell_idx"], size=nt, replace=False)
src = torch.from_numpy(get_X[src_idx]).float()
tgt = torch.from_numpy(get_X[tgt_idx]).float()
pert = torch.from_numpy(cond["pert_vec"]).float()
score = compute_flow_score(model, src, tgt, pert, smiles, device)
group.append(score)
result = {
"zero_smiles_mean_score": float(np.mean(zero_scores)) if zero_scores else 0,
"nonzero_smiles_mean_score": float(np.mean(nonzero_scores)) if nonzero_scores else 0,
"zero_smiles_count": len(zero_scores),
"nonzero_smiles_count": len(nonzero_scores),
}
print(f" Zero-SMILES conditions: {len(zero_scores)}")
print(f" Non-zero-SMILES conditions: {len(nonzero_scores)}")
print(f" Mean score (zero SMILES): {result['zero_smiles_mean_score']:.4f}")
print(f" Mean score (nonzero SMILES): {result['nonzero_smiles_mean_score']:.4f}")
if zero_scores and nonzero_scores:
diff = result['nonzero_smiles_mean_score'] - result['zero_smiles_mean_score']
print(f" Difference (nonzero - zero): {diff:.4f}")
result["score_difference"] = diff
return result
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--checkpoint", required=True)
parser.add_argument("--preprocessed", required=True)
parser.add_argument("--gene-map", default="")
parser.add_argument("--smiles", default="")
parser.add_argument("--target-num-genes", type=int, default=None)
parser.add_argument("--output", default="outputs/causal_flow_drug/eval_results.json")
parser.add_argument("--device", default="cuda")
args = parser.parse_args()
device = torch.device(args.device if torch.cuda.is_available() else "cpu")
print("=== Loading dataset ===")
ds = Sciplex3Dataset(
h5ad_path="",
n_hvg=2000,
preprocessed_path=args.preprocessed,
target_num_genes=args.target_num_genes,
drug_smiles_csv=args.smiles,
)
conditions = ds._conditions
get_smiles = lambda name: ds._get_drug_smiles(name)
get_X = ds._X
print(f" {len(conditions)} conditions, {ds.num_genes} genes")
print("\n=== Loading model ===")
model = load_model(args.checkpoint, ds.num_genes, device)
results = {}
print("\n=== 1. Drug class retrieval ===")
results["drug_class_retrieval"] = evaluate_drug_class_retrieval(
model, conditions, get_X, get_smiles, device, n_test=100
)
print("\n=== 2. Dose consistency ===")
results["dose_consistency"] = evaluate_dose_consistency(
model, conditions, get_X, get_smiles, device
)
print("\n=== 3. Zero vs non-zero embedding ===")
results["embedding_ablation"] = evaluate_zero_vs_nonzero_embeddings(
model, conditions, get_X, get_smiles, device
)
# Summary
print("\n=== Summary ===")
print(f" Drug gate: {torch.sigmoid(model.drug_gate).item():.4f}")
if "rank_improvement" in results.get("drug_class_retrieval", {}):
print(f" Class retrieval improvement: {results['drug_class_retrieval']['rank_improvement']:.1f}")
if "score_difference" in results.get("embedding_ablation", {}):
print(f" Embedding score difference: {results['embedding_ablation']['score_difference']:.4f}")
# Save results
os.makedirs(os.path.dirname(args.output) or ".", exist_ok=True)
with open(args.output, "w") as f:
json.dump(results, f, indent=2)
print(f"\nResults saved to {args.output}")
if __name__ == "__main__":
main()
|