zakerytclarke commited on
Commit
090e237
·
verified ·
1 Parent(s): f9e63d0

Update src/streamlit_app.py

Browse files
Files changed (1) hide show
  1. src/streamlit_app.py +68 -62
src/streamlit_app.py CHANGED
@@ -1,10 +1,11 @@
1
  import os
2
  import time
3
  import requests
 
4
 
5
  import streamlit as st
6
  import torch
7
- from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
8
 
9
  # Optional LangSmith (trace + feedback)
10
  try:
@@ -105,7 +106,6 @@ if "messages" not in st.session_state:
105
  if "seeded" not in st.session_state:
106
  st.session_state.seeded = False
107
 
108
- # Seed exactly once on first load
109
  if (not st.session_state.seeded) and (len(st.session_state.messages) == 0):
110
  st.session_state.messages = [SAMPLE_USER_MSG, SAMPLE_ASSISTANT_MSG]
111
  st.session_state.seeded = True
@@ -204,39 +204,6 @@ def count_tokens(text: str) -> int:
204
  return len(tokenizer.encode(text)) if text else 0
205
 
206
 
207
- # =========================
208
- # LANGSMITH-TRACED ANSWER FUNCTION
209
- # =========================
210
- if traceable:
211
-
212
- @traceable(name="teapot_answer")
213
- def traced_answer(context: str, system_prompt: str, question: str) -> str:
214
- prompt = f"{context}\n{system_prompt}\n{question}\n"
215
- inputs = tokenizer(prompt, return_tensors="pt").to(device)
216
- with torch.no_grad():
217
- out = model.generate(
218
- **inputs,
219
- max_new_tokens=MAX_NEW_TOKENS,
220
- do_sample=False,
221
- num_beams=1,
222
- )
223
- return tokenizer.decode(out[0], skip_special_tokens=True)
224
-
225
- else:
226
-
227
- def traced_answer(context: str, system_prompt: str, question: str) -> str:
228
- prompt = f"{context}\n{system_prompt}\n{question}\n"
229
- inputs = tokenizer(prompt, return_tensors="pt").to(device)
230
- with torch.no_grad():
231
- out = model.generate(
232
- **inputs,
233
- max_new_tokens=MAX_NEW_TOKENS,
234
- do_sample=False,
235
- num_beams=1,
236
- )
237
- return tokenizer.decode(out[0], skip_special_tokens=True)
238
-
239
-
240
  def get_trace_id_if_available() -> str | None:
241
  if not get_current_run_tree:
242
  return None
@@ -247,6 +214,51 @@ def get_trace_id_if_available() -> str | None:
247
  return None
248
 
249
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
250
  # =========================
251
  # FEEDBACK HANDLER (attached to trace_id)
252
  # =========================
@@ -257,7 +269,6 @@ def handle_feedback(idx: int):
257
  msg = st.session_state.messages[idx]
258
  trace_id = msg.get("trace_id")
259
 
260
- # Attach feedback to this traced run
261
  if ls_client and trace_id:
262
  score = 1 if val == "👍" else 0
263
  try:
@@ -272,9 +283,16 @@ def handle_feedback(idx: int):
272
 
273
 
274
  # =========================
275
- # RENDER HISTORY
276
- # Row 1: message + feedback
277
- # Row 2: inspect + debug metrics
 
 
 
 
 
 
 
278
  # =========================
279
  for i, msg in enumerate(st.session_state.messages):
280
  with st.chat_message(msg["role"]):
@@ -299,7 +317,6 @@ for i, msg in enumerate(st.session_state.messages):
299
 
300
  # Row 2
301
  inspect_col, metrics_col = st.columns([12, 1], vertical_alignment="center")
302
-
303
  with inspect_col:
304
  st.caption(
305
  f"🔎 {msg.get('search_time', 0.0):.2f}s (search) "
@@ -307,7 +324,6 @@ for i, msg in enumerate(st.session_state.messages):
307
  f"⚡ {msg.get('tps', 0.0):.1f} tok/s "
308
  f"🧾 {msg.get('input_tokens', 0)} input tokens • {msg.get('output_tokens', 0)} output tokens"
309
  )
310
-
311
  with metrics_col:
312
  with st.popover("ℹ️", help="Inspect"):
313
  st.markdown("**Context**")
@@ -319,15 +335,10 @@ for i, msg in enumerate(st.session_state.messages):
319
 
320
 
321
  # =========================
322
- # INPUT + GENERATE (NO RERUN / NO FLASH)
 
323
  # =========================
324
- query = st.chat_input("Ask a question...")
325
-
326
  if query:
327
- # Persist user message
328
- st.session_state.messages.append({"role": "user", "content": query})
329
-
330
- # Generate immediately in the same run (no st.rerun)
331
  question = query
332
 
333
  # Web search
@@ -338,9 +349,8 @@ if query:
338
  prompt = f"{context}\n{system_prompt}\n{question}\n"
339
  input_tokens = count_tokens(prompt)
340
 
341
- # Assistant response
342
  with st.chat_message("assistant"):
343
- # Row 1: message + feedback (feedback disabled until persisted)
344
  msg_col, fb_col = st.columns([14, 1], vertical_alignment="center")
345
  with msg_col:
346
  placeholder = st.empty()
@@ -348,20 +358,16 @@ if query:
348
  st.feedback("thumbs", key="live_fb", disabled=True)
349
 
350
  start = time.perf_counter()
351
- answer = traced_answer(context, system_prompt, question)
352
- trace_id = get_trace_id_if_available()
353
 
354
- # Typewriter render (reduce updates to avoid jitter)
355
  buf = ""
356
- for j, ch in enumerate(answer, 1):
357
- buf += ch
358
- if j % 6 == 0: # update every 6 chars
359
- placeholder.markdown(buf)
360
- time.sleep(0.001)
361
- placeholder.markdown(buf)
362
 
 
363
  gen_time = time.perf_counter() - start
364
- output_tokens = count_tokens(answer)
365
  tps = output_tokens / gen_time if gen_time > 0 else 0.0
366
 
367
  # Row 2: inspect + metrics
@@ -384,11 +390,11 @@ if query:
384
  st.markdown("**Prompt**")
385
  st.code(prompt, language="text")
386
 
387
- # Persist assistant message (so it shows on subsequent runs)
388
  st.session_state.messages.append(
389
  {
390
  "role": "assistant",
391
- "content": answer,
392
  "context": context,
393
  "system_prompt": system_prompt,
394
  "question": question,
 
1
  import os
2
  import time
3
  import requests
4
+ import threading
5
 
6
  import streamlit as st
7
  import torch
8
+ from transformers import AutoTokenizer, AutoModelForSeq2SeqLM, TextIteratorStreamer
9
 
10
  # Optional LangSmith (trace + feedback)
11
  try:
 
106
  if "seeded" not in st.session_state:
107
  st.session_state.seeded = False
108
 
 
109
  if (not st.session_state.seeded) and (len(st.session_state.messages) == 0):
110
  st.session_state.messages = [SAMPLE_USER_MSG, SAMPLE_ASSISTANT_MSG]
111
  st.session_state.seeded = True
 
204
  return len(tokenizer.encode(text)) if text else 0
205
 
206
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
207
  def get_trace_id_if_available() -> str | None:
208
  if not get_current_run_tree:
209
  return None
 
214
  return None
215
 
216
 
217
+ # =========================
218
+ # TRACED STREAMING GENERATION
219
+ # =========================
220
+ def stream_answer_tokens(context: str, system_prompt: str, question: str):
221
+ """
222
+ Yields decoded text increments as they are generated (real streaming).
223
+ """
224
+ prompt = f"{context}\n{system_prompt}\n{question}\n"
225
+ inputs = tokenizer(prompt, return_tensors="pt").to(device)
226
+
227
+ streamer = TextIteratorStreamer(
228
+ tokenizer,
229
+ skip_prompt=True,
230
+ skip_special_tokens=True,
231
+ )
232
+
233
+ gen_kwargs = dict(
234
+ **inputs,
235
+ max_new_tokens=MAX_NEW_TOKENS,
236
+ do_sample=False,
237
+ num_beams=1,
238
+ streamer=streamer,
239
+ )
240
+
241
+ # Run generate in a background thread so we can iterate streamer in the main thread.
242
+ t = threading.Thread(target=model.generate, kwargs=gen_kwargs, daemon=True)
243
+ t.start()
244
+
245
+ for text in streamer:
246
+ # text can be tiny chunks; yield as-is
247
+ yield text
248
+
249
+
250
+ if traceable:
251
+ # Wrap a traced function around the streaming loop (LangSmith will see a single run)
252
+ @traceable(name="teapot_answer_stream")
253
+ def traced_stream_answer(context: str, system_prompt: str, question: str):
254
+ for chunk in stream_answer_tokens(context, system_prompt, question):
255
+ yield chunk
256
+ else:
257
+ def traced_stream_answer(context: str, system_prompt: str, question: str):
258
+ for chunk in stream_answer_tokens(context, system_prompt, question):
259
+ yield chunk
260
+
261
+
262
  # =========================
263
  # FEEDBACK HANDLER (attached to trace_id)
264
  # =========================
 
269
  msg = st.session_state.messages[idx]
270
  trace_id = msg.get("trace_id")
271
 
 
272
  if ls_client and trace_id:
273
  score = 1 if val == "👍" else 0
274
  try:
 
283
 
284
 
285
  # =========================
286
+ # INPUT FIRST (so new user msg renders immediately)
287
+ # =========================
288
+ query = st.chat_input("Ask a question...")
289
+
290
+ if query:
291
+ st.session_state.messages.append({"role": "user", "content": query})
292
+
293
+
294
+ # =========================
295
+ # RENDER HISTORY (now includes latest user msg)
296
  # =========================
297
  for i, msg in enumerate(st.session_state.messages):
298
  with st.chat_message(msg["role"]):
 
317
 
318
  # Row 2
319
  inspect_col, metrics_col = st.columns([12, 1], vertical_alignment="center")
 
320
  with inspect_col:
321
  st.caption(
322
  f"🔎 {msg.get('search_time', 0.0):.2f}s (search) "
 
324
  f"⚡ {msg.get('tps', 0.0):.1f} tok/s "
325
  f"🧾 {msg.get('input_tokens', 0)} input tokens • {msg.get('output_tokens', 0)} output tokens"
326
  )
 
327
  with metrics_col:
328
  with st.popover("ℹ️", help="Inspect"):
329
  st.markdown("**Context**")
 
335
 
336
 
337
  # =========================
338
+ # GENERATE ONLY IF THIS RUN RECEIVED A NEW QUERY
339
+ # (We detect by: query is not None)
340
  # =========================
 
 
341
  if query:
 
 
 
 
342
  question = query
343
 
344
  # Web search
 
349
  prompt = f"{context}\n{system_prompt}\n{question}\n"
350
  input_tokens = count_tokens(prompt)
351
 
352
+ # Stream assistant response (real streaming)
353
  with st.chat_message("assistant"):
 
354
  msg_col, fb_col = st.columns([14, 1], vertical_alignment="center")
355
  with msg_col:
356
  placeholder = st.empty()
 
358
  st.feedback("thumbs", key="live_fb", disabled=True)
359
 
360
  start = time.perf_counter()
 
 
361
 
 
362
  buf = ""
363
+ placeholder.markdown("") # ensures first token updates a visible element
364
+ for chunk in traced_stream_answer(context, system_prompt, question):
365
+ buf += chunk
366
+ placeholder.markdown(buf)
 
 
367
 
368
+ trace_id = get_trace_id_if_available()
369
  gen_time = time.perf_counter() - start
370
+ output_tokens = count_tokens(buf)
371
  tps = output_tokens / gen_time if gen_time > 0 else 0.0
372
 
373
  # Row 2: inspect + metrics
 
390
  st.markdown("**Prompt**")
391
  st.code(prompt, language="text")
392
 
393
+ # Persist assistant message
394
  st.session_state.messages.append(
395
  {
396
  "role": "assistant",
397
+ "content": buf,
398
  "context": context,
399
  "system_prompt": system_prompt,
400
  "question": question,