yash184 commited on
Commit
4606d64
Β·
verified Β·
1 Parent(s): bc386b2

Create App.py

Browse files
Files changed (1) hide show
  1. App.py +214 -0
App.py ADDED
@@ -0,0 +1,214 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import gradio as gr
2
+ import torch
3
+ from PIL import Image
4
+ import logging
5
+ from typing import Optional
6
+ import time
7
+ from diffusers import StableDiffusionXLImg2ImgPipeline, StableDiffusionXLPipeline
8
+
9
+ logging.basicConfig(level=logging.INFO)
10
+ logger = logging.getLogger(__name__)
11
+
12
+ # ===== CONFIG =====
13
+ DEVICE = "cpu"
14
+ DTYPE = torch.float32
15
+
16
+ # ===== PIPELINE MANAGER =====
17
+ class PipelineManager:
18
+ def __init__(self):
19
+ self.txt2img_pipe = None
20
+ self.img2img_pipe = None
21
+ self.model_loaded = False
22
+ self.load_lock = False
23
+
24
+ def load_models(self):
25
+ """Load SDXL models"""
26
+ try:
27
+ logger.info("πŸ“₯ Loading models...")
28
+
29
+ self.txt2img_pipe = StableDiffusionXLPipeline.from_pretrained(
30
+ "stabilityai/stable-diffusion-xl-base-1.0",
31
+ torch_dtype=DTYPE,
32
+ use_safetensors=True
33
+ )
34
+ self.txt2img_pipe = self.txt2img_pipe.to(DEVICE)
35
+ self.txt2img_pipe.enable_attention_slicing()
36
+
37
+ self.img2img_pipe = StableDiffusionXLImg2ImgPipeline.from_pretrained(
38
+ "stabilityai/stable-diffusion-xl-base-1.0",
39
+ torch_dtype=DTYPE,
40
+ use_safetensors=True
41
+ )
42
+ self.img2img_pipe = self.img2img_pipe.to(DEVICE)
43
+ self.img2img_pipe.enable_attention_slicing()
44
+
45
+ self.model_loaded = True
46
+ logger.info("βœ… Models loaded!")
47
+ return True
48
+
49
+ except Exception as e:
50
+ logger.error(f"❌ Error loading models: {e}")
51
+ return False
52
+
53
+ def initialize(self):
54
+ if self.load_lock:
55
+ return
56
+ self.load_lock = True
57
+ self.load_models()
58
+ self.load_lock = False
59
+
60
+ def generate_txt2img(
61
+ self,
62
+ prompt: str,
63
+ negative_prompt: str = "",
64
+ num_steps: int = 20,
65
+ guidance: float = 7.5,
66
+ height: int = 768,
67
+ width: int = 768,
68
+ seed: int = -1
69
+ ) -> Image.Image:
70
+
71
+ if not self.model_loaded:
72
+ raise RuntimeError("Model not loaded")
73
+
74
+ if seed == -1:
75
+ seed = int(time.time())
76
+
77
+ generator = torch.Generator(device=DEVICE).manual_seed(seed)
78
+
79
+ logger.info(f"🎨 Generating: {prompt[:50]}...")
80
+
81
+ with torch.no_grad():
82
+ image = self.txt2img_pipe(
83
+ prompt=prompt,
84
+ negative_prompt=negative_prompt,
85
+ num_inference_steps=num_steps,
86
+ guidance_scale=guidance,
87
+ height=height,
88
+ width=width,
89
+ generator=generator
90
+ ).images[0]
91
+
92
+ return image
93
+
94
+ def generate_img2img(
95
+ self,
96
+ prompt: str,
97
+ image: Image.Image,
98
+ negative_prompt: str = "",
99
+ num_steps: int = 20,
100
+ guidance: float = 7.5,
101
+ strength: float = 0.8,
102
+ seed: int = -1
103
+ ) -> Image.Image:
104
+
105
+ if not self.model_loaded:
106
+ raise RuntimeError("Model not loaded")
107
+
108
+ if seed == -1:
109
+ seed = int(time.time())
110
+
111
+ generator = torch.Generator(device=DEVICE).manual_seed(seed)
112
+ image = image.resize((768, 768), Image.Resampling.LANCZOS)
113
+
114
+ logger.info(f"πŸ–ΌοΈ Transforming: {prompt[:50]}...")
115
+
116
+ with torch.no_grad():
117
+ image = self.img2img_pipe(
118
+ prompt=prompt,
119
+ image=image,
120
+ negative_prompt=negative_prompt,
121
+ num_inference_steps=num_steps,
122
+ guidance_scale=guidance,
123
+ strength=strength,
124
+ generator=generator
125
+ ).images[0]
126
+
127
+ return image
128
+
129
+ pipeline_manager = PipelineManager()
130
+
131
+ # ===== UI FUNCTIONS =====
132
+
133
+ def txt2img(prompt, neg_prompt, steps, guidance, height, width, seed):
134
+ try:
135
+ if not pipeline_manager.model_loaded:
136
+ return None, "❌ Model loading..."
137
+ image = pipeline_manager.generate_txt2img(prompt, neg_prompt, steps, guidance, height, width, seed)
138
+ return image, "βœ… Done!"
139
+ except Exception as e:
140
+ return None, f"❌ {str(e)}"
141
+
142
+ def img2img(prompt, input_image, neg_prompt, steps, guidance, strength, seed):
143
+ try:
144
+ if input_image is None:
145
+ return None, "❌ Upload image first"
146
+ if not pipeline_manager.model_loaded:
147
+ return None, "❌ Model loading..."
148
+ image = pipeline_manager.generate_img2img(prompt, input_image, neg_prompt, steps, guidance, strength, seed)
149
+ return image, "βœ… Done!"
150
+ except Exception as e:
151
+ return None, f"❌ {str(e)}"
152
+
153
+ # ===== GRADIO UI =====
154
+
155
+ with gr.Blocks(title="FLUX Generator") as demo:
156
+ gr.Markdown("# 🎨 FLUX - Image Generator")
157
+
158
+ with gr.Tabs():
159
+
160
+ with gr.Tab("πŸ“ Text-to-Image"):
161
+ with gr.Row():
162
+ with gr.Column():
163
+ prompt = gr.Textbox(label="Prompt", lines=3, placeholder="Describe image...")
164
+ neg_prompt = gr.Textbox(label="Negative", lines=2, placeholder="What to avoid...")
165
+
166
+ with gr.Row():
167
+ height = gr.Slider(256, 1024, 768, 64, label="Height")
168
+ width = gr.Slider(256, 1024, 768, 64, label="Width")
169
+
170
+ with gr.Row():
171
+ steps = gr.Slider(1, 50, 20, 1, label="Steps")
172
+ guidance = gr.Slider(1, 15, 7.5, 0.5, label="Guidance")
173
+
174
+ seed = gr.Number(-1, label="Seed (-1=random)", precision=0)
175
+ btn = gr.Button("🎨 Generate", variant="primary", size="lg")
176
+
177
+ with gr.Column():
178
+ output = gr.Image(label="Output")
179
+ status = gr.Textbox(interactive=False, label="Status")
180
+
181
+ btn.click(txt2img, [prompt, neg_prompt, steps, guidance, height, width, seed], [output, status])
182
+
183
+ with gr.Tab("πŸ–ΌοΈ Image-to-Image"):
184
+ with gr.Row():
185
+ with gr.Column():
186
+ img_input = gr.Image(label="Input Image", type="pil")
187
+ prompt2 = gr.Textbox(label="Prompt", lines=3, placeholder="Transform to...")
188
+ neg_prompt2 = gr.Textbox(label="Negative", lines=2)
189
+
190
+ with gr.Row():
191
+ steps2 = gr.Slider(1, 50, 20, 1, label="Steps")
192
+ guidance2 = gr.Slider(1, 15, 7.5, 0.5, label="Guidance")
193
+
194
+ strength = gr.Slider(0, 1, 0.8, 0.05, label="Strength")
195
+ seed2 = gr.Number(-1, label="Seed (-1=random)", precision=0)
196
+ btn2 = gr.Button("πŸ–ΌοΈ Generate", variant="primary", size="lg")
197
+
198
+ with gr.Column():
199
+ output2 = gr.Image(label="Output")
200
+ status2 = gr.Textbox(interactive=False, label="Status")
201
+
202
+ btn2.click(img2img, [prompt2, img_input, neg_prompt2, steps2, guidance2, strength, seed2], [output2, status2])
203
+
204
+ def on_load():
205
+ logger.info("πŸš€ Loading pipeline...")
206
+ pipeline_manager.initialize()
207
+ if pipeline_manager.model_loaded:
208
+ return "βœ… Ready!"
209
+ return "⏳ Loading models..."
210
+
211
+ gr.on_load(on_load)
212
+
213
+ if __name__ == "__main__":
214
+ demo.launch(server_name="0.0.0.0", server_port=7860, share=True)