#!/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()