Spaces:
Sleeping
Sleeping
| 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) | |