MagicFlow / app.py
lea97338's picture
Update app.py
132881e verified
Raw
History Blame Contribute Delete
3.26 kB
import torch
import traceback
import io
import gradio as gr
from PIL import Image
from diffusers import Flux2KleinPipeline
# 1. Chargement initial du modèle (exécuté une seule fois au démarrage)
print("Loading model...")
pipe = Flux2KleinPipeline.from_pretrained(
"black-forest-labs/FLUX.2-klein-base-4B",
torch_dtype=torch.bfloat16,
)
pipe.to("cuda")
pipe.load_lora_weights("./sking_v73_flux_4b_000027000.safetensors")
print("Model loaded successfully.")
# 2. Fonction principale appelée par l'interface Gradio
def generate_minecraft_image(input_image, prompt, guidance, seed, n_step):
if input_image is None:
raise gr.Error("Veuillez téléverser une image de départ.")
try:
# Prétraitement similaire à votre API (Conversion RGBA)
img = input_image.convert("RGBA")
# Exécution du pipeline Flux 2 Klein
pipeline_output = pipe(
image=img,
prompt=prompt,
height=768,
width=768,
num_inference_steps=int(n_step),
guidance_scale=float(guidance),
num_images_per_prompt=1,
generator=torch.Generator("cuda").manual_seed(int(seed))
)
# Retourne la liste des images générées à la galerie Gradio
return pipeline_output.images
except Exception as e:
traceback.print_exc()
print(f"Error during generation: {e}")
raise gr.Error(f"Erreur lors de la génération : {str(e)}")
# 3. Construction de l'interface graphique (Gradio)
with gr.Blocks(title="Minecraft IA - Flux.2 Klein", theme=gr.themes.Soft()) as demo:
gr.Markdown("# 🎮 Minecraft Pixel-to-Pixel IA Generator")
gr.Markdown("Transformez vos images ou appliquez votre style LoRA avec FLUX.2 Klein [4B].")
with gr.Row():
# Colonne de gauche : Entrées et paramètres
with gr.Column():
input_img = gr.Image(label="Image source (Image Input)", type="pil")
prompt_txt = gr.Textbox(label="Prompt (Mot-clé / Description)", value="", placeholder="Ex: minecraft skin, isometric block...")
with gr.Accordion("Paramètres avancés", open=True):
guidance_slider = gr.Slider(minimum=1.0, maximum=20.0, value=4.0, step=0.5, label="Guidance Scale")
steps_slider = gr.Slider(minimum=1, maximum=150, value=100, step=1, label="Nombre de pas (Inference Steps)")
seed_number = gr.Number(value=42, label="Seed", precision=0)
submit_btn = gr.Button("🚀 Générer l'image", variant="primary")
# Colonne de droite : Résultats
with gr.Column():
output_gallery = gr.Gallery(label="Images Générées", columns=1, height="auto", object_fit="contain")
# Liaison du bouton à la fonction de génération
submit_btn.click(
fn=generate_minecraft_image,
inputs=[input_img, prompt_txt, guidance_slider, seed_number, steps_slider],
outputs=output_gallery
)
# 4. Lancement du serveur Web local
if __name__ == "__main__":
# Lance l'interface sur le port 10012 similaire à votre configuration uvicorn
demo.launch(server_name="0.0.0.0", server_port=10012)