zakerytclarke commited on
Commit
42556c6
·
verified ·
1 Parent(s): 78a9bf0

Update src/streamlit_app.py

Browse files
Files changed (1) hide show
  1. src/streamlit_app.py +174 -216
src/streamlit_app.py CHANGED
@@ -1,180 +1,167 @@
1
  import os
2
- import re
3
  import time
4
  import threading
5
- from typing import List, Dict, Optional
6
-
7
  import streamlit as st
8
  import torch
9
- import requests
10
-
11
  from transformers import AutoTokenizer, AutoModelForSeq2SeqLM, TextIteratorStreamer
12
 
13
- # Fast RAG
14
- from sklearn.feature_extraction.text import TfidfVectorizer
15
- import numpy as np
16
-
17
- # File parsing
18
- import io
19
- try:
20
- import docx
21
- except:
22
- docx = None
23
- try:
24
- import PyPDF2
25
- except:
26
- PyPDF2 = None
27
-
28
- # LangSmith
29
  try:
30
  from langsmith import Client as LangSmithClient
31
  except:
32
  LangSmithClient = None
33
 
34
 
35
- # -----------------------
36
  # CONFIG
37
- # -----------------------
38
  MODEL_NAME = "teapotai/tinyteapot"
39
  MAX_INPUT_TOKENS = 512
40
  MAX_NEW_TOKENS = 192
41
  TOP_K_SEARCH = 3
42
- TOP_K_RAG = 3
43
-
44
- st.set_page_config(page_title="TeapotAI Chat", page_icon="🫖", layout="centered")
45
 
 
 
 
 
 
46
 
47
- # -----------------------
48
- # MODEL LOADING (CACHED)
49
- # -----------------------
50
  @st.cache_resource
51
  def load_model():
52
  tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
53
  model = AutoModelForSeq2SeqLM.from_pretrained(MODEL_NAME)
54
  device = "cuda" if torch.cuda.is_available() else "cpu"
55
- model.to(device).eval()
 
56
  return tokenizer, model, device
57
 
58
-
59
  tokenizer, model, device = load_model()
60
 
61
 
62
- # -----------------------
63
  # LANGSMITH
64
- # -----------------------
65
  @st.cache_resource
66
  def get_langsmith():
67
- if LangSmithClient and (os.getenv("LANGCHAIN_API_KEY") or os.getenv("LANGSMITH_API_KEY")):
 
68
  return LangSmithClient()
69
  return None
70
 
71
  ls_client = get_langsmith()
72
 
73
 
74
- # -----------------------
75
- # FAST FILE → TEXT
76
- # -----------------------
77
- def extract_text(file) -> str:
78
- name = file.name.lower()
 
 
 
 
 
 
 
 
 
 
 
79
 
80
- if name.endswith(".txt") or name.endswith(".md"):
81
- return file.read().decode("utf-8", errors="ignore")
 
 
 
82
 
83
- if name.endswith(".pdf") and PyPDF2:
84
- reader = PyPDF2.PdfReader(file)
85
- return "\n".join(page.extract_text() or "" for page in reader.pages)
 
 
 
 
 
 
 
 
86
 
87
- if name.endswith(".docx") and docx:
88
- doc = docx.Document(file)
89
- return "\n".join(p.text for p in doc.paragraphs)
 
 
 
90
 
91
- return ""
92
 
93
 
94
- def chunk_by_paragraph(text: str) -> List[str]:
95
- chunks = [p.strip() for p in re.split(r"\n{2,}|\n", text) if len(p.strip()) > 30]
96
- return chunks[:2000] # safety cap
 
 
 
 
97
 
 
 
 
 
98
 
99
- # -----------------------
100
- # FAST TFIDF RAG (CACHED)
101
- # -----------------------
102
- @st.cache_resource
103
- def build_tfidf(chunks: List[str]):
104
- if not chunks:
105
- return None, None
106
- vectorizer = TfidfVectorizer(stop_words="english", max_features=20000)
107
- matrix = vectorizer.fit_transform(chunks)
108
- return vectorizer, matrix
109
-
110
-
111
- def retrieve_top_chunks(query: str, chunks: List[str], vectorizer, matrix, k=3):
112
- if not chunks or vectorizer is None:
113
- return []
114
- q_vec = vectorizer.transform([query])
115
- scores = (matrix @ q_vec.T).toarray().ravel()
116
- top_idx = np.argsort(scores)[-k:][::-1]
117
- return [chunks[i] for i in top_idx if scores[i] > 0]
118
-
119
-
120
- # -----------------------
121
- # SEARCH (FAST)
122
- # -----------------------
123
- def web_search(query):
124
- key = os.getenv("BRAVE_API_KEY") or st.secrets.get("BRAVE_API_KEY", None)
125
- if not key:
126
- return [], ""
127
-
128
- headers = {"X-Subscription-Token": key, "Accept": "application/json"}
129
  params = {"q": query, "count": TOP_K_SEARCH}
130
 
131
  t0 = time.perf_counter()
132
- r = requests.get(
133
- "https://api.search.brave.com/res/v1/web/search",
134
- headers=headers,
135
- params=params,
136
- timeout=8,
137
- )
138
- data = r.json()
 
 
 
139
  t1 = time.perf_counter()
140
 
141
- results = []
142
- ctx_blocks = []
143
-
144
  for i, item in enumerate(data.get("web", {}).get("results", [])[:TOP_K_SEARCH], 1):
145
  title = item.get("title", "")
146
  url = item.get("url", "")
147
- desc = item.get("description", "").replace("<strong>", "").replace("</strong>", "")
148
-
149
- results.append({"title": title, "url": url, "snippet": desc})
150
- ctx_blocks.append(f"[{i}] {title}\nURL: {url}\nSnippet: {desc}")
151
-
152
- return results, "\n\n".join(ctx_blocks), (t1 - t0)
153
 
 
 
154
 
155
- # -----------------------
156
- # PROMPT + TRUNCATION
157
- # -----------------------
158
- def build_prompt(context, system, question):
159
- return f"{context}\n{system}\n{question}\n"
160
 
161
-
162
- def truncate_context(context, system, question):
163
- base = build_prompt("", system, question)
164
- base_tokens = tokenizer.encode(base)
 
 
165
  budget = MAX_INPUT_TOKENS - len(base_tokens)
166
 
167
  ctx_tokens = tokenizer.encode(context)
168
  if len(ctx_tokens) <= budget:
169
  return context
170
 
171
- return tokenizer.decode(ctx_tokens[-budget:], skip_special_tokens=True)
 
 
172
 
173
 
174
- # -----------------------
175
  # STREAM GENERATION
176
- # -----------------------
177
- def stream_generate(prompt):
178
  inputs = tokenizer(prompt, return_tensors="pt").to(device)
179
  streamer = TextIteratorStreamer(tokenizer, skip_special_tokens=True)
180
 
@@ -183,6 +170,7 @@ def stream_generate(prompt):
183
  **inputs,
184
  max_new_tokens=MAX_NEW_TOKENS,
185
  do_sample=False,
 
186
  streamer=streamer,
187
  )
188
 
@@ -195,70 +183,34 @@ def stream_generate(prompt):
195
  yield text
196
 
197
 
198
- # -----------------------
199
- # SESSION STATE
200
- # -----------------------
201
- if "messages" not in st.session_state:
202
- st.session_state.messages = []
203
-
204
- if "rag_chunks" not in st.session_state:
205
- st.session_state.rag_chunks = []
206
 
207
- if "vectorizer" not in st.session_state:
208
- st.session_state.vectorizer = None
209
 
210
- if "matrix" not in st.session_state:
211
- st.session_state.matrix = None
212
 
213
-
214
- # -----------------------
215
- # HEADER
216
- # -----------------------
217
- st.markdown("## 🫖 TeapotAI Chat")
218
- st.caption("Fast, grounded answers with hybrid RAG")
219
-
220
-
221
- # -----------------------
222
- # SETTINGS SIDEBAR (UPGRADED)
223
- # -----------------------
224
- with st.sidebar:
225
- st.markdown("### Settings")
226
-
227
- system_prompt = st.text_area(
228
- "System Prompt",
229
- value=(
230
- "You are Teapot, an open-source AI assistant optimized for low-end devices. "
231
- "Answer using the provided context only. "
232
- "If the context does not answer the question, reply exactly: "
233
- "'I am sorry but I don't have any information on that'."
234
- ),
235
- height=150,
236
- )
237
-
238
- st.markdown("### Custom Context (RAG)")
239
- pasted = st.text_area("Paste context text (optional)", height=150)
240
- uploaded = st.file_uploader("Or upload file (.txt, .pdf, .docx, .md)", type=["txt", "pdf", "docx", "md"])
241
-
242
-
243
- # Build RAG index (VERY FAST)
244
- combined_text = ""
245
- if pasted:
246
- combined_text += pasted + "\n"
247
- if uploaded:
248
- combined_text += extract_text(uploaded)
249
-
250
- if combined_text:
251
- chunks = chunk_by_paragraph(combined_text)
252
- vec, mat = build_tfidf(chunks)
253
- st.session_state.rag_chunks = chunks
254
- st.session_state.vectorizer = vec
255
- st.session_state.matrix = mat
256
- st.success(f"Indexed {len(chunks)} context chunks")
257
 
258
 
259
- # -----------------------
260
- # CHAT HISTORY
261
- # -----------------------
262
  for i, msg in enumerate(st.session_state.messages):
263
  with st.chat_message(msg["role"]):
264
  st.markdown(msg["content"])
@@ -271,67 +223,70 @@ for i, msg in enumerate(st.session_state.messages):
271
  f"🧮 {msg['tokens']} tokens"
272
  )
273
 
274
- # Inline clean thumbs (better UX)
275
- col1, col2, _ = st.columns([1, 1, 6])
276
- if col1.button("👍", key=f"up_{i}", disabled=msg.get("feedback") is not None):
277
- msg["feedback"] = 1
278
- if ls_client and msg.get("run_id"):
279
- ls_client.create_feedback(msg["run_id"], key="feedback", score=1.0)
280
- if col2.button("👎", key=f"down_{i}", disabled=msg.get("feedback") is not None):
281
- msg["feedback"] = -1
282
- if ls_client and msg.get("run_id"):
283
- ls_client.create_feedback(msg["run_id"], key="feedback", score=-1.0)
284
-
285
-
286
- # -----------------------
287
- # INPUT
288
- # -----------------------
289
  query = st.chat_input("Ask a question...")
290
 
291
  if query:
292
  st.session_state.messages.append({"role": "user", "content": query})
293
 
294
- # 1️⃣ Web Search (timed)
295
- results, web_context, search_time = web_search(query)
296
-
297
- # 2️⃣ TFIDF RAG (VERY FAST ~1-3ms)
298
- rag_context = ""
299
- if st.session_state.vectorizer:
300
- top_chunks = retrieve_top_chunks(
301
- query,
302
- st.session_state.rag_chunks,
303
- st.session_state.vectorizer,
304
- st.session_state.matrix,
305
- k=TOP_K_RAG,
306
- )
307
- rag_context = "\n\n".join(top_chunks)
 
 
 
 
308
 
309
- # 3️⃣ Hybrid Context
310
- full_context = (rag_context + "\n\n" + web_context).strip()
311
- truncated_context = truncate_context(full_context, system_prompt, query)
312
- prompt = build_prompt(truncated_context, system_prompt, query)
313
 
314
- # LangSmith Run
315
  run_id = None
316
  if ls_client:
317
- run = ls_client.create_run(
318
- name="teapot_chat",
319
- run_type="llm",
320
- inputs={
321
- "context": truncated_context,
322
- "system_prompt": system_prompt,
323
- "question": query,
324
- },
325
- )
326
- run_id = run.id
 
 
 
327
 
328
- # 4️⃣ Stream Generation (timed)
329
  with st.chat_message("assistant"):
330
  placeholder = st.empty()
331
 
332
  gen_start = time.perf_counter()
333
  final_text = ""
334
-
335
  for partial in stream_generate(prompt):
336
  final_text = partial
337
  placeholder.markdown(final_text)
@@ -348,7 +303,10 @@ if query:
348
  )
349
 
350
  if ls_client and run_id:
351
- ls_client.update_run(run_id, outputs={"answer": final_text})
 
 
 
352
 
353
  st.session_state.messages.append(
354
  {
 
1
  import os
 
2
  import time
3
  import threading
4
+ import requests
 
5
  import streamlit as st
6
  import torch
 
 
7
  from transformers import AutoTokenizer, AutoModelForSeq2SeqLM, TextIteratorStreamer
8
 
9
+ # Optional LangSmith
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
10
  try:
11
  from langsmith import Client as LangSmithClient
12
  except:
13
  LangSmithClient = None
14
 
15
 
16
+ # =========================
17
  # CONFIG
18
+ # =========================
19
  MODEL_NAME = "teapotai/tinyteapot"
20
  MAX_INPUT_TOKENS = 512
21
  MAX_NEW_TOKENS = 192
22
  TOP_K_SEARCH = 3
23
+ LOGO_URL = "https://teapotai.com/assets/logo.gif"
 
 
24
 
25
+ st.set_page_config(
26
+ page_title="TeapotAI Chat",
27
+ page_icon="🫖",
28
+ layout="centered"
29
+ )
30
 
31
+ # =========================
32
+ # LOAD MODEL (CACHED)
33
+ # =========================
34
  @st.cache_resource
35
  def load_model():
36
  tokenizer = AutoTokenizer.from_pretrained(MODEL_NAME)
37
  model = AutoModelForSeq2SeqLM.from_pretrained(MODEL_NAME)
38
  device = "cuda" if torch.cuda.is_available() else "cpu"
39
+ model.to(device)
40
+ model.eval()
41
  return tokenizer, model, device
42
 
 
43
  tokenizer, model, device = load_model()
44
 
45
 
46
+ # =========================
47
  # LANGSMITH
48
+ # =========================
49
  @st.cache_resource
50
  def get_langsmith():
51
+ api_key = os.getenv("LANGCHAIN_API_KEY") or os.getenv("LANGSMITH_API_KEY")
52
+ if api_key and LangSmithClient:
53
  return LangSmithClient()
54
  return None
55
 
56
  ls_client = get_langsmith()
57
 
58
 
59
+ # =========================
60
+ # SESSION STATE
61
+ # =========================
62
+ if "messages" not in st.session_state:
63
+ st.session_state.messages = []
64
+
65
+ # =========================
66
+ # HEADER (LOGO RESTORED)
67
+ # =========================
68
+ col1, col2 = st.columns([1, 6])
69
+ with col1:
70
+ st.image(LOGO_URL, use_column_width=True)
71
+ with col2:
72
+ st.markdown("## TeapotAI Chat")
73
+ st.caption("Fast, grounded answers with web context")
74
+
75
 
76
+ # =========================
77
+ # SIDEBAR SETTINGS
78
+ # =========================
79
+ with st.sidebar:
80
+ st.markdown("### Settings")
81
 
82
+ system_prompt = st.text_area(
83
+ "System Prompt",
84
+ value=(
85
+ "You are Teapot, an open-source AI assistant optimized for low-end devices, "
86
+ "providing short, accurate responses without hallucinating while excelling at "
87
+ "information extraction and text summarization. "
88
+ "If the context does not answer the question, reply exactly: "
89
+ "'I am sorry but I don't have any information on that'."
90
+ ),
91
+ height=180
92
+ )
93
 
94
+ st.markdown("### Extra Context (Optional)")
95
+ user_context = st.text_area(
96
+ "Paste context to append to web results",
97
+ height=150,
98
+ placeholder="Add any custom context here..."
99
+ )
100
 
101
+ use_web = st.checkbox("Use web search", value=True)
102
 
103
 
104
+ # =========================
105
+ # WEB SEARCH (FAST)
106
+ # =========================
107
+ def web_search(query: str):
108
+ api_key = os.getenv("BRAVE_API_KEY") or st.secrets.get("BRAVE_API_KEY", None)
109
+ if not api_key:
110
+ return "", 0.0
111
 
112
+ headers = {
113
+ "X-Subscription-Token": api_key,
114
+ "Accept": "application/json"
115
+ }
116
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
117
  params = {"q": query, "count": TOP_K_SEARCH}
118
 
119
  t0 = time.perf_counter()
120
+ try:
121
+ r = requests.get(
122
+ "https://api.search.brave.com/res/v1/web/search",
123
+ headers=headers,
124
+ params=params,
125
+ timeout=6,
126
+ )
127
+ data = r.json()
128
+ except:
129
+ return "", 0.0
130
  t1 = time.perf_counter()
131
 
132
+ blocks = []
 
 
133
  for i, item in enumerate(data.get("web", {}).get("results", [])[:TOP_K_SEARCH], 1):
134
  title = item.get("title", "")
135
  url = item.get("url", "")
136
+ desc = item.get("description", "")
137
+ desc = desc.replace("<strong>", "").replace("</strong>", "")
138
+ blocks.append(f"[{i}] {title}\nURL: {url}\nSnippet: {desc}")
 
 
 
139
 
140
+ context = "\n\n".join(blocks)
141
+ return context, (t1 - t0)
142
 
 
 
 
 
 
143
 
144
+ # =========================
145
+ # TRUNCATE TO LAST 512 TOKENS
146
+ # =========================
147
+ def truncate_to_512(context: str, system: str, question: str):
148
+ base_prompt = f"\n{system}\n{question}\n"
149
+ base_tokens = tokenizer.encode(base_prompt)
150
  budget = MAX_INPUT_TOKENS - len(base_tokens)
151
 
152
  ctx_tokens = tokenizer.encode(context)
153
  if len(ctx_tokens) <= budget:
154
  return context
155
 
156
+ # Keep MOST RECENT tokens (tail truncation)
157
+ truncated = ctx_tokens[-budget:]
158
+ return tokenizer.decode(truncated, skip_special_tokens=True)
159
 
160
 
161
+ # =========================
162
  # STREAM GENERATION
163
+ # =========================
164
+ def stream_generate(prompt: str):
165
  inputs = tokenizer(prompt, return_tensors="pt").to(device)
166
  streamer = TextIteratorStreamer(tokenizer, skip_special_tokens=True)
167
 
 
170
  **inputs,
171
  max_new_tokens=MAX_NEW_TOKENS,
172
  do_sample=False,
173
+ num_beams=1,
174
  streamer=streamer,
175
  )
176
 
 
183
  yield text
184
 
185
 
186
+ # =========================
187
+ # LANGSMITH FEEDBACK HANDLER
188
+ # =========================
189
+ def handle_feedback(idx: int):
190
+ val = st.session_state[f"feedback_{idx}"]
191
+ msg = st.session_state.messages[idx]
 
 
192
 
193
+ if val is None:
194
+ return
195
 
196
+ msg["feedback"] = val
 
197
 
198
+ if ls_client and msg.get("run_id"):
199
+ score = 1 if val == "👍" else 0
200
+ try:
201
+ ls_client.create_feedback(
202
+ run_id=msg["run_id"],
203
+ key="thumb_rating",
204
+ score=score,
205
+ comment="thumbs_up" if score else "thumbs_down",
206
+ )
207
+ except Exception as e:
208
+ print("LangSmith feedback error:", e)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
209
 
210
 
211
+ # =========================
212
+ # RENDER CHAT
213
+ # =========================
214
  for i, msg in enumerate(st.session_state.messages):
215
  with st.chat_message(msg["role"]):
216
  st.markdown(msg["content"])
 
223
  f"🧮 {msg['tokens']} tokens"
224
  )
225
 
226
+ feedback_key = f"feedback_{i}"
227
+ st.session_state.setdefault(feedback_key, msg.get("feedback"))
228
+
229
+ st.feedback(
230
+ "thumbs",
231
+ key=feedback_key,
232
+ disabled=msg.get("feedback") is not None,
233
+ on_change=handle_feedback,
234
+ args=(i,),
235
+ )
236
+
237
+
238
+ # =========================
239
+ # CHAT INPUT
240
+ # =========================
241
  query = st.chat_input("Ask a question...")
242
 
243
  if query:
244
  st.session_state.messages.append({"role": "user", "content": query})
245
 
246
+ # ---- WEB SEARCH ----
247
+ web_context = ""
248
+ search_time = 0.0
249
+ if use_web:
250
+ web_context, search_time = web_search(query)
251
+
252
+ # ---- COMBINED CONTEXT (WEB + USER BOX) ----
253
+ combined_context = ""
254
+ if user_context:
255
+ combined_context += user_context.strip() + "\n\n"
256
+ if web_context:
257
+ combined_context += web_context
258
+
259
+ truncated_context = truncate_to_512(
260
+ combined_context,
261
+ system_prompt,
262
+ query
263
+ )
264
 
265
+ prompt = f"{truncated_context}\n{system_prompt}\n{query}\n"
 
 
 
266
 
267
+ # ---- LANGSMITH RUN ----
268
  run_id = None
269
  if ls_client:
270
+ try:
271
+ run = ls_client.create_run(
272
+ name="teapot_chat",
273
+ run_type="llm",
274
+ inputs={
275
+ "context": truncated_context,
276
+ "system_prompt": system_prompt,
277
+ "question": query,
278
+ },
279
+ )
280
+ run_id = run.id
281
+ except:
282
+ pass
283
 
284
+ # ---- STREAM OUTPUT ----
285
  with st.chat_message("assistant"):
286
  placeholder = st.empty()
287
 
288
  gen_start = time.perf_counter()
289
  final_text = ""
 
290
  for partial in stream_generate(prompt):
291
  final_text = partial
292
  placeholder.markdown(final_text)
 
303
  )
304
 
305
  if ls_client and run_id:
306
+ try:
307
+ ls_client.update_run(run_id, outputs={"answer": final_text})
308
+ except:
309
+ pass
310
 
311
  st.session_state.messages.append(
312
  {