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)