merged-AI / app.py
Muyumba's picture
app.py
7fcd80e verified
Raw
History Blame Contribute Delete
1.58 kB
import torch
from transformers import GPT2Tokenizer, GPT2LMHeadModel
import gradio as gr
# Chargement du modèle depuis Hugging Face
tokenizer = GPT2Tokenizer.from_pretrained("Muyumba/gpt2-merged")
model = GPT2LMHeadModel.from_pretrained("Muyumba/gpt2-merged")
model.eval()
# Fonction d'inférence simple
def generate_response(message, history, temperature, max_new_tokens, top_p):
prompt = ""
for user_input, bot_reply in history:
prompt += f"User: {user_input}\nAI: {bot_reply}\n"
prompt += f"User: {message}\nAI:"
inputs = tokenizer.encode(prompt, return_tensors="pt")
outputs = model.generate(
inputs,
max_new_tokens=max_new_tokens,
do_sample=True,
temperature=temperature,
top_p=top_p,
pad_token_id=tokenizer.eos_token_id,
)
response = tokenizer.decode(outputs[0], skip_special_tokens=True)
response = response.split("AI:")[-1].strip()
return response
# Interface Gradio avec historique de chat
chat = gr.ChatInterface(
fn=generate_response,
title="Merged AI Chatbot",
description="Un chatbot basé sur le modèle GPT2 fusionné.",
chatbot=gr.Chatbot(),
textbox=gr.Textbox(placeholder="Pose ta question ici..."),
additional_inputs=[
gr.Slider(50, 1024, value=128, label="Max new tokens"),
gr.Slider(0.1, 1.5, value=0.7, step=0.1, label="Temperature"),
gr.Slider(0.1, 1.0, value=0.95, step=0.05, label="Top-p"),
],
)
if __name__ == "__main__":
chat.launch(server_name="0.0.0.0", server_port=7860, share=True, ssr=False)