Spaces:
Sleeping
Sleeping
File size: 6,278 Bytes
1dd0d05 ada094c 660a97d 9c0cdb0 660a97d 1dd0d05 487797a 660a97d 487797a 660a97d 487797a 660a97d 487797a 660a97d 487797a 660a97d 1dd0d05 487797a 660a97d 487797a 660a97d 487797a 660a97d 487797a 660a97d 487797a 660a97d 487797a 660a97d 487797a 660a97d 1dd0d05 487797a 660a97d 1dd0d05 487797a 660a97d 487797a 660a97d 487797a 660a97d 8f463d0 660a97d 487797a 660a97d 487797a 660a97d 487797a 660a97d 487797a 660a97d 487797a 660a97d 487797a 660a97d 487797a 660a97d 487797a 660a97d 487797a 660a97d 487797a 660a97d 487797a 660a97d 487797a 660a97d 1dd0d05 487797a 1dd0d05 660a97d 487797a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 | 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)
|