| import streamlit as st |
| from transformers import AutoTokenizer, AutoModelForCausalLM |
| import torch |
|
|
| |
| @st.cache_resource |
| def load_model(): |
| model_name = "sshleifer/tiny-gpt2" |
| tokenizer = AutoTokenizer.from_pretrained(model_name) |
| model = AutoModelForCausalLM.from_pretrained(model_name) |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| model.to(device) |
| model.eval() |
| return tokenizer, model, device |
|
|
| tokenizer, model, device = load_model() |
|
|
| |
| if "messages" not in st.session_state: |
| st.session_state.messages = [] |
|
|
| |
| st.title("π€ Tiny GPT-2 Chatbot") |
| st.markdown("Ask me anything! This bot runs locally with no API key.") |
|
|
| |
| for msg in st.session_state.messages: |
| role = "π§βπ»" if msg["role"] == "user" else "π€" |
| st.markdown(f"**{role}:** {msg['content']}") |
|
|
| |
| user_input = st.text_input("Type your message...", key="user_input") |
|
|
| if user_input: |
| |
| st.session_state.messages.append({"role": "user", "content": user_input}) |
|
|
| |
| prompt = "\n".join([m["content"] for m in st.session_state.messages]) |
|
|
| |
| inputs = tokenizer(prompt, return_tensors="pt").to(device) |
| outputs = model.generate( |
| **inputs, |
| max_new_tokens=100, |
| temperature=0.7, |
| top_p=0.95, |
| do_sample=True, |
| pad_token_id=tokenizer.eos_token_id |
| ) |
|
|
| output_text = tokenizer.decode(outputs[0], skip_special_tokens=True) |
| reply = output_text[len(prompt):].strip().split("\n")[0] |
|
|
| |
| st.session_state.messages.append({"role": "assistant", "content": reply}) |
|
|
| |
| st.experimental_rerun() |
|
|