#!/usr/bin/env python3 """Drug recommendation baselines for CausalFlow-ID comparison. Implements multiple baseline methods for drug recommendation: 1. Random: uniform random scores 2. GenePrior: cosine similarity between drug target genes and causal genes 3. DrugEmb: cosine similarity between Morgan fingerprint embeddings 4. DrugPrior: cosine similarity between drug embeddings predicted by the model 5. PerfectOracle: uses ground truth drug-disease association Usage: python scripts/drug_baselines.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/baseline_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 def load_baseline_data(preprocessed_path, gene_map_path, smiles_path, target_num_genes): """Load dataset and extract baseline features.""" ds = Sciplex3Dataset( h5ad_path="", n_hvg=2000, preprocessed_path=preprocessed_path, target_num_genes=target_num_genes, drug_smiles_csv=smiles_path, ) conditions = ds._conditions get_emb = ds._get_drug_embedding get_X = ds._X # Extract features for each condition rng = np.random.default_rng(42) cond_features = [] for cond in conditions: # Drug embedding (Morgan fingerprint) drug_emb = get_emb(cond["drug_name"]).numpy() # Perturbation vector (target genes) pert_vec = cond["pert_vec"] # Drug name drug_name = cond["drug_name"] # Target category target = cond.get("target", "") # Dose dose = cond["dose"] cond_features.append({ "drug_name": drug_name, "target": target, "dose": dose, "drug_emb": drug_emb, "pert_vec": pert_vec, }) return cond_features, ds.num_genes def baseline_random(conditions, n_test=100, seed=42): """Random baseline: uniform random scores.""" rng = np.random.default_rng(seed) test_indices = rng.choice(len(conditions), size=min(n_test, len(conditions)), replace=False) all_ranks = [] all_recall = {k: [] for k in [1, 5, 10, 20, 50]} for test_idx in test_indices: # Random scores for all conditions scores = rng.random(len(conditions)) ranked_indices = np.argsort(-scores) true_rank = int(np.where(ranked_indices == test_idx)[0][0]) + 1 all_ranks.append(true_rank) for k in all_recall: all_recall[k].append(1 if true_rank <= k else 0) metrics = { "median_rank": float(np.median(all_ranks)), "mean_rank": float(np.mean(all_ranks)), } for k in all_recall: metrics[f"recall@{k}"] = float(np.mean(all_recall[k])) return metrics def baseline_drug_emb_similarity(conditions, n_test=100, seed=42): """DrugEmb baseline: rank by cosine similarity of Morgan fingerprints. For each test condition, score all conditions by cosine similarity between their drug embeddings (higher similarity = higher rank). """ rng = np.random.default_rng(seed) test_indices = rng.choice(len(conditions), size=min(n_test, len(conditions)), replace=False) # Pre-compute drug embeddings drug_embs = {} for cond in conditions: name = cond["drug_name"] if name not in drug_embs: drug_embs[name] = cond["drug_emb"] all_ranks = [] all_recall = {k: [] for k in [1, 5, 10, 20, 50]} for test_idx in test_indices: test_cond = conditions[test_idx] test_emb = drug_embs[test_cond["drug_name"]] # Score by cosine similarity scores = [] for i, cond in enumerate(conditions): other_emb = drug_embs[cond["drug_name"]] # Cosine similarity sim = np.dot(test_emb, other_emb) / (np.linalg.norm(test_emb) * np.linalg.norm(other_emb) + 1e-8) scores.append(sim) scores = np.array(scores) ranked_indices = np.argsort(-scores) true_rank = int(np.where(ranked_indices == test_idx)[0][0]) + 1 all_ranks.append(true_rank) for k in all_recall: all_recall[k].append(1 if true_rank <= k else 0) metrics = { "median_rank": float(np.median(all_ranks)), "mean_rank": float(np.mean(all_ranks)), } for k in all_recall: metrics[f"recall@{k}"] = float(np.mean(all_recall[k])) return metrics def baseline_gene_prior(conditions, n_test=100, seed=42): """GenePrior baseline: cosine similarity between drug target genes and known target genes. For each drug, we have a pert_vec indicating target genes. Score = cosine similarity between test drug's target genes and candidate's target genes. """ rng = np.random.default_rng(seed) test_indices = rng.choice(len(conditions), size=min(n_test, len(conditions)), replace=False) all_ranks = [] all_recall = {k: [] for k in [1, 5, 10, 20, 50]} for test_idx in test_indices: test_cond = conditions[test_idx] test_vec = test_cond["pert_vec"] # Score by cosine similarity of pert_vecs scores = [] for i, cond in enumerate(conditions): other_vec = cond["pert_vec"] sim = np.dot(test_vec, other_vec) / (np.linalg.norm(test_vec) * np.linalg.norm(other_vec) + 1e-8) scores.append(sim) scores = np.array(scores) ranked_indices = np.argsort(-scores) true_rank = int(np.where(ranked_indices == test_idx)[0][0]) + 1 all_ranks.append(true_rank) for k in all_recall: all_recall[k].append(1 if true_rank <= k else 0) metrics = { "median_rank": float(np.median(all_ranks)), "mean_rank": float(np.mean(all_ranks)), } for k in all_recall: metrics[f"recall@{k}"] = float(np.mean(all_recall[k])) return metrics def baseline_drug_class_match(conditions, n_test=100, seed=42): """Drug class baseline: rank by target category match. Score = 1 if same target category, 0 otherwise. Breaks ties randomly. """ rng = np.random.default_rng(seed) test_indices = rng.choice(len(conditions), size=min(n_test, len(conditions)), replace=False) all_ranks = [] all_recall = {k: [] for k in [1, 5, 10, 20, 50]} for test_idx in test_indices: test_cond = conditions[test_idx] test_target = test_cond.get("target", "") # Score by target category match scores = [] for i, cond in enumerate(conditions): other_target = cond.get("target", "") score = 1.0 if test_target == other_target and test_target != "" else 0.0 # Add small random noise to break ties score += rng.random() * 0.01 scores.append(score) scores = np.array(scores) ranked_indices = np.argsort(-scores) true_rank = int(np.where(ranked_indices == test_idx)[0][0]) + 1 all_ranks.append(true_rank) for k in all_recall: all_recall[k].append(1 if true_rank <= k else 0) metrics = { "median_rank": float(np.median(all_ranks)), "mean_rank": float(np.mean(all_ranks)), } for k in all_recall: metrics[f"recall@{k}"] = float(np.mean(all_recall[k])) return metrics def main(): parser = argparse.ArgumentParser(description="Drug recommendation baselines") parser.add_argument("--checkpoint", default="", help="Model checkpoint (not used for baselines)") 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("--n-test", type=int, default=100) parser.add_argument("--output", default="outputs/causal_flow_drug/baseline_results.json") parser.add_argument("--device", default="cuda") args = parser.parse_args() print("=== Loading baseline data ===") conditions, num_genes = load_baseline_data( args.preprocessed, args.gene_map, args.smiles, args.target_num_genes ) print(f" {len(conditions)} conditions, {num_genes} genes") results = {} print(f"\n=== Running baselines (n_test={args.n_test}) ===") print("\n1. Random baseline...") results["random"] = baseline_random(conditions, n_test=args.n_test) print("\n2. DrugEmb (Morgan fingerprint cosine similarity)...") results["drug_emb"] = baseline_drug_emb_similarity(conditions, n_test=args.n_test) print("\n3. GenePrior (target gene cosine similarity)...") results["gene_prior"] = baseline_gene_prior(conditions, n_test=args.n_test) print("\n4. DrugClassMatch (same target category)...") results["drug_class_match"] = baseline_drug_class_match(conditions, n_test=args.n_test) # Summary print("\n=== Baseline Summary ===") print(f" Random: Recall@1={results['random']['recall@1']:.4f}, MedRank={results['random']['median_rank']:.1f}") print(f" DrugEmb: Recall@1={results['drug_emb']['recall@1']:.4f}, MedRank={results['drug_emb']['median_rank']:.1f}") print(f" GenePrior: Recall@1={results['gene_prior']['recall@1']:.4f}, MedRank={results['gene_prior']['median_rank']:.1f}") print(f" DrugClassMatch: Recall@1={results['drug_class_match']['recall@1']:.4f}, MedRank={results['drug_class_match']['median_rank']:.1f}") # 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"\nBaseline results saved to {args.output}") if __name__ == "__main__": main()