GID-Flow / PDGrapher /scripts /diagnose_alignment.py
Boom5426's picture
Upload GID-Flow project snapshot (deduped: code + key artifacts)
07fcdfe verified
Raw
History Blame Contribute Delete
11.9 kB
"""Diagnostic probe for DrugRank-Flow weak retrieval.
Dumps hard statistics to results/diagnostics/alignment_diag.json so we can
reason about WHY gap_emb <-> drug_proj alignment is only ~2-3x chance:
1. Transcriptional signal strength by dose: ||target_mean - source_mean||
(is 10nM basically vehicle-noise?)
2. gap_emb collapse: pairwise cosine spread + per-dim std over many conditions
3. drug_proj separation: pairwise cosine among the 189 gallery drugs
4. Same-drug cross-dose consistency of gap_emb (tests dose-confounding)
5. Retrieval Hit@10 / median-rank stratified by dose
6. gap_emb vs drug_proj: for a batch, cosine(true pair) vs cosine(best wrong)
Run:
python scripts/diagnose_alignment.py \
--checkpoint outputs/drug_rank/phase2_best.pt \
--config configs/drug_rank_phase2.yaml
"""
from __future__ import annotations
import argparse
import json
import logging
import os
import sys
from pathlib import Path
import numpy as np
import torch
import torch.nn.functional as F
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
from gidflow.data.sciplex_dataset import Sciplex3Dataset
from gidflow.models.population_encoder import PopulationEncoder
from gidflow.models.gap_encoder import GapEncoder
from gidflow.models.drug_encoder import DrugEncoder
from gidflow.models.drug_gene_bridge import DrugGeneBridge
logging.basicConfig(level=logging.INFO, format="%(asctime)s | %(message)s", datefmt="%H:%M:%S")
log = logging.getLogger(__name__)
ANNOTATION_DIR = Path("/data/boom/ICLR/data/annotation")
CELL_LINE_MAP = {"A549": 1, "K562": 2, "MCF7": 3}
def build_and_load(cfg, ckpt_path, device):
m = cfg["model"]
num_proteins = len(json.load(open(ANNOTATION_DIR / "protein_target_vocab.json")))
source_enc = PopulationEncoder(num_genes=m["num_genes"], hidden_dim=m["encoder_hidden"], output_dim=m["encoder_output"]).to(device)
target_enc = PopulationEncoder(num_genes=m["num_genes"], hidden_dim=m["encoder_hidden"], output_dim=m["encoder_output"]).to(device)
gap_enc = GapEncoder(input_dim=m["encoder_output"], hidden_dim=m["gap_hidden"], output_dim=m["gap_output"],
proj_dim=m["gap_proj_dim"], num_cell_lines=m["num_cell_lines"], num_genes=m["num_genes"]).to(device)
drug_enc = DrugEncoder(encoding=m.get("drug_encoder", "morgan"), emb_dim=m["drug_emb_dim"], freeze=True).to(device)
bridge = DrugGeneBridge(num_proteins=num_proteins, drug_emb_dim=m["drug_emb_dim"], hidden_dim=m["bridge_hidden_dim"],
proj_dim=m["bridge_proj_dim"], protein_emb_dim=m["bridge_protein_emb_dim"]).to(device)
ck = torch.load(ckpt_path, map_location=device, weights_only=False)
source_enc.load_state_dict(ck["source_enc"]); target_enc.load_state_dict(ck["target_enc"])
gap_enc.load_state_dict(ck["gap_enc"]); drug_enc.load_state_dict(ck["drug_enc"]); bridge.load_state_dict(ck["bridge"])
for mod in (source_enc, target_enc, gap_enc, drug_enc, bridge):
mod.eval()
return source_enc, target_enc, gap_enc, drug_enc, bridge
@torch.no_grad()
def encode_conditions(dataset, conds, models, device, pass_cell_line):
source_enc, target_enc, gap_enc, drug_enc, bridge = models
X = dataset._X
gaps, dnames, doses, cls, deltas = [], [], [], [], []
for cond in conds:
veh = np.asarray(cond["vehicle_cell_idx"]); drg = np.asarray(cond["drug_cell_idx"])
if len(veh) == 0 or len(drg) == 0:
continue
ns = min(64, len(veh)); nt = min(64, len(drg))
rng = np.random.default_rng(0)
s = rng.choice(veh, ns, replace=False); t = rng.choice(drg, nt, replace=False)
src = torch.from_numpy(np.asarray(X[s], np.float32))[None].to(device)
tgt = torch.from_numpy(np.asarray(X[t], np.float32))[None].to(device)
sm = torch.ones(1, ns, dtype=torch.bool, device=device)
tm = torch.ones(1, nt, dtype=torch.bool, device=device)
z_s = source_enc(src, sm); z_t = target_enc(tgt, tm)
if pass_cell_line:
cl = torch.tensor([CELL_LINE_MAP.get(cond["cell_line"], 0)], device=device)
g = gap_enc(z_s, z_t, cell_line_ids=cl)["gap_emb"]
else:
g = gap_enc(z_s, z_t)["gap_emb"]
gaps.append(g[0].cpu().numpy())
dnames.append(cond["drug_name"]); doses.append(float(cond["dose"])); cls.append(cond["cell_line"])
deltas.append(float(np.linalg.norm(np.asarray(X[t]).mean(0) - np.asarray(X[s]).mean(0))))
return np.array(gaps), dnames, np.array(doses), cls, np.array(deltas)
@torch.no_grad()
def build_gallery(drug_order, smiles_map, drug_enc, bridge, device):
projs = []
for start in range(0, len(drug_order), 32):
names = drug_order[start:start+32]
smis = [smiles_map.get(n, "C") or "C" for n in names]
emb = drug_enc(smis); projs.append(bridge(emb)["drug_proj"].cpu().numpy())
return np.concatenate(projs, 0)
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--checkpoint", default="outputs/drug_rank/phase2_best.pt")
ap.add_argument("--config", default="configs/drug_rank_phase2.yaml")
ap.add_argument("--n-conditions", type=int, default=400)
args = ap.parse_args()
import yaml
cfg = yaml.safe_load(open(args.config))
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
models = build_and_load(cfg, args.checkpoint, device)
source_enc, target_enc, gap_enc, drug_enc, bridge = models
dc = cfg["data"]; mc = cfg["model"]
smiles_csv = os.path.join(dc["annotation_dir"], "drug_annotation_master.csv")
dataset = Sciplex3Dataset(h5ad_path=dc["sciplex3_h5ad"], n_hvg=mc["num_genes"],
max_source_cells=64, max_target_cells=64, seed=42,
drug_emb_dim=mc["drug_emb_dim"], preprocessed_path=None,
drug_smiles_csv=smiles_csv)
drug_order = json.load(open(ANNOTATION_DIR / "drug_order.json"))
drug_to_idx = {n: i for i, n in enumerate(drug_order)}
import csv as _csv
smiles_map = {}
with open(os.path.join(dc["annotation_dir"], "drug_annotation_master.csv")) as f:
for row in _csv.DictReader(f):
if row.get("drug_name") and row.get("smiles"):
smiles_map[row["drug_name"].strip()] = row["smiles"].strip()
out = {}
# ---- Signal strength by dose (all conditions) ----
# Cache vehicle means per cell line ONCE (vehicle_cell_idx is shared per line).
veh_mean_cache = {}
def _veh_mean(cell_line, veh_idx):
if cell_line not in veh_mean_cache:
veh_mean_cache[cell_line] = dataset._X[np.asarray(veh_idx)].mean(0)
return veh_mean_cache[cell_line]
by_dose = {}
for cond in dataset._conditions:
veh = np.asarray(cond["vehicle_cell_idx"]); drg = np.asarray(cond["drug_cell_idx"])
if len(veh) == 0 or len(drg) == 0:
continue
d = float(np.linalg.norm(dataset._X[drg].mean(0) - _veh_mean(cond["cell_line"], veh)))
by_dose.setdefault(float(cond["dose"]), []).append(d)
out["signal_by_dose"] = {str(k): {"mean_delta_norm": round(float(np.mean(v)), 4),
"std": round(float(np.std(v)), 4), "n": len(v)}
for k, v in sorted(by_dose.items())}
# vehicle-vehicle noise floor: split vehicle cells of one cell line in half
veh_all = dataset._conditions[0]["vehicle_cell_idx"]
rng = np.random.default_rng(1)
noise = []
for _ in range(20):
perm = rng.permutation(np.asarray(veh_all)); h = len(perm)//2
noise.append(float(np.linalg.norm(dataset._X[perm[:h]].mean(0) - dataset._X[perm[h:2*h]].mean(0))))
out["vehicle_noise_floor"] = round(float(np.mean(noise)), 4)
# ---- Encode a sample of conditions (BOTH with and without cell_line) ----
rng2 = np.random.default_rng(7)
sample = list(rng2.choice(len(dataset._conditions), min(args.n_conditions, len(dataset._conditions)), replace=False))
conds = [dataset._conditions[i] for i in sample]
for tag, pass_cl in [("no_cellline_TRAINMODE", False), ("with_cellline_EVALMODE", True)]:
gaps, dnames, doses, cls, deltas = encode_conditions(dataset, conds, models, device, pass_cl)
gaps_n = gaps / (np.linalg.norm(gaps, axis=1, keepdims=True) + 1e-8)
# gap collapse
sim = gaps_n @ gaps_n.T
off = sim[~np.eye(len(sim), dtype=bool)]
# gallery
gallery = build_gallery(drug_order, smiles_map, drug_enc, bridge, device)
gallery_n = gallery / (np.linalg.norm(gallery, axis=1, keepdims=True) + 1e-8)
scores = gaps_n @ gallery_n.T # [N, 189]
true_idx = np.array([drug_to_idx.get(n, -1) for n in dnames])
valid = true_idx >= 0
ranks = []
for i in np.where(valid)[0]:
order = np.argsort(-scores[i])
ranks.append(int(np.where(order == true_idx[i])[0][0]) + 1)
ranks = np.array(ranks)
# by dose retrieval
vdoses = doses[valid]
dose_ret = {}
for dv in sorted(set(vdoses.tolist())):
rr = ranks[vdoses == dv]
dose_ret[str(dv)] = {"hit@10": round(float((rr <= 10).mean()), 4),
"median_rank": float(np.median(rr)), "n": int(len(rr))}
out[tag] = {
"gap_offdiag_cosine_mean": round(float(off.mean()), 4),
"gap_offdiag_cosine_std": round(float(off.std()), 4),
"gap_perdim_std_mean": round(float(gaps.std(0).mean()), 4),
"true_pair_cosine_mean": round(float(np.mean([scores[i, true_idx[i]] for i in np.where(valid)[0]])), 4),
"overall_hit@10": round(float((ranks <= 10).mean()), 4),
"overall_median_rank": float(np.median(ranks)),
"retrieval_by_dose": dose_ret,
}
# ---- drug_proj separation (shared) ----
gallery = build_gallery(drug_order, smiles_map, drug_enc, bridge, device)
gallery_n = gallery / (np.linalg.norm(gallery, axis=1, keepdims=True) + 1e-8)
gsim = gallery_n @ gallery_n.T
goff = gsim[~np.eye(len(gsim), dtype=bool)]
out["drug_proj_separation"] = {
"pairwise_cosine_mean": round(float(goff.mean()), 4),
"pairwise_cosine_std": round(float(goff.std()), 4),
"pairwise_cosine_max": round(float(goff.max()), 4),
"n_drugs": len(gallery),
}
# ---- Same-drug cross-dose gap consistency (train mode encoding) ----
gaps, dnames, doses, cls, deltas = encode_conditions(dataset, conds, models, device, pass_cell_line=False)
gaps_n = gaps / (np.linalg.norm(gaps, axis=1, keepdims=True) + 1e-8)
from collections import defaultdict
drug_groups = defaultdict(list)
for i, n in enumerate(dnames):
drug_groups[n].append(i)
within, across = [], []
for n, idxs in drug_groups.items():
if len(idxs) >= 2:
for a in range(len(idxs)):
for b in range(a+1, len(idxs)):
within.append(float(gaps_n[idxs[a]] @ gaps_n[idxs[b]]))
# random across-drug pairs
rng3 = np.random.default_rng(3)
for _ in range(2000):
i, j = rng3.integers(0, len(gaps_n), 2)
if dnames[i] != dnames[j]:
across.append(float(gaps_n[i] @ gaps_n[j]))
out["same_drug_gap_consistency"] = {
"within_drug_cosine_mean": round(float(np.mean(within)), 4) if within else None,
"across_drug_cosine_mean": round(float(np.mean(across)), 4) if across else None,
"n_within_pairs": len(within),
"note": "within should be >> across if gap_emb is drug-specific; if similar, dose/noise dominates",
}
os.makedirs("results/diagnostics", exist_ok=True)
outpath = "results/diagnostics/alignment_diag.json"
json.dump(out, open(outpath, "w"), indent=2)
log.info("Saved %s", outpath)
print(json.dumps(out, indent=2))
if __name__ == "__main__":
main()