coding / App.py
yash184's picture
Create App.py
4606d64 verified
Raw
History Blame Contribute Delete
7.55 kB
import gradio as gr
import torch
from PIL import Image
import logging
from typing import Optional
import time
from diffusers import StableDiffusionXLImg2ImgPipeline, StableDiffusionXLPipeline
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
# ===== CONFIG =====
DEVICE = "cpu"
DTYPE = torch.float32
# ===== PIPELINE MANAGER =====
class PipelineManager:
def __init__(self):
self.txt2img_pipe = None
self.img2img_pipe = None
self.model_loaded = False
self.load_lock = False
def load_models(self):
"""Load SDXL models"""
try:
logger.info("πŸ“₯ Loading models...")
self.txt2img_pipe = StableDiffusionXLPipeline.from_pretrained(
"stabilityai/stable-diffusion-xl-base-1.0",
torch_dtype=DTYPE,
use_safetensors=True
)
self.txt2img_pipe = self.txt2img_pipe.to(DEVICE)
self.txt2img_pipe.enable_attention_slicing()
self.img2img_pipe = StableDiffusionXLImg2ImgPipeline.from_pretrained(
"stabilityai/stable-diffusion-xl-base-1.0",
torch_dtype=DTYPE,
use_safetensors=True
)
self.img2img_pipe = self.img2img_pipe.to(DEVICE)
self.img2img_pipe.enable_attention_slicing()
self.model_loaded = True
logger.info("βœ… Models loaded!")
return True
except Exception as e:
logger.error(f"❌ Error loading models: {e}")
return False
def initialize(self):
if self.load_lock:
return
self.load_lock = True
self.load_models()
self.load_lock = False
def generate_txt2img(
self,
prompt: str,
negative_prompt: str = "",
num_steps: int = 20,
guidance: float = 7.5,
height: int = 768,
width: int = 768,
seed: int = -1
) -> Image.Image:
if not self.model_loaded:
raise RuntimeError("Model not loaded")
if seed == -1:
seed = int(time.time())
generator = torch.Generator(device=DEVICE).manual_seed(seed)
logger.info(f"🎨 Generating: {prompt[:50]}...")
with torch.no_grad():
image = self.txt2img_pipe(
prompt=prompt,
negative_prompt=negative_prompt,
num_inference_steps=num_steps,
guidance_scale=guidance,
height=height,
width=width,
generator=generator
).images[0]
return image
def generate_img2img(
self,
prompt: str,
image: Image.Image,
negative_prompt: str = "",
num_steps: int = 20,
guidance: float = 7.5,
strength: float = 0.8,
seed: int = -1
) -> Image.Image:
if not self.model_loaded:
raise RuntimeError("Model not loaded")
if seed == -1:
seed = int(time.time())
generator = torch.Generator(device=DEVICE).manual_seed(seed)
image = image.resize((768, 768), Image.Resampling.LANCZOS)
logger.info(f"πŸ–ΌοΈ Transforming: {prompt[:50]}...")
with torch.no_grad():
image = self.img2img_pipe(
prompt=prompt,
image=image,
negative_prompt=negative_prompt,
num_inference_steps=num_steps,
guidance_scale=guidance,
strength=strength,
generator=generator
).images[0]
return image
pipeline_manager = PipelineManager()
# ===== UI FUNCTIONS =====
def txt2img(prompt, neg_prompt, steps, guidance, height, width, seed):
try:
if not pipeline_manager.model_loaded:
return None, "❌ Model loading..."
image = pipeline_manager.generate_txt2img(prompt, neg_prompt, steps, guidance, height, width, seed)
return image, "βœ… Done!"
except Exception as e:
return None, f"❌ {str(e)}"
def img2img(prompt, input_image, neg_prompt, steps, guidance, strength, seed):
try:
if input_image is None:
return None, "❌ Upload image first"
if not pipeline_manager.model_loaded:
return None, "❌ Model loading..."
image = pipeline_manager.generate_img2img(prompt, input_image, neg_prompt, steps, guidance, strength, seed)
return image, "βœ… Done!"
except Exception as e:
return None, f"❌ {str(e)}"
# ===== GRADIO UI =====
with gr.Blocks(title="FLUX Generator") as demo:
gr.Markdown("# 🎨 FLUX - Image Generator")
with gr.Tabs():
with gr.Tab("πŸ“ Text-to-Image"):
with gr.Row():
with gr.Column():
prompt = gr.Textbox(label="Prompt", lines=3, placeholder="Describe image...")
neg_prompt = gr.Textbox(label="Negative", lines=2, placeholder="What to avoid...")
with gr.Row():
height = gr.Slider(256, 1024, 768, 64, label="Height")
width = gr.Slider(256, 1024, 768, 64, label="Width")
with gr.Row():
steps = gr.Slider(1, 50, 20, 1, label="Steps")
guidance = gr.Slider(1, 15, 7.5, 0.5, label="Guidance")
seed = gr.Number(-1, label="Seed (-1=random)", precision=0)
btn = gr.Button("🎨 Generate", variant="primary", size="lg")
with gr.Column():
output = gr.Image(label="Output")
status = gr.Textbox(interactive=False, label="Status")
btn.click(txt2img, [prompt, neg_prompt, steps, guidance, height, width, seed], [output, status])
with gr.Tab("πŸ–ΌοΈ Image-to-Image"):
with gr.Row():
with gr.Column():
img_input = gr.Image(label="Input Image", type="pil")
prompt2 = gr.Textbox(label="Prompt", lines=3, placeholder="Transform to...")
neg_prompt2 = gr.Textbox(label="Negative", lines=2)
with gr.Row():
steps2 = gr.Slider(1, 50, 20, 1, label="Steps")
guidance2 = gr.Slider(1, 15, 7.5, 0.5, label="Guidance")
strength = gr.Slider(0, 1, 0.8, 0.05, label="Strength")
seed2 = gr.Number(-1, label="Seed (-1=random)", precision=0)
btn2 = gr.Button("πŸ–ΌοΈ Generate", variant="primary", size="lg")
with gr.Column():
output2 = gr.Image(label="Output")
status2 = gr.Textbox(interactive=False, label="Status")
btn2.click(img2img, [prompt2, img_input, neg_prompt2, steps2, guidance2, strength, seed2], [output2, status2])
def on_load():
logger.info("πŸš€ Loading pipeline...")
pipeline_manager.initialize()
if pipeline_manager.model_loaded:
return "βœ… Ready!"
return "⏳ Loading models..."
gr.on_load(on_load)
if __name__ == "__main__":
demo.launch(server_name="0.0.0.0", server_port=7860, share=True)