Spaces:
Sleeping
Sleeping
| import streamlit as st | |
| import torch | |
| from transformers import AutoModelForCausalLM, AutoTokenizer | |
| def load_model(): | |
| tokenizer = AutoTokenizer.from_pretrained("microsoft/DialoGPT-medium") | |
| model = AutoModelForCausalLM.from_pretrained("microsoft/DialoGPT-medium") | |
| return tokenizer, model | |
| st.title("ChatBot with DialoGPT") | |
| tokenizer, model = load_model() | |
| # initialize history | |
| if 'chat_history_ids' not in st.session_state: | |
| st.session_state.chat_history_ids = None | |
| user_input = st.text_input("You:", "") | |
| if user_input: | |
| new_input = tokenizer.encode(user_input + tokenizer.eos_token, return_tensors='pt') | |
| bot_input = torch.cat([st.session_state.chat_history_ids, new_input], dim=-1) \ | |
| if st.session_state.chat_history_ids is not None else new_input | |
| st.session_state.chat_history_ids = model.generate(bot_input, max_length=1000, pad_token_id=tokenizer.eos_token_id) | |
| response = tokenizer.decode( | |
| st.session_state.chat_history_ids[:, bot_input.shape[-1]:][0], | |
| skip_special_tokens=True | |
| ) | |
| st.write(f"**Bot:** {response}") | |