afg1 commited on
Commit
7a7d25b
Β·
verified Β·
1 Parent(s): 8a0195c

UMAP projection + filtered plot + expandable abstracts + cardiovascular example queries

Browse files
README.md CHANGED
@@ -21,8 +21,9 @@ side by side over a corpus of ~4,000 Europe PMC abstracts (microRNA & disease):
21
  3. **Hybrid** β€” Reciprocal Rank Fusion (RRF, k=60) over the BM25 and dense rankings
22
 
23
  Plus a **metadata filter** (year / journal) that restricts the candidate pool *before*
24
- retrieval β€” the distinctly vector-DB feature β€” and a **PCA scatter plot** for spatial
25
  intuition, with connector lines from the query to the documents each method retrieved.
 
26
 
27
  ## How it's built (read this before the talk)
28
 
@@ -48,9 +49,11 @@ mislead rather than teach.
48
 
49
  ## Teaching caveat on the plot
50
 
51
- The 2D PCA projection distorts true high-dimensional distances, so the retrieved points
52
- may not be the visually-closest dots. **The ranked lists are authoritative; the plot is
53
- only for intuition.**
 
 
54
 
55
  ## Rebuild the index locally
56
 
 
21
  3. **Hybrid** β€” Reciprocal Rank Fusion (RRF, k=60) over the BM25 and dense rankings
22
 
23
  Plus a **metadata filter** (year / journal) that restricts the candidate pool *before*
24
+ retrieval β€” the distinctly vector-DB feature β€” and a **UMAP scatter plot** for spatial
25
  intuition, with connector lines from the query to the documents each method retrieved.
26
+ Click any result to expand its full abstract.
27
 
28
  ## How it's built (read this before the talk)
29
 
 
49
 
50
  ## Teaching caveat on the plot
51
 
52
+ The plot uses **UMAP** for the 2D projection (it shows local cluster structure better
53
+ than PCA). Two honest caveats: UMAP warps global distances, and the live query is placed
54
+ by an *approximate* out-of-sample `transform`, so its position is only indicative.
55
+ Retrieved points may not be the visually-closest dots. **The ranked lists are
56
+ authoritative; the plot is only for intuition.**
57
 
58
  ## Rebuild the index locally
59
 
app.py CHANGED
@@ -3,12 +3,13 @@
3
  A teaching app for EMBL-EBI researchers. Reads pre-built artifacts from ./data/ (made by
4
  build_index.py) and only embeds the user's LIVE query at runtime. Read top-to-bottom:
5
 
6
- artifacts -> embed query (GPU) -> filter -> 3 rankings -> RRF -> plot
7
 
8
  ZeroGPU note: the embedding model is loaded on CPU at import. GPU is touched ONLY inside
9
  the @spaces.GPU functions. Do not move the model to CUDA anywhere else.
10
  """
11
 
 
12
  import json
13
  import os
14
 
@@ -30,8 +31,8 @@ from text_utils import tokenize
30
  # ====================================================================================
31
  D = config.DATA_DIR
32
  EMBEDDINGS = np.load(os.path.join(D, "embeddings.npy")) # (N, dim), L2-normalised
33
- PCA_COORDS = np.load(os.path.join(D, "pca_coords.npy")) # (N, 2)
34
- PCA = joblib.load(os.path.join(D, "pca.joblib"))
35
  META = pd.read_parquet(os.path.join(D, "metadata.parquet"))
36
  with open(os.path.join(D, "bm25_tokens.json")) as f:
37
  BM25 = BM25Okapi(json.load(f))
@@ -54,7 +55,9 @@ METHOD_COLORS = {"BM25": "#ff7f0e", "Dense": "#1f77b4", "Hybrid": "#2ca02c"}
54
  # Embedding model β€” loaded on CPU. GPU is used ONLY inside @spaces.GPU below.
55
  # ====================================================================================
56
  MODEL = SentenceTransformer(config.EMBEDDING_MODEL, device="cpu")
57
- assert MODEL.get_sentence_embedding_dimension() == EMBEDDINGS.shape[1], (
 
 
58
  "Model dim doesn't match committed embeddings β€” did you swap EMBEDDING_MODEL "
59
  "without re-running build_index.py?"
60
  )
@@ -124,29 +127,42 @@ def rrf_fuse(bm25: np.ndarray, dense: np.ndarray, cand: np.ndarray, k: int):
124
 
125
 
126
  def format_results(method: str, results: list[tuple[int, float]], score_label: str) -> str:
127
- """Render a ranked list as Markdown. This list is the SOURCE OF TRUTH for retrieval."""
 
 
 
128
  if not results:
129
- return f"### {method}\n\n_No results._"
130
- lines = [f"### {method}"]
131
  for rank, (doc, score) in enumerate(results, start=1):
132
- snippet = ABSTRACTS[doc][:200].rsplit(" ", 1)[0] + "…"
133
- lines.append(
134
- f"**{rank}. {TITLES[doc]}** \n"
135
- f"`{score_label}={score:.3f}` Β· {int(YEARS[doc])} Β· *{JOURNALS[doc]}* \n"
136
- f"{snippet}\n"
 
 
 
 
 
137
  )
138
- return "\n".join(lines)
 
139
 
 
 
140
 
141
- def make_plot(query_coord=None, retrieved=None, extra_point=None) -> go.Figure:
142
- """PCA scatter of the whole corpus, plus query point + connector lines to hits."""
 
143
  fig = go.Figure()
144
- # Whole corpus as grey background β€” "the space the database searches through".
 
145
  fig.add_trace(
146
  go.Scatter(
147
- x=PCA_COORDS[:, 0], y=PCA_COORDS[:, 1], mode="markers",
148
- marker=dict(size=4, color="lightgrey"), text=TITLES, hoverinfo="text",
149
- name="corpus",
150
  )
151
  )
152
  if retrieved and query_coord is not None:
@@ -154,16 +170,16 @@ def make_plot(query_coord=None, retrieved=None, extra_point=None) -> go.Figure:
154
  color = METHOD_COLORS[method]
155
  xs, ys = [], []
156
  for doc, _ in results: # connector lines query -> each hit
157
- xs += [query_coord[0], PCA_COORDS[doc, 0], None]
158
- ys += [query_coord[1], PCA_COORDS[doc, 1], None]
159
  fig.add_trace(
160
  go.Scatter(x=xs, y=ys, mode="lines", line=dict(color=color, width=1),
161
  opacity=0.5, name=f"{method} links", hoverinfo="skip")
162
  )
163
  fig.add_trace(
164
  go.Scatter(
165
- x=[PCA_COORDS[d, 0] for d, _ in results],
166
- y=[PCA_COORDS[d, 1] for d, _ in results],
167
  mode="markers",
168
  marker=dict(size=11, color=color, symbol="circle-open", line=dict(width=2)),
169
  text=[TITLES[d] for d, _ in results], hoverinfo="text", name=f"{method} hits",
@@ -212,8 +228,10 @@ def run_search(query, k, year_lo, year_hi, journal):
212
  dense_top = top_k(dense_scores, cand, k)
213
  hybrid_top = rrf_fuse(bm25_scores, dense_scores, cand, k)
214
 
215
- query_coord = PCA.transform(qvec.reshape(1, -1))[0]
216
- fig = make_plot(query_coord, {"BM25": bm25_top, "Dense": dense_top, "Hybrid": hybrid_top})
 
 
217
 
218
  return (
219
  format_results("BM25", bm25_top, "score"),
@@ -229,39 +247,45 @@ def run_embed_own(text):
229
  if not text:
230
  return make_plot()
231
  vec = embed_text(text) # <-- GPU work
232
- coord = PCA.transform(vec.reshape(1, -1))[0]
233
  return make_plot(extra_point=(coord, text[:80]))
234
 
235
 
236
  # ====================================================================================
237
- # Example queries β€” TODO: the presenter fills these in.
238
- # Each placeholder is meant to make ONE method visibly win; pick real strings against the
239
- # current microRNA/disease corpus before the talk.
240
  # ====================================================================================
241
- # BM25 should win: a precise token/ID that appears VERBATIM in abstracts.
242
- EXACT_ID_QUERY = "TODO_EXACT_ID: e.g. a specific miRNA like 'miR-21' or a gene symbol"
243
- # Dense should win: a conceptual paraphrase using NONE of the corpus's exact words.
244
- PARAPHRASE_QUERY = "TODO_PARAPHRASE: e.g. 'small RNAs that switch genes off in tumours'"
245
- # Lexical vs semantic gap: an acronym whose expansion is what's written in the text.
246
- ACRONYM_QUERY = "TODO_ACRONYM: e.g. an acronym vs its full form"
247
- # BM25 strong on rare exact tokens.
248
- RARE_TERM_QUERY = "TODO_RARE_TERM: e.g. a rare, very specific technical term"
249
- # Hybrid should win: a broad topic where fusing lexical + semantic beats either alone.
250
- BROAD_CONCEPT_QUERY = "TODO_BROAD: e.g. 'microRNA biomarkers for early cancer detection'"
 
 
 
 
 
 
251
 
252
  EXAMPLES = [
253
  ("Exact ID (BM25)", EXACT_ID_QUERY),
254
  ("Paraphrase (Dense)", PARAPHRASE_QUERY),
255
- ("Acronym", ACRONYM_QUERY),
256
- ("Rare term (BM25)", RARE_TERM_QUERY),
257
  ("Broad concept (Hybrid)", BROAD_CONCEPT_QUERY),
258
  ]
259
 
260
  PLOT_CAVEAT = (
261
- "⚠️ **The 2D projection distorts true distances.** This PCA view captures only a "
262
- "sliver of the 384-dimensional space, so retrieved points may *not* be the "
263
- "visually-closest dots. **The ranked lists above are authoritative** β€” the plot is "
264
- "only for spatial intuition."
 
265
  )
266
 
267
  # ====================================================================================
@@ -294,12 +318,13 @@ with gr.Blocks(title="RAG retrieval: BM25 vs Dense vs Hybrid") as demo:
294
  search_btn = gr.Button("Search", variant="primary")
295
  filter_info = gr.Markdown("β€”")
296
 
 
297
  with gr.Row():
298
- bm25_out = gr.Markdown(label="BM25")
299
- dense_out = gr.Markdown(label="Dense")
300
- hybrid_out = gr.Markdown(label="Hybrid")
301
 
302
- gr.Markdown("## Vector space (PCA projection)")
303
  plot = gr.Plot(value=make_plot())
304
  gr.Markdown(PLOT_CAVEAT)
305
 
 
3
  A teaching app for EMBL-EBI researchers. Reads pre-built artifacts from ./data/ (made by
4
  build_index.py) and only embeds the user's LIVE query at runtime. Read top-to-bottom:
5
 
6
+ artifacts -> embed query (GPU) -> filter -> 3 rankings -> RRF -> UMAP plot
7
 
8
  ZeroGPU note: the embedding model is loaded on CPU at import. GPU is touched ONLY inside
9
  the @spaces.GPU functions. Do not move the model to CUDA anywhere else.
10
  """
11
 
12
+ import html
13
  import json
14
  import os
15
 
 
31
  # ====================================================================================
32
  D = config.DATA_DIR
33
  EMBEDDINGS = np.load(os.path.join(D, "embeddings.npy")) # (N, dim), L2-normalised
34
+ COORDS = np.load(os.path.join(D, "umap_coords.npy")) # (N, 2) UMAP projection
35
+ REDUCER = joblib.load(os.path.join(D, "umap.joblib")) # projects new points via .transform
36
  META = pd.read_parquet(os.path.join(D, "metadata.parquet"))
37
  with open(os.path.join(D, "bm25_tokens.json")) as f:
38
  BM25 = BM25Okapi(json.load(f))
 
55
  # Embedding model β€” loaded on CPU. GPU is used ONLY inside @spaces.GPU below.
56
  # ====================================================================================
57
  MODEL = SentenceTransformer(config.EMBEDDING_MODEL, device="cpu")
58
+ # Prefer the new method name, fall back to the deprecated one across versions.
59
+ _get_dim = getattr(MODEL, "get_embedding_dimension", None) or MODEL.get_sentence_embedding_dimension
60
+ assert _get_dim() == EMBEDDINGS.shape[1], (
61
  "Model dim doesn't match committed embeddings β€” did you swap EMBEDDING_MODEL "
62
  "without re-running build_index.py?"
63
  )
 
127
 
128
 
129
  def format_results(method: str, results: list[tuple[int, float]], score_label: str) -> str:
130
+ """Render a ranked list as HTML. Each result is a <details> β€” click the title to
131
+ expand the FULL abstract (helpful for seeing *why* something ranked highly).
132
+ This list is the SOURCE OF TRUTH for retrieval.
133
+ """
134
  if not results:
135
+ return f"<h3>{method}</h3><p><em>No results.</em></p>"
136
+ parts = [f"<h3>{method}</h3>"]
137
  for rank, (doc, score) in enumerate(results, start=1):
138
+ title = html.escape(TITLES[doc]) # abstracts contain <, >, & β€” must escape
139
+ abstract = html.escape(ABSTRACTS[doc])
140
+ journal = html.escape(str(JOURNALS[doc]))
141
+ parts.append(
142
+ "<details style='margin-bottom:10px;border-bottom:1px solid #ddd;padding-bottom:6px;'>"
143
+ f"<summary style='cursor:pointer;'><b>{rank}. {title}</b><br>"
144
+ f"<code>{score_label}={score:.3f}</code> Β· {int(YEARS[doc])} Β· <em>{journal}</em>"
145
+ "</summary>"
146
+ f"<p style='margin-top:6px;font-size:0.9em;line-height:1.4;'>{abstract}</p>"
147
+ "</details>"
148
  )
149
+ return "".join(parts)
150
+
151
 
152
+ def make_plot(query_coord=None, retrieved=None, extra_point=None, cand=None) -> go.Figure:
153
+ """UMAP scatter of the corpus, plus query point + connector lines to hits.
154
 
155
+ `cand` (optional) = the indices that passed the metadata filter; when given, only
156
+ those points are drawn as the grey background, so the plot reflects the filter.
157
+ """
158
  fig = go.Figure()
159
+ bg = np.arange(len(COORDS)) if cand is None else cand
160
+ # Grey background β€” "the space the database searches through" (after filtering).
161
  fig.add_trace(
162
  go.Scatter(
163
+ x=COORDS[bg, 0], y=COORDS[bg, 1], mode="markers",
164
+ marker=dict(size=4, color="lightgrey"),
165
+ text=[TITLES[i] for i in bg], hoverinfo="text", name=f"corpus ({len(bg)})",
166
  )
167
  )
168
  if retrieved and query_coord is not None:
 
170
  color = METHOD_COLORS[method]
171
  xs, ys = [], []
172
  for doc, _ in results: # connector lines query -> each hit
173
+ xs += [query_coord[0], COORDS[doc, 0], None]
174
+ ys += [query_coord[1], COORDS[doc, 1], None]
175
  fig.add_trace(
176
  go.Scatter(x=xs, y=ys, mode="lines", line=dict(color=color, width=1),
177
  opacity=0.5, name=f"{method} links", hoverinfo="skip")
178
  )
179
  fig.add_trace(
180
  go.Scatter(
181
+ x=[COORDS[d, 0] for d, _ in results],
182
+ y=[COORDS[d, 1] for d, _ in results],
183
  mode="markers",
184
  marker=dict(size=11, color=color, symbol="circle-open", line=dict(width=2)),
185
  text=[TITLES[d] for d, _ in results], hoverinfo="text", name=f"{method} hits",
 
228
  dense_top = top_k(dense_scores, cand, k)
229
  hybrid_top = rrf_fuse(bm25_scores, dense_scores, cand, k)
230
 
231
+ query_coord = REDUCER.transform(qvec.reshape(1, -1))[0]
232
+ fig = make_plot(
233
+ query_coord, {"BM25": bm25_top, "Dense": dense_top, "Hybrid": hybrid_top}, cand=cand
234
+ )
235
 
236
  return (
237
  format_results("BM25", bm25_top, "score"),
 
247
  if not text:
248
  return make_plot()
249
  vec = embed_text(text) # <-- GPU work
250
+ coord = REDUCER.transform(vec.reshape(1, -1))[0]
251
  return make_plot(extra_point=(coord, text[:80]))
252
 
253
 
254
  # ====================================================================================
255
+ # Example queries β€” tuned to the cardiovascular slice of this microRNA/disease corpus.
256
+ # Each is chosen (and verified against the corpus) to make ONE method visibly win.
 
257
  # ====================================================================================
258
+ # BM25 wins: an exact miRNA identifier. The literal token "miR-208a" is matched precisely;
259
+ # dense retrieval drifts to other, merely-similar cardiac miRNA papers.
260
+ EXACT_ID_QUERY = "miR-208a"
261
+ # Dense wins: a plain-English paraphrase using NONE of the corpus jargon (no "microRNA",
262
+ # "fibrosis", "myocardial"). BM25 has almost no tokens to match; dense maps the meaning
263
+ # to miRNA-in-cardiac-fibrosis papers.
264
+ PARAPHRASE_QUERY = "small RNA molecules that worsen scarring after a heart attack"
265
+ # BM25 wins: a bare acronym (acute myocardial infarction). "AMI" is a literal token BM25
266
+ # nails, but as a 3-letter string it carries little semantic signal, so dense flounders.
267
+ ACRONYM_QUERY = "AMI"
268
+ # BM25 wins: a specific molecule + mechanism. BM25 pins the exact miR-499 papers; dense
269
+ # returns topically-similar cardiomyocyte-apoptosis papers about *other* miRNAs.
270
+ RARE_TERM_QUERY = "miR-499 cardiomyocyte apoptosis"
271
+ # Hybrid wins: a broad real query where lexical and semantic each surface good-but-
272
+ # different papers, and RRF fuses them into the strongest combined ranking.
273
+ BROAD_CONCEPT_QUERY = "circulating microRNA biomarkers for cardiovascular disease"
274
 
275
  EXAMPLES = [
276
  ("Exact ID (BM25)", EXACT_ID_QUERY),
277
  ("Paraphrase (Dense)", PARAPHRASE_QUERY),
278
+ ("Acronym (BM25)", ACRONYM_QUERY),
279
+ ("Specific molecule (BM25)", RARE_TERM_QUERY),
280
  ("Broad concept (Hybrid)", BROAD_CONCEPT_QUERY),
281
  ]
282
 
283
  PLOT_CAVEAT = (
284
+ "⚠️ **The 2D UMAP projection distorts true distances.** UMAP preserves rough local "
285
+ "neighbourhoods but warps global distances, and the live query is placed by an "
286
+ "*approximate* out-of-sample fit β€” so its position (and the connector lines) are only "
287
+ "indicative. Retrieved points may *not* be the visually-closest dots. "
288
+ "**The ranked lists above are authoritative**; the plot is only for intuition."
289
  )
290
 
291
  # ====================================================================================
 
318
  search_btn = gr.Button("Search", variant="primary")
319
  filter_info = gr.Markdown("β€”")
320
 
321
+ gr.Markdown("_Click any result to expand its full abstract._")
322
  with gr.Row():
323
+ bm25_out = gr.HTML(label="BM25")
324
+ dense_out = gr.HTML(label="Dense")
325
+ hybrid_out = gr.HTML(label="Hybrid")
326
 
327
+ gr.Markdown("## Vector space (UMAP projection)")
328
  plot = gr.Plot(value=make_plot())
329
  gr.Markdown(PLOT_CAVEAT)
330
 
build_index.py CHANGED
@@ -15,8 +15,8 @@ import joblib
15
  import numpy as np
16
  import pandas as pd
17
  import requests
 
18
  from sentence_transformers import SentenceTransformer
19
- from sklearn.decomposition import PCA
20
 
21
  import config
22
  from text_utils import tokenize
@@ -106,13 +106,18 @@ def main():
106
  np.save(os.path.join(config.DATA_DIR, "embeddings.npy"), embeddings)
107
  print(f"Saved embeddings: {embeddings.shape}")
108
 
109
- # 3. PCA ------------------------------------------------------------------------
110
- print("Fitting 2D PCA...")
111
- pca = PCA(n_components=2, random_state=0)
112
- coords = pca.fit_transform(embeddings).astype(np.float32)
113
- joblib.dump(pca, os.path.join(config.DATA_DIR, "pca.joblib"))
114
- np.save(os.path.join(config.DATA_DIR, "pca_coords.npy"), coords)
115
- print(f"Saved PCA + corpus coords: {coords.shape}")
 
 
 
 
 
116
 
117
  # 4. BM25 tokens ----------------------------------------------------------------
118
  print("Tokenising for BM25...")
 
15
  import numpy as np
16
  import pandas as pd
17
  import requests
18
+ import umap
19
  from sentence_transformers import SentenceTransformer
 
20
 
21
  import config
22
  from text_utils import tokenize
 
106
  np.save(os.path.join(config.DATA_DIR, "embeddings.npy"), embeddings)
107
  print(f"Saved embeddings: {embeddings.shape}")
108
 
109
+ # 3. UMAP -----------------------------------------------------------------------
110
+ # UMAP shows local cluster structure better than PCA. It also supports projecting
111
+ # new/out-of-sample points via .transform() β€” we use that to place the live query.
112
+ # Caveat: single-point transform is an approximate fit, so query placement is rough.
113
+ print("Fitting 2D UMAP (cosine metric)...")
114
+ reducer = umap.UMAP(
115
+ n_components=2, metric="cosine", n_neighbors=15, min_dist=0.1, random_state=42
116
+ )
117
+ coords = reducer.fit_transform(embeddings).astype(np.float32)
118
+ joblib.dump(reducer, os.path.join(config.DATA_DIR, "umap.joblib"))
119
+ np.save(os.path.join(config.DATA_DIR, "umap_coords.npy"), coords)
120
+ print(f"Saved UMAP reducer + corpus coords: {coords.shape}")
121
 
122
  # 4. BM25 tokens ----------------------------------------------------------------
123
  print("Tokenising for BM25...")
data/{pca.joblib β†’ umap.joblib} RENAMED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:f13f7a85bd6fe24bf75709961bcdefb4e2ca4587e3e377cf67f29a7b765baefd
3
- size 5583
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:2690f51123d5179bf011fad8d62e5ee5856f31fb0ed78f2a1038b2d4d356b24a
3
+ size 6914713
data/{pca_coords.npy β†’ umap_coords.npy} RENAMED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:ca78a6240daa3c539a57f76faf6efdaf2419b7031d683f780adc112e5eef9746
3
  size 32072
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:a1a93e57503710b82ebe52758bfddfa338e93e089fcc90de3b516655f6c006a1
3
  size 32072
requirements.txt CHANGED
@@ -3,6 +3,7 @@ spaces
3
  sentence-transformers
4
  rank-bm25
5
  scikit-learn
 
6
  numpy
7
  pandas
8
  pyarrow
 
3
  sentence-transformers
4
  rank-bm25
5
  scikit-learn
6
+ umap-learn
7
  numpy
8
  pandas
9
  pyarrow