Gilastro-AI / app.py
Muyumba's picture
app.py
487797a verified
Raw
History Blame Contribute Delete
6.28 kB
import gradio as gr
import torch
from transformers import GPT2Tokenizer, GPT2LMHeadModel
import os
# 🎯 Configuration du modèle
MODEL_NAME = "Muyumba/oprimus_ai" # Ou votre modèle fine-tuné
DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# 📝 Variables globales pour le cache
tokenizer = None
model = None
def load_model():
"""Charge le modèle et le tokenizer"""
global tokenizer, model
if tokenizer is None or model is None:
print("🔄 Chargement du modèle...")
tokenizer = GPT2Tokenizer.from_pretrained(MODEL_NAME)
model = GPT2LMHeadModel.from_pretrained(MODEL_NAME)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
model.to(DEVICE)
model.eval()
print("✅ Modèle chargé avec succès!")
return tokenizer, model
def generate_text(prompt, max_length=1000, temperature=0.8, top_p=0.9, repetition_penalty=1.1):
"""
Génère du texte à partir d'un prompt, sans limite stricte de tokens.
Utilise une génération incrémentale si nécessaire.
"""
try:
tokenizer, model = load_model()
if len(prompt.strip()) == 0:
return "❌ Veuillez entrer un prompt valide."
# Encodage du prompt
input_ids = tokenizer.encode(prompt, return_tensors="pt").to(DEVICE)
generated_ids = input_ids
total_length = input_ids.shape[1]
target_length = total_length + max_length
# Génération par morceaux
with torch.no_grad():
while generated_ids.shape[1] < target_length:
outputs = model.generate(
generated_ids,
max_length=min(generated_ids.shape[1] + 256, target_length),
temperature=max(0.1, min(2.0, temperature)),
top_p=max(0.1, min(1.0, top_p)),
repetition_penalty=max(1.0, min(2.0, repetition_penalty)),
do_sample=True,
pad_token_id=tokenizer.eos_token_id,
no_repeat_ngram_size=2,
)
generated_ids = outputs
# Arrêt si le modèle prédit EOS
if generated_ids[0, -1].item() == tokenizer.eos_token_id:
break
generated_text = tokenizer.decode(generated_ids[0], skip_special_tokens=True)
if generated_text.startswith(prompt):
generated_text = prompt + "\n\n" + "📝 **Génération:**\n" + generated_text[len(prompt):].strip()
return generated_text
except Exception as e:
return f"❌ Erreur lors de la génération: {str(e)}"
def clear_text():
return "", ""
# 🎨 Interface Gradio
def create_interface():
css = """
.gradio-container { font-family: 'Arial', sans-serif; max-width: 900px; }
.gr-button-primary { background: linear-gradient(45deg, #007bff, #0056b3); color: white; border: none; border-radius: 8px; padding: 10px 20px; font-weight: bold; }
.gr-textbox { border-radius: 8px; border: 2px solid #e0e0e0; }
"""
with gr.Blocks(css=css, title="🤖 Générateur de Texte IA") as demo:
gr.Markdown("""
# 🤖 Générateur de Texte Intelligent
### Propulsé par Poetry-AI
Entrez votre texte et laissez l'IA continuer votre histoire !
""")
with gr.Row():
with gr.Column(scale=1):
prompt_input = gr.Textbox(
label="📝 Votre prompt",
placeholder="Il était une fois...",
lines=5,
max_lines=10
)
with gr.Accordion("⚙️ Paramètres avancés", open=False):
max_length_slider = gr.Slider(
minimum=100,
maximum=5000,
value=1000,
step=50,
label="📏 Longueur maximale"
)
temperature_slider = gr.Slider(0.1, 2.0, value=0.8, step=0.1, label="🌡️ Créativité")
top_p_slider = gr.Slider(0.1, 1.0, value=0.9, step=0.05, label="🎯 Diversité")
rep_penalty_slider = gr.Slider(1.0, 2.0, value=1.1, step=0.05, label="🔄 Anti-répétition")
with gr.Row():
generate_btn = gr.Button("🚀 Générer", variant="primary")
clear_btn = gr.Button("🗑️ Effacer", variant="secondary")
with gr.Column(scale=1):
output_text = gr.Textbox(
label="✨ Texte généré",
lines=20,
max_lines=30,
interactive=False
)
gr.Markdown("### 💡 Exemples de prompts:")
gr.Examples(
examples=[
["Il était une fois, dans un royaume lointain..."],
["Le scientifique découvrit quelque chose d'extraordinaire..."],
["Par une nuit d'orage, Marie entendit un bruit étrange..."],
["L'intelligence artificielle du futur pourrait..."],
["Dans les rues de Paris, un mystère se cachait..."]
],
inputs=prompt_input
)
generate_btn.click(
fn=generate_text,
inputs=[prompt_input, max_length_slider, temperature_slider, top_p_slider, rep_penalty_slider],
outputs=output_text
)
clear_btn.click(fn=clear_text, inputs=[], outputs=[prompt_input, output_text])
gr.Markdown("""
---
**ℹ️ Conseils d'utilisation:**
- 🌡️ Temperature élevée = plus créatif mais moins cohérent
- 🎯 Top-p faible = plus focalisé sur les mots probables
- 📏 Longueur = nombre de tokens à générer (peut dépasser la limite native en générant par morceaux)
- 🔄 Anti-répétition = évite les boucles de mots
""")
return demo
if __name__ == "__main__":
print("🔄 Initialisation de l'application...")
load_model()
demo = create_interface()
demo.launch(server_name="0.0.0.0", server_port=7860, share=False, debug=True, show_error=True)