| """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 = {} |
|
|
| |
| |
| 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())} |
| |
| 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) |
|
|
| |
| 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) |
| |
| sim = gaps_n @ gaps_n.T |
| off = sim[~np.eye(len(sim), dtype=bool)] |
| |
| 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 |
| 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) |
| |
| 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, |
| } |
|
|
| |
| 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), |
| } |
|
|
| |
| 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]])) |
| |
| 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() |
|
|