ollibolli commited on
Commit
d5ba468
·
1 Parent(s): 4b9c5bf

move to umap

Browse files
Files changed (2) hide show
  1. app.py +66 -22
  2. requirements.txt +2 -1
app.py CHANGED
@@ -3,6 +3,8 @@ import torch
3
  import numpy as np
4
  from transformers import AutoTokenizer, AutoModelForCausalLM
5
  from sklearn.decomposition import PCA
 
 
6
  import json
7
 
8
  # Model configuration
@@ -17,23 +19,40 @@ model = AutoModelForCausalLM.from_pretrained(MODEL_ID, dtype=torch.float32).to(d
17
  vocab_size = tokenizer.vocab_size if tokenizer.vocab_size is not None else len(tokenizer)
18
  vocab_tokens = tokenizer.convert_ids_to_tokens(list(range(vocab_size)))
19
 
20
- # Cache for embeddings and PCA (computed once at startup)
21
  embeddings_cache = None
22
- pca_cache = None
23
- pca_projections_cache = None
 
 
 
 
 
 
24
 
25
  def initialize_embeddings():
26
- """Initialize embeddings and PCA projections once at startup"""
27
- global embeddings_cache, pca_cache, pca_projections_cache
28
-
29
- # Get embeddings
30
- embeddings_cache = model.get_input_embeddings().weight.detach().cpu().numpy() # [V, d]
31
-
32
- # Compute PCA
33
- pca_cache = PCA(n_components=2, random_state=0)
34
- pca_projections_cache = pca_cache.fit_transform(embeddings_cache) # [V, 2]
35
-
36
- return embeddings_cache, pca_projections_cache
 
 
 
 
 
 
 
 
 
 
 
37
 
38
  # Initialize embeddings at startup
39
  initialize_embeddings()
@@ -145,10 +164,13 @@ def predict_comprehensive(
145
 
146
  # PCA projections if requested
147
  if include_pca:
148
- if pca_projections_cache is not None:
 
149
  result["pca"] = {
150
- "projections": pca_projections_cache.tolist(),
151
- "explained_variance_ratio": pca_cache.explained_variance_ratio_.tolist()
 
 
152
  }
153
 
154
  return result
@@ -206,7 +228,7 @@ with gr.Blocks(title="Token Probability Visualization API") as demo:
206
  with gr.Row():
207
  top_k_slider = gr.Slider(0, 200, step=1, value=20, label="Top-K tokens")
208
  include_embeddings = gr.Checkbox(False, label="Include Embeddings Sample")
209
- include_pca = gr.Checkbox(True, label="Include PCA Projections")
210
  include_unconditional = gr.Checkbox(True, label="Include Unconditional")
211
  use_logprobs = gr.Checkbox(True, label="Include Log Probabilities")
212
 
@@ -245,17 +267,39 @@ with gr.Blocks(title="Token Probability Visualization API") as demo:
245
  api_name="embeddings"
246
  )
247
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
248
  gr.Markdown("""
249
  ## API Usage
250
-
251
  ### Endpoints:
252
- - `/predict`: Main endpoint for token probabilities with optional embeddings and PCA
253
  - `/embeddings`: Get embeddings for specific token IDs
 
254
 
255
  ### Response includes:
256
  - Token probabilities (conditional and unconditional)
257
  - Log probabilities
258
- - PCA projections for all vocabulary tokens
259
  - Token embeddings (sample or specific)
260
  - Vocabulary mappings
261
 
@@ -267,4 +311,4 @@ with gr.Blocks(title="Token Probability Visualization API") as demo:
267
  """)
268
 
269
  if __name__ == "__main__":
270
- demo.launch(server_name="0.0.0.0", server_port=7860, share=False)
 
3
  import numpy as np
4
  from transformers import AutoTokenizer, AutoModelForCausalLM
5
  from sklearn.decomposition import PCA
6
+ from sklearn.preprocessing import normalize
7
+ from umap import UMAP
8
  import json
9
 
10
  # Model configuration
 
19
  vocab_size = tokenizer.vocab_size if tokenizer.vocab_size is not None else len(tokenizer)
20
  vocab_tokens = tokenizer.convert_ids_to_tokens(list(range(vocab_size)))
21
 
22
+ # Cache for embeddings and UMAP (computed once at startup)
23
  embeddings_cache = None
24
+ umap_projections_cache = None
25
+ umap_params = {
26
+ "n_neighbors": 75,
27
+ "min_dist": 0.15,
28
+ "metric": "cosine",
29
+ "random_state": 0,
30
+ "n_components": 2,
31
+ }
32
 
33
  def initialize_embeddings():
34
+ """Initialize embeddings and UMAP projections once at startup"""
35
+ global embeddings_cache, umap_projections_cache
36
+
37
+ # Get embeddings and cast to float32 for performance
38
+ embeddings_cache = (
39
+ model.get_input_embeddings().weight.detach().cpu().numpy().astype(np.float32)
40
+ ) # [V, d]
41
+
42
+ # Normalize rows (cosine distance works best with normalized vectors)
43
+ norm_embeds = normalize(embeddings_cache, norm="l2", axis=1)
44
+
45
+ # Optional PCA(50) for speed/stability before UMAP (not a fallback layout)
46
+ d = norm_embeds.shape[1]
47
+ n_components = min(50, d)
48
+ pca50 = PCA(n_components=n_components, svd_solver="randomized", random_state=0)
49
+ reduced = pca50.fit_transform(norm_embeds)
50
+
51
+ # UMAP to 2D (full vocab)
52
+ umap_model = UMAP(**umap_params)
53
+ umap_projections_cache = umap_model.fit_transform(reduced).astype(np.float32) # [V, 2]
54
+
55
+ return embeddings_cache, umap_projections_cache
56
 
57
  # Initialize embeddings at startup
58
  initialize_embeddings()
 
164
 
165
  # PCA projections if requested
166
  if include_pca:
167
+ if umap_projections_cache is not None:
168
+ # For compatibility we keep the `pca` key but fill with UMAP projections
169
  result["pca"] = {
170
+ "projections": umap_projections_cache.tolist(),
171
+ "explained_variance_ratio": [0.0, 0.0], # Placeholder; UMAP has no variance ratio
172
+ "method": "umap",
173
+ "umap_params": umap_params,
174
  }
175
 
176
  return result
 
228
  with gr.Row():
229
  top_k_slider = gr.Slider(0, 200, step=1, value=20, label="Top-K tokens")
230
  include_embeddings = gr.Checkbox(False, label="Include Embeddings Sample")
231
+ include_pca = gr.Checkbox(False, label="Include UMAP Projections (2D)")
232
  include_unconditional = gr.Checkbox(True, label="Include Unconditional")
233
  use_logprobs = gr.Checkbox(True, label="Include Log Probabilities")
234
 
 
267
  api_name="embeddings"
268
  )
269
 
270
+ with gr.Tab("Layout API"):
271
+ gr.Markdown("### Get 2D token layout (UMAP) once and cache client-side")
272
+ with gr.Row():
273
+ get_layout_btn = gr.Button("Get UMAP Layout", variant="primary")
274
+ layout_output = gr.JSON(label="UMAP Layout Response")
275
+
276
+ def get_layout_endpoint():
277
+ return {
278
+ "method": "umap",
279
+ "umap_params": umap_params,
280
+ "tokens": [nice_tok(t) for t in vocab_tokens],
281
+ "projections": umap_projections_cache.tolist() if umap_projections_cache is not None else None,
282
+ }
283
+
284
+ get_layout_btn.click(
285
+ fn=get_layout_endpoint,
286
+ inputs=None,
287
+ outputs=layout_output,
288
+ api_name="layout",
289
+ )
290
+
291
  gr.Markdown("""
292
  ## API Usage
293
+
294
  ### Endpoints:
295
+ - `/predict`: Main endpoint for token probabilities with optional embeddings and UMAP
296
  - `/embeddings`: Get embeddings for specific token IDs
297
+ - `/layout`: Get 2D UMAP projections (fetch once, reuse client-side)
298
 
299
  ### Response includes:
300
  - Token probabilities (conditional and unconditional)
301
  - Log probabilities
302
+ - UMAP projections for all vocabulary tokens (if requested or via `/layout`)
303
  - Token embeddings (sample or specific)
304
  - Vocabulary mappings
305
 
 
311
  """)
312
 
313
  if __name__ == "__main__":
314
+ demo.launch(server_name="0.0.0.0", server_port=7860, share=False)
requirements.txt CHANGED
@@ -3,4 +3,5 @@ torch>=2.0.0
3
  transformers>=4.30.0
4
  scikit-learn>=1.3.0
5
  numpy>=1.24.0
6
- spaces
 
 
3
  transformers>=4.30.0
4
  scikit-learn>=1.3.0
5
  numpy>=1.24.0
6
+ spaces
7
+ umap-learn>=0.5.6