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)