-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathbuild_semantic_indices.py
More file actions
128 lines (113 loc) · 5.71 KB
/
Copy pathbuild_semantic_indices.py
File metadata and controls
128 lines (113 loc) · 5.71 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
"""
build_semantic_indices.py
=========================
Pre-compute MedEmbed embeddings for semantic fallback matching.
Generates three embedding sets:
1. Symptom canonical names → hpo_data_final/symptom_embeddings.npy + .json
2. MONDO disease synonyms → hpo_data_final/disease_embeddings.npy + .json
3. KG disease names → hpo_data_final/kg_disease_embeddings.npy + .json
Each .json index is a list of [id, name] in row-order of the .npy array.
"""
import json
import numpy as np
from pathlib import Path
from sentence_transformers import SentenceTransformer
MONDO_DIR = Path(__file__).parent / "mondo_data_new"
OUT_DIR = Path(__file__).parent / "hpo_data_final"
OUT_DIR.mkdir(parents=True, exist_ok=True)
print("Loading MedEmbed model...")
model = SentenceTransformer("abhinand/MedEmbed-large-v0.1", device="cpu")
# ---------------------------------------------------------------------------
# 1. Symptom canonical names (skip if already built)
# ---------------------------------------------------------------------------
sym_npy = OUT_DIR / "symptom_embeddings.npy"
sym_json = OUT_DIR / "symptom_embedding_index.json"
if sym_npy.exists() and sym_json.exists():
print("\n[1/3] Symptom embeddings already exist, skipping.")
else:
print("\n[1/3] Symptom canonical names...")
sc = json.load(open(MONDO_DIR / "symptom_catalog.json"))
id_to_symptom = {hid: entry["label"] for hid, entry in sc.items()}
sym_labels = [(hid, name) for hid, name in id_to_symptom.items()]
sym_names = [name for _, name in sym_labels]
print(f" Encoding {len(sym_names):,} names...")
sym_emb = model.encode(sym_names, normalize_embeddings=True, show_progress_bar=True)
np.save(sym_npy, sym_emb)
with open(sym_json, "w") as f:
json.dump(sym_labels, f)
print(f" Saved symptom_embeddings.npy ({sym_emb.shape})")
print(f" Saved symptom_embedding_index.json ({len(sym_labels)} entries)")
# ---------------------------------------------------------------------------
# 2. MONDO disease canonical labels (skip if already built)
# ---------------------------------------------------------------------------
dis_npy = OUT_DIR / "disease_embeddings.npy"
dis_json = OUT_DIR / "disease_embedding_index.json"
if dis_npy.exists() and dis_json.exists():
print("\n[2/3] MONDO disease embeddings already exist, skipping.")
else:
print("\n[2/3] MONDO disease canonical labels...")
dc = json.load(open(MONDO_DIR / "disease_catalog.json"))
dis_labels = [(did, entry["label"]) for did, entry in dc.items()]
dis_names = [name for _, name in dis_labels]
print(f" Encoding {len(dis_names):,} names...")
dis_emb = model.encode(dis_names, normalize_embeddings=True, show_progress_bar=True)
np.save(dis_npy, dis_emb)
with open(dis_json, "w") as f:
json.dump(dis_labels, f)
print(f" Saved disease_embeddings.npy ({dis_emb.shape})")
print(f" Saved disease_embedding_index.json ({len(dis_labels)} entries)")
# ---------------------------------------------------------------------------
# 3. KG disease names — only diseases with scorable (HPO-bridgeable) phenotypes
# ---------------------------------------------------------------------------
kg_npy = OUT_DIR / "kg_disease_embeddings.npy"
kg_json = OUT_DIR / "kg_disease_embedding_index.json"
if kg_npy.exists() and kg_json.exists():
print("\n[3/3] KG disease embeddings already exist, skipping.")
else:
print("\n[3/3] KG disease names (scorable only)...")
# Load kg.feather to build phenotype x_id → symptom index
import pandas as pd
kg_df = pd.read_feather(str(Path(__file__).parent / "kg.feather"))
pheno = kg_df[kg_df["relation"] == "disease_phenotype_positive"]
# Build parent/child map for dedup (same as scanner)
pp = kg_df[kg_df["relation"] == "phenotype_phenotype"]
parent_of: dict[str, str] = {}
for _, row in pp.iterrows():
parent_of[str(row["y_name"]).strip().lower()] = str(row["x_name"]).strip().lower()
# Load canonical symptom name → HPO ID mapping for bridging
sc = json.load(open(MONDO_DIR / "symptom_catalog.json"))
canon_name_to_hpo = {entry["label"].lower(): hid for hid, entry in sc.items()}
# Group by x_id: dedup symptoms, then bridge to HPO IDs
kg_labels = []
for xid, group in pheno.groupby("x_id"):
disease_name = str(group["x_name"].iloc[0]).strip()
if not disease_name:
continue
raw = set(str(s).strip() for s in group["y_name"].tolist())
# Dedup (same logic as scanner._dedup_symptoms)
lower_map = {s.lower(): s for s in raw}
deduped = set()
for s_lower, s_orig in lower_map.items():
if s_lower in parent_of and parent_of[s_lower] in lower_map:
continue
deduped.add(s_orig)
# Bridge to HPO IDs
bridged = {canon_name_to_hpo.get(s.lower()) for s in deduped if canon_name_to_hpo.get(s.lower())}
if bridged:
kg_labels.append((xid, disease_name))
print(f" {len(kg_labels):,} scorable KG x_ids")
if kg_labels:
kg_names = [name for _, name in kg_labels]
print(f" Encoding {len(kg_names):,} names...")
kg_emb = model.encode(kg_names, normalize_embeddings=True, show_progress_bar=True)
np.save(kg_npy, kg_emb)
with open(kg_json, "w") as f:
json.dump(kg_labels, f)
print(f" Saved kg_disease_embeddings.npy ({kg_emb.shape})")
print(f" Saved kg_disease_embedding_index.json ({len(kg_labels)} entries)")
else:
print(" WARNING: no KG diseases with scorable phenotypes found!")
np.save(kg_npy, np.zeros((0, 1024), dtype="float32"))
with open(kg_json, "w") as f:
json.dump([], f)
print("\nDone — all semantic indices built.")