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)