Spaces:
Sleeping
Sleeping
Update app.py
Browse files
app.py
CHANGED
|
@@ -1,72 +1,90 @@
|
|
| 1 |
# app.py
|
| 2 |
-
#
|
| 3 |
-
|
| 4 |
import streamlit as st
|
| 5 |
-
from transformers import AutoTokenizer, AutoModelForCausalLM,
|
|
|
|
| 6 |
import torch
|
| 7 |
|
| 8 |
-
st.set_page_config(page_title="PromptPilot", layout="centered")
|
| 9 |
-
st.title("π€ PromptPilot")
|
| 10 |
|
| 11 |
-
st.markdown("
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 12 |
|
| 13 |
-
# Track mode with session state
|
| 14 |
if 'mode' not in st.session_state:
|
| 15 |
st.session_state.mode = 'Code Generator'
|
| 16 |
|
| 17 |
-
|
| 18 |
-
|
| 19 |
-
|
| 20 |
-
|
| 21 |
-
|
| 22 |
-
|
| 23 |
-
|
| 24 |
-
|
| 25 |
-
|
| 26 |
-
if
|
| 27 |
-
st.session_state.mode =
|
| 28 |
|
| 29 |
st.markdown(f"### π Current Mode: **{st.session_state.mode}**")
|
|
|
|
| 30 |
|
| 31 |
-
# Text input
|
| 32 |
-
user_input = st.text_input("Enter your question or prompt below:")
|
| 33 |
-
|
| 34 |
-
# Load models
|
| 35 |
@st.cache_resource
|
| 36 |
def load_codegen():
|
| 37 |
-
|
| 38 |
-
|
| 39 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 40 |
|
| 41 |
@st.cache_resource
|
| 42 |
-
def
|
| 43 |
-
|
| 44 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 45 |
return tokenizer, model
|
| 46 |
|
| 47 |
-
|
| 48 |
-
if user_input and st.session_state.mode:
|
| 49 |
with st.spinner("Generating response..."):
|
| 50 |
-
|
| 51 |
-
|
| 52 |
-
|
| 53 |
-
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
st.code(
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 57 |
|
| 58 |
-
elif
|
| 59 |
-
|
| 60 |
-
|
| 61 |
-
|
| 62 |
-
|
| 63 |
-
st.subheader("π Answer:")
|
| 64 |
-
st.write(result)
|
| 65 |
|
| 66 |
-
elif
|
| 67 |
-
|
| 68 |
-
|
| 69 |
-
|
| 70 |
-
|
| 71 |
-
st.subheader("π Bayesian-style Answer (Sampled):")
|
| 72 |
-
st.write(result)
|
|
|
|
| 1 |
# app.py
|
| 2 |
+
# PromptPilot: Multi-Model AI Chatbot for AML-3304
|
|
|
|
| 3 |
import streamlit as st
|
| 4 |
+
from transformers import (AutoTokenizer, AutoModelForCausalLM,
|
| 5 |
+
AutoModelForSeq2SeqLM, RagTokenizer, RagRetriever, RagSequenceForGeneration)
|
| 6 |
import torch
|
| 7 |
|
| 8 |
+
st.set_page_config(page_title="PromptPilot: Unified AI Chatbot", layout="centered")
|
| 9 |
+
st.title("π€ PromptPilot: Unified AI Chatbot")
|
| 10 |
|
| 11 |
+
st.markdown("""
|
| 12 |
+
**MVP Chatbot Demo for AML-3304**
|
| 13 |
+
Supports:
|
| 14 |
+
- Basic Code Generation
|
| 15 |
+
- General Q&A
|
| 16 |
+
- Bayesian-style Q&A
|
| 17 |
+
- Creative Writing
|
| 18 |
+
- DeepSeekβR1 RAG (Retrieval-Augmented Q&A)
|
| 19 |
+
""")
|
| 20 |
|
|
|
|
| 21 |
if 'mode' not in st.session_state:
|
| 22 |
st.session_state.mode = 'Code Generator'
|
| 23 |
|
| 24 |
+
cols = st.columns(5)
|
| 25 |
+
labels = [
|
| 26 |
+
("π§βπ» Basic Code", 'Code Generator'),
|
| 27 |
+
("π General Q&A", 'General Q&A'),
|
| 28 |
+
("π Bayesian Q&A", 'Bayesian Q&A'),
|
| 29 |
+
("βοΈ Creative Writing", 'Creative Writing'),
|
| 30 |
+
("π RAG Q&A", 'RAG Q&A'),
|
| 31 |
+
]
|
| 32 |
+
for col, (emoji, mode_name) in zip(cols, labels):
|
| 33 |
+
if col.button(emoji):
|
| 34 |
+
st.session_state.mode = mode_name
|
| 35 |
|
| 36 |
st.markdown(f"### π Current Mode: **{st.session_state.mode}**")
|
| 37 |
+
user_input = st.text_input("Enter your prompt/question:")
|
| 38 |
|
|
|
|
|
|
|
|
|
|
|
|
|
| 39 |
@st.cache_resource
|
| 40 |
def load_codegen():
|
| 41 |
+
return (AutoTokenizer.from_pretrained("Salesforce/codegen-350M-mono"),
|
| 42 |
+
AutoModelForCausalLM.from_pretrained("Salesforce/codegen-350M-mono"))
|
| 43 |
+
|
| 44 |
+
@st.cache_resource
|
| 45 |
+
def load_flan():
|
| 46 |
+
return (AutoTokenizer.from_pretrained("google/flan-t5-base"),
|
| 47 |
+
AutoModelForSeq2SeqLM.from_pretrained("google/flan-t5-base"))
|
| 48 |
|
| 49 |
@st.cache_resource
|
| 50 |
+
def load_falcon():
|
| 51 |
+
return (AutoTokenizer.from_pretrained("tiiuae/falcon-rw-1b"),
|
| 52 |
+
AutoModelForCausalLM.from_pretrained("tiiuae/falcon-rw-1b"))
|
| 53 |
+
|
| 54 |
+
@st.cache_resource
|
| 55 |
+
def load_rag():
|
| 56 |
+
tokenizer = RagTokenizer.from_pretrained("deepseek-ai/deepseek-r1-rag")
|
| 57 |
+
retriever = RagRetriever.from_pretrained("deepseek-ai/deepseek-r1-rag", index_name="custom")
|
| 58 |
+
model = RagSequenceForGeneration.from_pretrained("deepseek-ai/deepseek-r1-rag", retriever=retriever)
|
| 59 |
return tokenizer, model
|
| 60 |
|
| 61 |
+
if user_input:
|
|
|
|
| 62 |
with st.spinner("Generating response..."):
|
| 63 |
+
mode = st.session_state.mode
|
| 64 |
+
|
| 65 |
+
if mode == 'Code Generator':
|
| 66 |
+
tok, mod = load_codegen()
|
| 67 |
+
inp = tok(user_input, return_tensors="pt")
|
| 68 |
+
out = mod.generate(inp.input_ids, max_new_tokens=128, temperature=0.7, do_sample=True)
|
| 69 |
+
st.code(tok.decode(out[0], skip_special_tokens=True), language="python")
|
| 70 |
+
|
| 71 |
+
elif mode in ('General Q&A', 'Bayesian Q&A'):
|
| 72 |
+
tok, mod = load_flan()
|
| 73 |
+
inp = tok(user_input, return_tensors="pt")
|
| 74 |
+
if mode == 'General Q&A':
|
| 75 |
+
out = mod.generate(inp.input_ids, max_new_tokens=100)
|
| 76 |
+
else:
|
| 77 |
+
out = mod.generate(inp.input_ids, max_new_tokens=100, do_sample=True, temperature=0.9, top_k=50)
|
| 78 |
+
st.write(tok.decode(out[0], skip_special_tokens=True))
|
| 79 |
|
| 80 |
+
elif mode == 'Creative Writing':
|
| 81 |
+
tok, mod = load_falcon()
|
| 82 |
+
inp = tok(user_input, return_tensors="pt")
|
| 83 |
+
out = mod.generate(inp.input_ids, max_new_tokens=150, do_sample=True, top_p=0.95, temperature=1.0)
|
| 84 |
+
st.write(tok.decode(out[0], skip_special_tokens=True))
|
|
|
|
|
|
|
| 85 |
|
| 86 |
+
elif mode == 'RAG Q&A':
|
| 87 |
+
tok, mod = load_rag()
|
| 88 |
+
inp = tok(user_input, return_tensors="pt")
|
| 89 |
+
out = mod.generate(input_ids=inp["input_ids"], context_input_ids=inp["context_input_ids"], max_new_tokens=100)
|
| 90 |
+
st.write(tok.batch_decode(out, skip_special_tokens=True)[0])
|
|
|
|
|
|