| |
| """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"] |
|
|
| |
| has_drug = "drug_gate" in sd and "drug_encoder.projection.0.weight" in sd |
|
|
| |
| model_cfg = ckpt.get("config", {}).get("model", {}) |
| sc_cfg = ckpt.get("config", {}).get("dataset", {}).get("sciplex", {}) |
|
|
| if not model_cfg: |
| |
| print(" No config in checkpoint, inferring from state dict...") |
|
|
| |
| encoder_output = sd["source_encoder.mlp.6.weight"].shape[0] |
| encoder_hidden = sd["source_encoder.mlp.0.weight"].shape[0] |
| |
| enc_input = sd["source_encoder.mlp.0.weight"].shape[1] |
| num_genes_inferred = enc_input // 2 |
|
|
| |
| 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] |
| 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_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}") |
|
|
| |
| 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, |
| target_cells_batch: torch.Tensor, |
| pert_vec_batch: torch.Tensor, |
| drug_smiles_list: list, |
| 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 = [] |
|
|
| |
| drug_embs = model.drug_encoder(drug_smiles_list).to(device) |
|
|
| 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"] |
|
|
| |
| tgt_mean = tgt.mean(dim=1, keepdim=True) |
| pred_mean = pred_cells.mean(dim=1, keepdim=True) |
|
|
| |
| mse = ((pred_mean - tgt_mean) ** 2).mean(dim=(1, 2)) |
| |
| 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.""" |
| |
| 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) |
|
|
| |
| 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) |
|
|
| |
| 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] |
|
|
| |
| 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): |
| |
| test_src = src_batch[test_idx:test_idx+1].expand(n_cond, -1, -1) |
| test_tgt = tgt_batch[test_idx:test_idx+1].expand(n_cond, -1, -1) |
| 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) |
| test_smiles = all_drug_smiles |
|
|
| |
| scores = score_batch( |
| model, test_src, test_tgt, test_pert, test_smiles, device, |
| batch_size=batch_size, |
| ) |
|
|
| |
| 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") |
|
|
| |
| 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}") |
|
|
| |
| 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 |
|
|
| |
| 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:,}") |
|
|
| |
| 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("\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}") |
|
|
| |
| 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() |
|
|