GID-Flow / PDGrapher /scripts /drug_baselines.py
Boom5426's picture
Upload GID-Flow project snapshot (deduped: code + key artifacts)
07fcdfe verified
Raw
History Blame Contribute Delete
10.2 kB
#!/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()