Spaces:
Sleeping
Sleeping
| import gradio as gr | |
| from transformers import AutoModelForCausalLM, AutoTokenizer | |
| import torch | |
| # ========================= | |
| # LOAD MODEL 1 | |
| # ========================= | |
| model_1_name = "microsoft/DialoGPT-small" | |
| tokenizer_1 = AutoTokenizer.from_pretrained(model_1_name) | |
| model_1 = AutoModelForCausalLM.from_pretrained(model_1_name) | |
| # ========================= | |
| # LOAD MODEL 2 | |
| # ========================= | |
| model_2_name = "microsoft/DialoGPT-medium" | |
| tokenizer_2 = AutoTokenizer.from_pretrained(model_2_name) | |
| model_2 = AutoModelForCausalLM.from_pretrained(model_2_name) | |
| device = "cpu" | |
| model_1.to(device) | |
| model_2.to(device) | |
| # ========================= | |
| # GENERATION FUNCTION | |
| # ========================= | |
| def generate(model, tokenizer, message, history_ids): | |
| input_ids = tokenizer.encode(message + tokenizer.eos_token, return_tensors="pt") | |
| if history_ids is not None: | |
| bot_input = torch.cat([history_ids, input_ids], dim=-1) | |
| else: | |
| bot_input = input_ids | |
| output = model.generate( | |
| bot_input, | |
| max_length=200, | |
| pad_token_id=tokenizer.eos_token_id, | |
| do_sample=True, | |
| top_k=50, | |
| top_p=0.95, | |
| temperature=0.9 | |
| ) | |
| reply = tokenizer.decode( | |
| output[:, bot_input.shape[-1]:][0], | |
| skip_special_tokens=True | |
| ) | |
| return reply, output | |
| # ========================= | |
| # CHAT LOOP FUNCTION | |
| # ========================= | |
| def dual_chat(user_input, history): | |
| if history is None: | |
| history = {"msg": "Hello", "h1": None, "h2": None} | |
| msg = user_input | |
| chat_log = "" | |
| for _ in range(3): # rounds | |
| chat_log += f"π΅ Model 1: {msg}\n" | |
| reply2, history["h2"] = generate(model_2, tokenizer_2, msg, history["h2"]) | |
| chat_log += f"π΄ Model 2: {reply2}\n" | |
| reply1, history["h1"] = generate(model_1, tokenizer_1, reply2, history["h1"]) | |
| msg = reply1 | |
| return chat_log, history | |
| # ========================= | |
| # GRADIO UI | |
| # ========================= | |
| with gr.Blocks() as demo: | |
| gr.Markdown("# π€ AI vs AI Chat (Hugging Face)") | |
| input_box = gr.Textbox(label="Start Conversation") | |
| output_box = gr.Textbox(label="Chat Output") | |
| state = gr.State() | |
| btn = gr.Button("Run AI Chat") | |
| btn.click( | |
| dual_chat, | |
| inputs=[input_box, state], | |
| outputs=[output_box, state] | |
| ) | |
| demo.launch() |