mrs83 commited on
Commit
2de1465
·
verified ·
1 Parent(s): afe0e8f

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +136 -57
app.py CHANGED
@@ -26,22 +26,41 @@ class StringStoppingCriteria(StoppingCriteria):
26
  def create_demo(api_url, model_id=None):
27
  # Standalone mode check
28
  hf_model_id = model_id or os.environ.get("HF_MODEL_ID", "ethicalabs/Echo-DSRN-Small-Instruct")
29
-
30
  local_model = None
31
  local_tokenizer = None
32
-
 
33
  if hf_model_id:
34
  print(f"📦 Standalone Mode: Loading {hf_model_id} via Transformers...")
35
  try:
36
  local_tokenizer = AutoTokenizer.from_pretrained(hf_model_id, trust_remote_code=True)
37
  local_model = AutoModelForCausalLM.from_pretrained(
38
- hf_model_id,
39
  trust_remote_code=True,
40
  torch_dtype=torch.bfloat16 if torch.cuda.is_available() else torch.float32,
41
- device_map="auto"
42
  )
43
  local_model.eval()
44
  print("✅ Model loaded successfully.")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
45
  except Exception as e:
46
  print(f"❌ Failed to load model: {e}")
47
  traceback.print_exc()
@@ -51,43 +70,49 @@ def create_demo(api_url, model_id=None):
51
  font-family: 'JetBrains Mono', 'Courier New', monospace;
52
  line-height: 1.8;
53
  padding: 24px;
54
- background: #0f172a;
55
  border-radius: 12px;
56
  color: #e2e8f0;
57
  min-height: 250px;
58
  white-space: pre-wrap;
59
  overflow-wrap: break-word;
60
  border: 1px solid #1e293b;
 
 
61
  }
62
  .token {
63
- display: inline;
64
- padding: 1px 2px;
65
- border-radius: 4px;
66
  position: relative;
67
  cursor: help;
68
- transition: background 0.2s;
69
  }
70
  .token:hover {
71
- background: rgba(255, 255, 255, 0.1);
 
72
  }
73
  .token-tooltip {
74
- visibility: hidden;
75
  position: absolute;
76
- bottom: 120%;
77
  left: 50%;
78
  transform: translateX(-50%);
79
  background: #1e293b;
80
  color: white;
81
- padding: 6px 10px;
82
- border-radius: 6px;
83
- font-size: 0.8rem;
84
- z-index: 100;
85
  white-space: nowrap;
86
- box-shadow: 0 4px 12px rgba(0,0,0,0.5);
87
- border: 1px solid #334155;
 
 
 
88
  }
89
  .token:hover .token-tooltip {
90
- visibility: visible;
91
  }
92
  .mode-badge {
93
  display: inline-block;
@@ -99,21 +124,46 @@ def create_demo(api_url, model_id=None):
99
  }
100
  .mode-standalone { background: #059669; color: white; }
101
  .mode-api { background: #4f46e5; color: white; }
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
102
  """
103
 
104
- def generate_html(prompt, max_tokens, temperature, top_k):
105
  try:
106
  if local_model and local_tokenizer:
107
  # Standalone Mode Logic
108
  input_tokens = local_tokenizer(prompt, return_tensors="pt").to(local_model.device)
109
  input_ids = input_tokens.input_ids
110
-
111
  stop_strings = ["<|im_end|>", "<|end|>", "<|user|>"]
112
- stopping_criteria = StoppingCriteriaList([StringStoppingCriteria(local_tokenizer, stop_strings)])
113
-
 
 
114
  all_tokens = []
115
  all_logprobs = []
116
-
117
  # Echo (Prompt) Logprobs
118
  with torch.no_grad():
119
  outputs = local_model(input_ids)
@@ -122,13 +172,13 @@ def create_demo(api_url, model_id=None):
122
  token_id = input_ids[0, i].item()
123
  token_text = local_tokenizer.decode([token_id])
124
  all_tokens.append(token_text)
125
-
126
  if i == 0:
127
  all_logprobs.append(None)
128
  else:
129
- lp = torch.nn.functional.log_softmax(logits[0, i-1, :], dim=-1)
130
  all_logprobs.append(lp[token_id].item())
131
-
132
  # Generation
133
  with torch.no_grad():
134
  gen_out = local_model.generate(
@@ -137,27 +187,27 @@ def create_demo(api_url, model_id=None):
137
  temperature=temperature,
138
  top_k=top_k,
139
  top_p=0.9,
140
- repetition_penalty=1.2,
141
  do_sample=temperature > 0,
142
  use_cache=False,
143
  output_scores=True,
144
  return_dict_in_generate=True,
145
  stopping_criteria=stopping_criteria,
146
  pad_token_id=local_tokenizer.pad_token_id or 32000,
147
- eos_token_id=[32000, 32007, 32011]
148
  )
149
-
150
  output_ids = gen_out.sequences
151
  scores = gen_out.scores
152
- generated_ids = output_ids[0, input_ids.shape[1]:]
153
-
154
  for i, token_id in enumerate(generated_ids):
155
  token_id = token_id.item()
156
  token_text = local_tokenizer.decode([token_id])
157
  all_tokens.append(token_text)
158
  lp = torch.nn.functional.log_softmax(scores[i][0, :], dim=-1)
159
  all_logprobs.append(lp[token_id].item())
160
-
161
  logprobs = {"tokens": all_tokens, "token_logprobs": all_logprobs}
162
  else:
163
  # API Mode
@@ -166,6 +216,7 @@ def create_demo(api_url, model_id=None):
166
  "max_tokens": max_tokens,
167
  "temperature": temperature,
168
  "top_k": top_k,
 
169
  "logprobs": 1,
170
  "echo": True,
171
  }
@@ -190,14 +241,21 @@ def create_demo(api_url, model_id=None):
190
  color = "#e2e8f0" # Default
191
  if lp is not None:
192
  if prob > 0.8:
193
- color = "#4ade80" # Green
194
  elif prob < 0.2:
195
- color = "#f87171" # Red
196
 
197
  prob_pct = f"{(prob * 100):.1f}%" if lp is not None else "N/A"
 
198
 
199
  html += f'<span class="token" style="color: {color}">{display_token}'
200
- html += f'<span class="token-tooltip">Prob: {prob_pct}</span></span>'
 
 
 
 
 
 
201
 
202
  html += "</div>"
203
  return html
@@ -205,22 +263,38 @@ def create_demo(api_url, model_id=None):
205
  traceback.print_exc()
206
  return f'<div class="token-container" style="color: #f87171">❌ Error: {str(e)}</div>'
207
 
208
- with gr.Blocks(title="Echo-DSRN Predictor") as demo:
209
- gr.Markdown(
210
- """
211
- # 🌫️ Echo-DSRN Next Word Prediction
212
- ### Responsible AI Development by ethicalabs.ai
213
- """
214
- )
 
 
 
 
 
 
 
 
 
 
 
 
 
215
 
216
- mode_str = "Standalone Mode" if local_model else f"Connected to API: {api_url}"
217
- mode_cls = "mode-standalone" if local_model else "mode-api"
218
- gr.HTML(f'<div class="mode-badge {mode_cls}">{mode_str}</div>')
 
 
 
 
219
 
220
  gr.Markdown(
221
  """
222
- This application shows the internal probabilities of the model.
223
- Hover over tokens to see their prediction confidence.
224
  """
225
  )
226
 
@@ -238,6 +312,9 @@ def create_demo(api_url, model_id=None):
238
  )
239
  with gr.Row():
240
  top_k = gr.Slider(minimum=0, maximum=100, value=40, step=1, label="Top K")
 
 
 
241
 
242
  submit_btn = gr.Button("Generate with Probabilities", variant="primary")
243
 
@@ -246,16 +323,18 @@ def create_demo(api_url, model_id=None):
246
  output_html = gr.HTML(label="Visualized Output")
247
 
248
  submit_btn.click(
249
- generate_html, inputs=[prompt, max_tokens, temp, top_k], outputs=output_html
 
 
250
  )
251
 
252
  gr.Examples(
253
  examples=[
254
- ["The capital of France is", 128, 0.1, 40],
255
- ["Responsible AI development requires", 128, 0.7, 40],
256
- ["To be or not to be, that is the", 128, 0.1, 40],
257
  ],
258
- inputs=[prompt, max_tokens, temp, top_k],
259
  )
260
 
261
  return demo, css
@@ -267,10 +346,10 @@ if __name__ == "__main__":
267
  "--api_url", type=str, default="http://localhost:5000", help="URL of the Echo-DSRN API"
268
  )
269
  parser.add_argument(
270
- "--model_id",
271
- type=str,
272
- default="ethicalabs/Echo-DSRN-Small-Instruct",
273
- help="Local path or HF ID to the model for Standalone mode"
274
  )
275
  parser.add_argument("--port", type=int, default=7860, help="Port to run the Gradio app on")
276
  args = parser.parse_args()
 
26
  def create_demo(api_url, model_id=None):
27
  # Standalone mode check
28
  hf_model_id = model_id or os.environ.get("HF_MODEL_ID", "ethicalabs/Echo-DSRN-Small-Instruct")
29
+
30
  local_model = None
31
  local_tokenizer = None
32
+ model_metadata = None
33
+
34
  if hf_model_id:
35
  print(f"📦 Standalone Mode: Loading {hf_model_id} via Transformers...")
36
  try:
37
  local_tokenizer = AutoTokenizer.from_pretrained(hf_model_id, trust_remote_code=True)
38
  local_model = AutoModelForCausalLM.from_pretrained(
39
+ hf_model_id,
40
  trust_remote_code=True,
41
  torch_dtype=torch.bfloat16 if torch.cuda.is_available() else torch.float32,
42
+ device_map="auto",
43
  )
44
  local_model.eval()
45
  print("✅ Model loaded successfully.")
46
+
47
+ # Calculate metadata
48
+ total_params = sum(p.numel() for p in local_model.parameters())
49
+ trainable_params = sum(p.numel() for p in local_model.parameters() if p.requires_grad)
50
+ config = local_model.config
51
+
52
+ model_metadata = {
53
+ "Total Parameters": f"{total_params:,}",
54
+ "Trainable Parameters": f"{trainable_params:,}",
55
+ "Layers": getattr(
56
+ config, "num_hidden_layers", getattr(config, "num_layers", "N/A")
57
+ ),
58
+ "Hidden Size": getattr(config, "hidden_size", "N/A"),
59
+ "MLP Ratio": getattr(config, "mlp_ratio", "N/A"),
60
+ "Vocab Size": getattr(config, "vocab_size", "N/A"),
61
+ "Surprise Lambda Init": getattr(config, "surprise_lambda_init", "N/A"),
62
+ "Model Type": getattr(config, "model_type", "echo"),
63
+ }
64
  except Exception as e:
65
  print(f"❌ Failed to load model: {e}")
66
  traceback.print_exc()
 
70
  font-family: 'JetBrains Mono', 'Courier New', monospace;
71
  line-height: 1.8;
72
  padding: 24px;
73
+ background: #020617;
74
  border-radius: 12px;
75
  color: #e2e8f0;
76
  min-height: 250px;
77
  white-space: pre-wrap;
78
  overflow-wrap: break-word;
79
  border: 1px solid #1e293b;
80
+ position: relative;
81
+ overflow: visible !important;
82
  }
83
  .token {
84
+ display: inline-block;
85
+ padding: 0 1px;
86
+ border-radius: 3px;
87
  position: relative;
88
  cursor: help;
89
+ transition: all 0.2s;
90
  }
91
  .token:hover {
92
+ background: rgba(255, 255, 255, 0.15) !important;
93
+ z-index: 100;
94
  }
95
  .token-tooltip {
96
+ display: none;
97
  position: absolute;
98
+ bottom: 130%;
99
  left: 50%;
100
  transform: translateX(-50%);
101
  background: #1e293b;
102
  color: white;
103
+ padding: 12px;
104
+ border-radius: 10px;
105
+ font-size: 0.85rem;
106
+ z-index: 1000;
107
  white-space: nowrap;
108
+ box-shadow: 0 20px 25px -5px rgba(0, 0, 0, 0.5);
109
+ border: 1px solid rgba(255, 255, 255, 0.2);
110
+ pointer-events: none;
111
+ text-align: center;
112
+ min-width: 120px;
113
  }
114
  .token:hover .token-tooltip {
115
+ display: block;
116
  }
117
  .mode-badge {
118
  display: inline-block;
 
124
  }
125
  .mode-standalone { background: #059669; color: white; }
126
  .mode-api { background: #4f46e5; color: white; }
127
+
128
+ .metadata-card {
129
+ background: #0f172a;
130
+ padding: 15px;
131
+ border-radius: 8px;
132
+ border: 1px solid #1e293b;
133
+ font-family: inherit;
134
+ }
135
+ .metadata-grid {
136
+ display: grid;
137
+ grid-template-columns: repeat(auto-fill, minmax(180px, 1fr));
138
+ gap: 10px;
139
+ margin-top: 10px;
140
+ }
141
+ .metadata-item {
142
+ font-size: 0.85rem;
143
+ color: #94a3b8;
144
+ }
145
+ .metadata-value {
146
+ font-weight: 600;
147
+ color: #f1f5f9;
148
+ display: block;
149
+ }
150
  """
151
 
152
+ def generate_html(prompt, max_tokens, temperature, top_k, rep_penalty):
153
  try:
154
  if local_model and local_tokenizer:
155
  # Standalone Mode Logic
156
  input_tokens = local_tokenizer(prompt, return_tensors="pt").to(local_model.device)
157
  input_ids = input_tokens.input_ids
158
+
159
  stop_strings = ["<|im_end|>", "<|end|>", "<|user|>"]
160
+ stopping_criteria = StoppingCriteriaList(
161
+ [StringStoppingCriteria(local_tokenizer, stop_strings)]
162
+ )
163
+
164
  all_tokens = []
165
  all_logprobs = []
166
+
167
  # Echo (Prompt) Logprobs
168
  with torch.no_grad():
169
  outputs = local_model(input_ids)
 
172
  token_id = input_ids[0, i].item()
173
  token_text = local_tokenizer.decode([token_id])
174
  all_tokens.append(token_text)
175
+
176
  if i == 0:
177
  all_logprobs.append(None)
178
  else:
179
+ lp = torch.nn.functional.log_softmax(logits[0, i - 1, :], dim=-1)
180
  all_logprobs.append(lp[token_id].item())
181
+
182
  # Generation
183
  with torch.no_grad():
184
  gen_out = local_model.generate(
 
187
  temperature=temperature,
188
  top_k=top_k,
189
  top_p=0.9,
190
+ repetition_penalty=rep_penalty,
191
  do_sample=temperature > 0,
192
  use_cache=False,
193
  output_scores=True,
194
  return_dict_in_generate=True,
195
  stopping_criteria=stopping_criteria,
196
  pad_token_id=local_tokenizer.pad_token_id or 32000,
197
+ eos_token_id=[32000, 32007, 32011],
198
  )
199
+
200
  output_ids = gen_out.sequences
201
  scores = gen_out.scores
202
+ generated_ids = output_ids[0, input_ids.shape[1] :]
203
+
204
  for i, token_id in enumerate(generated_ids):
205
  token_id = token_id.item()
206
  token_text = local_tokenizer.decode([token_id])
207
  all_tokens.append(token_text)
208
  lp = torch.nn.functional.log_softmax(scores[i][0, :], dim=-1)
209
  all_logprobs.append(lp[token_id].item())
210
+
211
  logprobs = {"tokens": all_tokens, "token_logprobs": all_logprobs}
212
  else:
213
  # API Mode
 
216
  "max_tokens": max_tokens,
217
  "temperature": temperature,
218
  "top_k": top_k,
219
+ "repetition_penalty": rep_penalty,
220
  "logprobs": 1,
221
  "echo": True,
222
  }
 
241
  color = "#e2e8f0" # Default
242
  if lp is not None:
243
  if prob > 0.8:
244
+ color = "#4ade80" # Vibrant Green
245
  elif prob < 0.2:
246
+ color = "#f87171" # Soft Red
247
 
248
  prob_pct = f"{(prob * 100):.1f}%" if lp is not None else "N/A"
249
+ logit_str = f"{lp:.4f}" if lp is not None else "N/A"
250
 
251
  html += f'<span class="token" style="color: {color}">{display_token}'
252
+ html += f'<span class="token-tooltip">'
253
+ html += f'<span style="opacity: 0.7; font-size: 0.7rem">CONFIDENCE</span><br>'
254
+ html += f'<span style="font-size: 1.1rem; font-weight: 800; color: {color}">{prob_pct}</span><br>'
255
+ html += f'<hr style="margin: 5px 0; opacity: 0.1; border: none; border-top: 1px solid white;">'
256
+ html += f'<span style="opacity: 0.7; font-size: 0.7rem">LOGIT</span><br>'
257
+ html += f'<span style="font-family: monospace">{logit_str}</span>'
258
+ html += f'</span></span>'
259
 
260
  html += "</div>"
261
  return html
 
263
  traceback.print_exc()
264
  return f'<div class="token-container" style="color: #f87171">❌ Error: {str(e)}</div>'
265
 
266
+ with gr.Blocks(title="Echo-DSRN-Small-400m - Next Word Prediction") as demo:
267
+ with gr.Row():
268
+ with gr.Column(scale=4):
269
+ gr.Markdown(
270
+ """
271
+ # 🌫️ Echo-DSRN-Small-400m - Next Word Prediction
272
+ ### Responsible AI Development by ethicalabs.ai
273
+
274
+ **Echo-DSRN** is a high-performance **Dual-State Recurrent Network** inspired by the architecture in [Titans: Learning to Memorize at Test Time (Google Research)](https://arxiv.org/abs/2501.00663). It combines the efficiency of linear recurrence with optimized gating and a novel **Surprise Mechanism** to enable infinite context extrapolation.
275
+
276
+ <p style="margin-top: 15px; border-top: 1px solid rgba(255,255,255,0.1); padding-top: 10px;">
277
+ <span style="font-weight: bold; color: #f87171;">⚠️ Proprietary Technology Disclaimer</span><br>
278
+ <span style="font-style: italic; font-size: 0.9rem; color: #94a3b8;">This research is currently in an internal-only alpha phase. The model source code and weights can be released only when we have the capacity to scale up and expand and fine-tune and safety align it according to our strict responsibility standards.</span>
279
+ </p>
280
+ """
281
+ )
282
+ with gr.Column(scale=1):
283
+ mode_str = "Standalone Mode" if local_model else f"Connected to API: {api_url}"
284
+ mode_cls = "mode-standalone" if local_model else "mode-api"
285
+ gr.HTML(f'<div class="mode-badge {mode_cls}">{mode_str}</div>')
286
 
287
+ if model_metadata:
288
+ gr.Markdown("### 📦 Model Architecture Details")
289
+ metadata_html = '<div class="metadata-card"><div class="metadata-grid">'
290
+ for key, value in model_metadata.items():
291
+ metadata_html += f'<div class="metadata-item">{key}<span class="metadata-value">{value}</span></div>'
292
+ metadata_html += "</div></div>"
293
+ gr.HTML(metadata_html)
294
 
295
  gr.Markdown(
296
  """
297
+ Hover over generated tokens to visualize the model's confidence and raw logit values.
 
298
  """
299
  )
300
 
 
312
  )
313
  with gr.Row():
314
  top_k = gr.Slider(minimum=0, maximum=100, value=40, step=1, label="Top K")
315
+ rep_penalty = gr.Slider(
316
+ minimum=1.0, maximum=2.0, value=1.2, step=0.05, label="Repetition Penalty"
317
+ )
318
 
319
  submit_btn = gr.Button("Generate with Probabilities", variant="primary")
320
 
 
323
  output_html = gr.HTML(label="Visualized Output")
324
 
325
  submit_btn.click(
326
+ generate_html,
327
+ inputs=[prompt, max_tokens, temp, top_k, rep_penalty],
328
+ outputs=output_html,
329
  )
330
 
331
  gr.Examples(
332
  examples=[
333
+ ["The capital of France is", 128, 0.1, 40, 1.2],
334
+ ["Responsible AI development requires", 128, 0.7, 40, 1.2],
335
+ ["To be or not to be, that is the", 128, 0.1, 40, 1.2],
336
  ],
337
+ inputs=[prompt, max_tokens, temp, top_k, rep_penalty],
338
  )
339
 
340
  return demo, css
 
346
  "--api_url", type=str, default="http://localhost:5000", help="URL of the Echo-DSRN API"
347
  )
348
  parser.add_argument(
349
+ "--model_id",
350
+ type=str,
351
+ default=None,
352
+ help="Local path or HF ID to the model for Standalone mode",
353
  )
354
  parser.add_argument("--port", type=int, default=7860, help="Port to run the Gradio app on")
355
  args = parser.parse_args()