| |
| """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 |
|
|
| |
| rng = np.random.default_rng(42) |
| cond_features = [] |
| for cond in conditions: |
| |
| drug_emb = get_emb(cond["drug_name"]).numpy() |
|
|
| |
| pert_vec = cond["pert_vec"] |
|
|
| |
| drug_name = cond["drug_name"] |
|
|
| |
| target = cond.get("target", "") |
|
|
| |
| 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: |
| |
| 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) |
|
|
| |
| 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"]] |
|
|
| |
| scores = [] |
| for i, cond in enumerate(conditions): |
| other_emb = drug_embs[cond["drug_name"]] |
| |
| 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"] |
|
|
| |
| 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", "") |
|
|
| |
| 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 |
| |
| 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) |
|
|
| |
| 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}") |
|
|
| |
| 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() |
|
|