GID-Flow / PDGrapher /scripts /evaluate_drug_loco.py
Boom5426's picture
Upload GID-Flow project snapshot (deduped: code + key artifacts)
07fcdfe verified
Raw
History Blame Contribute Delete
16.3 kB
#!/usr/bin/env python3
"""Drug LOCO evaluation for CausalFlow-ID — batched version.
For each drug×dose condition i:
1. Get source (vehicle) cells for condition i
2. Score condition i against ALL N conditions (including itself)
using batched forward passes through the model
3. Check if condition i's true drug×dose is in top-K
Metrics:
- Recall@K: fraction where true condition is in top-K
- MRR: mean reciprocal rank
- Median rank
Usage:
python scripts/evaluate_drug_loco.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 \
--k 1 5 10 \
--output outputs/causal_flow_drug/loco_results.csv
"""
import argparse
import csv
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 scipy.spatial.distance import cosine
from gidflow.data.sciplex_dataset import Sciplex3Dataset, sciplex_collate_fn
from gidflow.models import CausalFlowGIDModel
def load_model(checkpoint_path: str, num_genes: int, device: torch.device) -> CausalFlowGIDModel:
"""Load trained CausalFlowGIDModel from checkpoint.
Handles both config-based construction (new checkpoints) and
state-dict inference (old checkpoints without config).
"""
ckpt = torch.load(checkpoint_path, map_location=device, weights_only=False)
sd = ckpt["model_state_dict"]
# Check if model was trained with drug conditioning
has_drug = "drug_gate" in sd and "drug_encoder.projection.0.weight" in sd
# Try to load config from checkpoint
model_cfg = ckpt.get("config", {}).get("model", {})
sc_cfg = ckpt.get("config", {}).get("dataset", {}).get("sciplex", {})
if not model_cfg:
# Infer architecture from state dict
print(" No config in checkpoint, inferring from state dict...")
# PopulationEncoder: source_encoder.mlp.6.weight → [encoder_output, encoder_hidden]
encoder_output = sd["source_encoder.mlp.6.weight"].shape[0]
encoder_hidden = sd["source_encoder.mlp.0.weight"].shape[0]
# Input to PopulationEncoder is 2*num_genes (source + target concatenated)
enc_input = sd["source_encoder.mlp.0.weight"].shape[1]
num_genes_inferred = enc_input // 2
# GapEncoder
gap_output = sd["gap_encoder.mlp.6.weight"].shape[0]
gap_hidden = sd["gap_encoder.mlp.0.weight"].shape[0]
# CausalPlanner per_gene_mlp
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
# Infer drug dimensions from flow_response shapes
pert_enc_in = sd["flow_response.pert_encoder.0.weight"].shape[1] # G + D_d
has_drug = "drug_gate" in sd
if has_drug and pert_enc_in > num_genes_inferred:
drug_emb_dim = pert_enc_in - num_genes_inferred
else:
drug_emb_dim = 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 embedding dim
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_cfg = {
"num_genes": num_genes_inferred,
"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,
"planner_topk": None,
"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,
"drug_gate_init": 0.1,
"lambda_flow": 1.0,
"lambda_cycle": 0.5,
"lambda_causal": 1.0,
"lambda_sparse": 0.01,
"lambda_recon": 1.0,
"lambda_drug_target": 1.0,
"lambda_drug_contrastive": 0.1,
"lambda_drug_dose": 0.01,
}
print(f" Inferred: num_genes={num_genes_inferred}, drug_emb_dim={drug_emb_dim}, use_drug={has_drug}")
# Override num_genes if needed (alignment padding)
model_cfg["num_genes"] = num_genes
model = CausalFlowGIDModel(
num_genes=num_genes,
encoder_hidden=model_cfg.get("encoder_hidden", 512),
encoder_output=model_cfg.get("encoder_output", 256),
gap_hidden=model_cfg.get("gap_hidden", 256),
gap_output=model_cfg.get("gap_output", 256),
causal_gene_emb_dim=model_cfg.get("causal_gene_emb_dim", 128),
causal_n_heads=model_cfg.get("causal_n_heads", 4),
causal_n_layers=model_cfg.get("causal_n_layers", 2),
planner_hidden=model_cfg.get("planner_hidden", 512),
planner_n_layers=model_cfg.get("planner_n_layers", 2),
planner_topk=model_cfg.get("planner_topk"),
flow_latent_dim=model_cfg.get("flow_latent_dim", 256),
flow_hidden_dim=model_cfg.get("flow_hidden_dim", 512),
flow_n_layers=model_cfg.get("flow_n_layers", 3),
flow_time_embed_dim=model_cfg.get("flow_time_embed_dim", 128),
flow_pert_emb_dim=model_cfg.get("flow_pert_emb_dim", 256),
n_layers=model_cfg.get("n_layers", 2),
use_cooccurrence=model_cfg.get("use_cooccurrence", True),
use_latent=model_cfg.get("use_latent", False),
drug_emb_dim=model_cfg.get("drug_emb_dim", 128),
use_drug_encoder=model_cfg.get("use_drug_encoder", has_drug),
use_drug_gene_bridge=model_cfg.get("use_drug_gene_bridge", has_drug),
drug_gate_init=model_cfg.get("drug_gate_init", 0.1),
lambda_flow=model_cfg.get("lambda_flow", 1.0),
lambda_cycle=model_cfg.get("lambda_cycle", 0.5),
lambda_causal=model_cfg.get("lambda_causal", 1.0),
lambda_sparse=model_cfg.get("lambda_sparse", 0.01),
lambda_recon=model_cfg.get("lambda_recon", 1.0),
lambda_drug_target=model_cfg.get("lambda_drug_target", 1.0),
lambda_drug_contrastive=model_cfg.get("lambda_drug_contrastive", 0.1),
lambda_drug_dose=model_cfg.get("lambda_drug_dose", 0.01),
).to(device)
model.load_state_dict(ckpt["model_state_dict"], strict=False)
model.eval()
print(f" Loaded checkpoint from epoch {ckpt.get('epoch', '?')}")
return model
def score_batch(
model: CausalFlowGIDModel,
source_cells_batch: torch.Tensor, # [B, Ns, G]
target_cells_batch: torch.Tensor, # [B, Nt, G]
pert_vec_batch: torch.Tensor, # [B, G]
drug_smiles_list: list, # list of str, length B
device: torch.device,
batch_size: int = 32,
) -> np.ndarray:
"""Score a batch of conditions using negative MSE between predicted and target.
Uses model's DrugEncoder to compute drug embeddings from SMILES.
Lower MSE = better match → higher score.
"""
model.eval()
B = source_cells_batch.size(0)
scores = []
# Pre-compute drug embeddings from SMILES using model's DrugEncoder
drug_embs = model.drug_encoder(drug_smiles_list).to(device) # [B, D_d]
for i in range(0, B, batch_size):
j = min(i + batch_size, B)
src = source_cells_batch[i:j].to(device)
tgt = target_cells_batch[i:j].to(device)
pert = pert_vec_batch[i:j].to(device)
drug = drug_embs[i:j]
dose = torch.zeros(j - i, device=device)
with torch.no_grad():
out = model(src, tgt, true_perturbation=pert, drug_smiles=drug_smiles_list[i:j], dose=dose)
pred_cells = out["pred_cells"] # [b, Np, G]
# Align to same population size for MSE
tgt_mean = tgt.mean(dim=1, keepdim=True) # [b, 1, G]
pred_mean = pred_cells.mean(dim=1, keepdim=True) # [b, 1, G]
# MSE between predicted and target mean expression (lower = better)
mse = ((pred_mean - tgt_mean) ** 2).mean(dim=(1, 2)) # [b]
# Convert to score: negative MSE (higher = better)
score = -mse
scores.extend(score.cpu().numpy().tolist())
return np.array(scores, dtype=np.float32)
def run_loco_batched(dataset, model, device, ks=(1, 5, 10), batch_size: int = 32):
"""Batched LOCO evaluation using DrugEncoder for drug embeddings."""
# Handle Subset wrapper
if hasattr(dataset, 'dataset'):
raw_ds = dataset.dataset
indices = dataset.indices
conditions = [raw_ds._conditions[i] for i in indices]
num_genes = raw_ds.num_genes
get_X = raw_ds._X
get_smiles = lambda name: raw_ds._get_drug_smiles(name)
else:
conditions = dataset._conditions
num_genes = dataset.num_genes
get_X = dataset._X
get_smiles = lambda name: dataset._get_drug_smiles(name)
n_cond = len(conditions)
rng = np.random.default_rng(42)
# Pre-build all condition metadata
print(" Building condition metadata...")
all_drug_names = []
all_doses = []
all_drug_smiles = []
all_pert_vecs = []
all_src_cells = []
all_tgt_cells = []
all_src_masks = []
all_tgt_masks = []
for cond in conditions:
drug_name = cond["drug_name"]
dose = cond["dose"]
drug_smiles = get_smiles(drug_name)
pert_vec = torch.from_numpy(cond["pert_vec"]).float()
ns = min(64, len(cond["vehicle_cell_idx"]))
nt = min(64, 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()
src_mask = torch.ones(ns)
tgt_mask = torch.ones(nt)
all_drug_names.append(drug_name)
all_doses.append(dose)
all_drug_smiles.append(drug_smiles)
all_pert_vecs.append(pert_vec)
all_src_cells.append(src)
all_tgt_cells.append(tgt)
all_src_masks.append(src_mask)
all_tgt_masks.append(tgt_mask)
# Stack into tensors
print(f" Stacking {n_cond} conditions...")
max_ns = max(c.size(0) for c in all_src_cells)
max_nt = max(c.size(0) for c in all_tgt_cells)
G = num_genes
src_batch = torch.zeros(n_cond, max_ns, G)
tgt_batch = torch.zeros(n_cond, max_nt, G)
src_mask_batch = torch.zeros(n_cond, max_ns)
tgt_mask_batch = torch.zeros(n_cond, max_nt)
pert_batch = torch.zeros(n_cond, G)
for i in range(n_cond):
ns = all_src_cells[i].size(0)
nt = all_tgt_cells[i].size(0)
src_batch[i, :ns] = all_src_cells[i]
tgt_batch[i, :nt] = all_tgt_cells[i]
src_mask_batch[i, :ns] = all_src_masks[i]
tgt_mask_batch[i, :nt] = all_tgt_masks[i]
pert_batch[i] = all_pert_vecs[i]
# Score all conditions against each source
print(f" Scoring {n_cond} test conditions against {n_cond} candidates...")
all_true_ranks = []
all_recall = {k: [] for k in ks}
for test_idx in range(n_cond):
# Source cells for this test condition
test_src = src_batch[test_idx:test_idx+1].expand(n_cond, -1, -1) # [N, Ns, G]
test_tgt = tgt_batch[test_idx:test_idx+1].expand(n_cond, -1, -1) # [N, Nt, G]
test_src_mask = src_mask_batch[test_idx:test_idx+1].expand(n_cond, -1)
test_tgt_mask = tgt_mask_batch[test_idx:test_idx+1].expand(n_cond, -1)
test_pert = pert_batch[test_idx:test_idx+1].expand(n_cond, -1) # [N, G]
test_smiles = all_drug_smiles # [N] list of SMILES
# Score all N conditions (including true one) using DrugEncoder
scores = score_batch(
model, test_src, test_tgt, test_pert, test_smiles, device,
batch_size=batch_size,
)
# Rank by score descending
ranked_indices = np.argsort(-scores)
true_rank = int(np.where(ranked_indices == test_idx)[0][0]) + 1
all_true_ranks.append(true_rank)
for k in ks:
all_recall[k].append(1 if true_rank <= k else 0)
if (test_idx + 1) % 50 == 0 or test_idx == 0:
print(f"\r LOCO [{test_idx+1}/{n_cond}] "
f"drug={all_drug_names[test_idx][:20]:20s} "
f"rank={true_rank:4d}/{n_cond}", end="", flush=True)
print(f"\r LOCO [{n_cond}/{n_cond}] complete")
# Aggregate metrics
metrics = {
"n_conditions": n_cond,
"median_rank": float(np.median(all_true_ranks)),
"mean_rank": float(np.mean(all_true_ranks)),
}
for k in ks:
metrics[f"recall@{k}"] = float(np.mean(all_recall[k]))
return metrics
def main():
parser = argparse.ArgumentParser(description="Drug LOCO evaluation")
parser.add_argument("--checkpoint", required=True, help="Model checkpoint path")
parser.add_argument("--preprocessed", required=True, help="SciPlex3 preprocessed .pt")
parser.add_argument("--gene-map", default="", help="ENSEMBL→symbol JSON mapping")
parser.add_argument("--smiles", default="", help="ChEMBL SMILES CSV")
parser.add_argument("--target-num-genes", type=int, default=None)
parser.add_argument("--k", type=int, nargs="+", default=[1, 5, 10, 20, 50])
parser.add_argument("--max-conditions", type=int, default=None)
parser.add_argument("--output", default="outputs/causal_flow_drug/loco_results.csv")
parser.add_argument("--device", default="cuda")
parser.add_argument("--batch-size", type=int, default=32)
args = parser.parse_args()
device = torch.device(args.device if torch.cuda.is_available() else "cpu")
print(f"Device: {device}")
# Load dataset
print("\n=== Loading SciPlex3 dataset ===")
ds = Sciplex3Dataset(
h5ad_path="",
n_hvg=2000,
preprocessed_path=args.preprocessed,
target_num_genes=args.target_num_genes,
drug_smiles_csv=args.smiles,
)
print(f" Dataset: {len(ds)} conditions, {ds.num_genes} genes")
if args.max_conditions:
from torch.utils.data import Subset
indices = list(range(min(args.max_conditions, len(ds))))
ds_raw = ds
ds = Subset(ds, indices)
num_genes = ds_raw.num_genes
print(f" Limited to {len(ds)} conditions")
else:
num_genes = ds.num_genes
# Load model
print("\n=== Loading model ===")
model = load_model(args.checkpoint, num_genes, device)
total_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
print(f" Parameters: {total_params:,}")
# Run LOCO
print(f"\n=== Running LOCO (K={args.k}, batch_size={args.batch_size}) ===")
metrics = run_loco_batched(ds, model, device, ks=tuple(args.k), batch_size=args.batch_size)
# Print summary
print("\n=== LOCO Results ===")
for k in args.k:
key = f"recall@{k}"
if key in metrics:
print(f" Recall@{k:2d}: {metrics[key]:.4f}")
print(f" Median rank: {metrics['median_rank']:.1f}")
print(f" Mean rank: {metrics['mean_rank']:.1f}")
print(f" Random baseline Recall@1: {1.0/metrics['n_conditions']:.4f}")
# Save metrics
os.makedirs(os.path.dirname(args.output) or ".", exist_ok=True)
metrics_path = args.output.replace(".csv", "_metrics.json")
with open(metrics_path, "w") as f:
json.dump(metrics, f, indent=2)
print(f"\nMetrics saved to {metrics_path}")
if __name__ == "__main__":
main()