kootaro commited on
Commit
dc98e61
·
verified ·
1 Parent(s): 9150c4f

Create app.py

Browse files
Files changed (1) hide show
  1. app.py +362 -0
app.py ADDED
@@ -0,0 +1,362 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """
2
+ FLUX.2 Klein 9B (distilled) + optional uncensored TE + NSFW LoRAs
3
+ Lighter than SDXL UnrealVision / Klein-base-50-steps.
4
+ """
5
+ import os
6
+ import gc
7
+ import json
8
+ import random
9
+ import base64
10
+ from io import BytesIO
11
+ from pathlib import Path
12
+
13
+ import gradio as gr
14
+ from gradio import Server
15
+ from fastapi.responses import HTMLResponse
16
+ import numpy as np
17
+ import spaces
18
+ import torch
19
+ from PIL import Image
20
+
21
+ MAX_SEED = np.iinfo(np.int32).max
22
+ LANCZOS = getattr(Image, "Resampling", Image).LANCZOS
23
+
24
+ device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
25
+ print("device:", device, "cuda:", torch.cuda.is_available())
26
+
27
+ from diffusers import Flux2KleinPipeline
28
+
29
+ dtype = torch.bfloat16
30
+
31
+ # Distilled 9B = 4 steps, much lighter than base-9B (28–50 steps)
32
+ MODEL_ID = os.getenv("KLEIN_MODEL", "black-forest-labs/FLUX.2-klein-9B")
33
+ USE_UNCENSORED_TE = os.getenv("KLEIN_UNCENSORED_TE", "1").strip() not in ("0", "false", "no")
34
+ UNCENSORED_TE_REPO = os.getenv(
35
+ "KLEIN_TE_REPO",
36
+ "ponpoke/flux2-klein-9b-uncensored-text-encoder",
37
+ )
38
+
39
+ ADAPTER_SPECS = {
40
+ "None": {
41
+ "repo": None,
42
+ "weights": None,
43
+ "adapter_name": None,
44
+ "default_strength": 0.0,
45
+ "hint": "No LoRA",
46
+ },
47
+ "NSFW-Solo": {
48
+ "repo": "diroverflo/FLux_Klein_9B_NSFW",
49
+ "weights": None, # auto from repo
50
+ "adapter_name": "nsfw-solo",
51
+ "default_strength": 0.75,
52
+ "hint": "Solo NSFW focus (no faces in train set)",
53
+ },
54
+ "Consistency": {
55
+ "repo": "dx8152/Flux2-Klein-9B-Consistency",
56
+ "weights": None,
57
+ "adapter_name": "klein-consistency",
58
+ "default_strength": 0.8,
59
+ "hint": "Identity / edit consistency (not NSFW unlock)",
60
+ },
61
+ }
62
+
63
+ LOADED_ADAPTERS: set = set()
64
+ ADAPTER_NAMES = list(ADAPTER_SPECS.keys())
65
+
66
+ print(f"Loading Flux2KleinPipeline: {MODEL_ID}")
67
+ pipe = Flux2KleinPipeline.from_pretrained(
68
+ MODEL_ID,
69
+ torch_dtype=dtype,
70
+ )
71
+ if torch.cuda.is_available():
72
+ try:
73
+ pipe.enable_model_cpu_offload()
74
+ print("cpu_offload OK")
75
+ except Exception as e:
76
+ print("offload fail:", e)
77
+ pipe = pipe.to(device)
78
+ else:
79
+ pipe = pipe.to(device)
80
+
81
+ # Best-effort uncensored text encoder swap
82
+ if USE_UNCENSORED_TE:
83
+ try:
84
+ from transformers import AutoModel, AutoTokenizer
85
+
86
+ print(f"[TE] loading uncensored encoder: {UNCENSORED_TE_REPO}")
87
+ tok = AutoTokenizer.from_pretrained(
88
+ UNCENSORED_TE_REPO,
89
+ token=os.getenv("HF_TOKEN") or None,
90
+ )
91
+ te = AutoModel.from_pretrained(
92
+ UNCENSORED_TE_REPO,
93
+ torch_dtype=dtype,
94
+ token=os.getenv("HF_TOKEN") or None,
95
+ )
96
+ # Flux2Klein may expose text_encoder / tokenizer attributes
97
+ if hasattr(pipe, "text_encoder"):
98
+ pipe.text_encoder = te
99
+ if hasattr(pipe, "tokenizer"):
100
+ pipe.tokenizer = tok
101
+ print("[TE] uncensored text encoder attached (best-effort)")
102
+ except Exception as e:
103
+ print(f"[TE] skip uncensored TE ({type(e).__name__}: {e})")
104
+ print("[TE] continuing with official encoder + NSFW LoRAs")
105
+
106
+ print("Pipeline ready.")
107
+
108
+
109
+ def b64_to_pil_list(b64_json_str):
110
+ if not b64_json_str or str(b64_json_str).strip() in ("", "[]"):
111
+ return []
112
+ try:
113
+ b64_list = json.loads(b64_json_str)
114
+ except Exception:
115
+ return []
116
+ out = []
117
+ for b64_str in b64_list:
118
+ if not b64_str or not isinstance(b64_str, str):
119
+ continue
120
+ try:
121
+ if b64_str.startswith("data:image"):
122
+ _, data = b64_str.split(",", 1)
123
+ else:
124
+ data = b64_str
125
+ out.append(Image.open(BytesIO(base64.b64decode(data))).convert("RGB"))
126
+ except Exception as e:
127
+ print("decode error:", e)
128
+ return out
129
+
130
+
131
+ def pil_to_b64_png(image: Image.Image) -> str:
132
+ buf = BytesIO()
133
+ image.save(buf, format="PNG")
134
+ return f"data:image/png;base64,{base64.b64encode(buf.getvalue()).decode()}"
135
+
136
+
137
+ def update_dimensions(image: Image.Image, max_side: int = 1024):
138
+ if image is None:
139
+ return 1024, 1024
140
+ w, h = image.size
141
+ if w >= h:
142
+ nw = max_side
143
+ nh = int(nw * h / w)
144
+ else:
145
+ nh = max_side
146
+ nw = int(nh * w / h)
147
+ return max(16, (nw // 16) * 16), max(16, (nh // 16) * 16)
148
+
149
+
150
+ app = Server(title="FLUX2-Klein-9B-NSFW")
151
+
152
+
153
+ @app.api(name="check_safety", queue=False)
154
+ def check_safety(prompt: str) -> dict:
155
+ return {"status": "ok"}
156
+
157
+
158
+ @app.api(name="generate")
159
+ @spaces.GPU(duration=120)
160
+ def generate(
161
+ images_b64_json: str,
162
+ prompt: str,
163
+ lora_adapter: str,
164
+ seed: int,
165
+ randomize_seed: bool,
166
+ guidance_scale: float,
167
+ steps: int,
168
+ lora_strength: float = 0.75,
169
+ ) -> dict:
170
+ gc.collect()
171
+ if torch.cuda.is_available():
172
+ torch.cuda.empty_cache()
173
+
174
+ if not prompt or not str(prompt).strip():
175
+ raise gr.Error("Prompt is empty.")
176
+
177
+ lora_adapter = (lora_adapter or "NSFW-Solo").strip()
178
+ spec = ADAPTER_SPECS.get(lora_adapter) or ADAPTER_SPECS["None"]
179
+ strength = float(lora_strength if lora_strength is not None else spec["default_strength"])
180
+ strength = max(0.0, min(strength, 1.5))
181
+
182
+ if spec["adapter_name"] is None or strength <= 0:
183
+ try:
184
+ pipe.disable_lora()
185
+ except Exception:
186
+ pass
187
+ print("--- LoRA off ---")
188
+ else:
189
+ name = spec["adapter_name"]
190
+ if name not in LOADED_ADAPTERS:
191
+ print(f"--- Loading LoRA {lora_adapter} from {spec['repo']} ---")
192
+ try:
193
+ kwargs = {"adapter_name": name}
194
+ if spec.get("weights"):
195
+ kwargs["weight_name"] = spec["weights"]
196
+ pipe.load_lora_weights(spec["repo"], **kwargs)
197
+ LOADED_ADAPTERS.add(name)
198
+ except Exception as e:
199
+ raise gr.Error(f"LoRA load failed ({lora_adapter}): {e}")
200
+ try:
201
+ pipe.set_adapters([name], adapter_weights=[strength])
202
+ except Exception:
203
+ pipe.set_adapters([name])
204
+ print(f"--- LoRA {name} @ {strength} ---")
205
+
206
+ if randomize_seed:
207
+ seed = random.randint(0, MAX_SEED)
208
+ seed = int(seed)
209
+ generator = torch.Generator(
210
+ device="cuda" if torch.cuda.is_available() else "cpu"
211
+ ).manual_seed(seed)
212
+
213
+ # Distilled Klein defaults: 4 steps, guidance ~1.0–4.0
214
+ steps = int(max(1, min(int(steps or 4), 28)))
215
+ guidance_scale = float(max(0.0, min(float(guidance_scale or 1.0), 8.0)))
216
+
217
+ pil_images = b64_to_pil_list(images_b64_json)
218
+ if pil_images:
219
+ width, height = update_dimensions(pil_images[0])
220
+ processed = [im.resize((width, height), LANCZOS) for im in pil_images]
221
+ image_input = processed if len(processed) > 1 else processed[0]
222
+ print(f"I2I {width}x{height} steps={steps} cfg={guidance_scale}")
223
+ else:
224
+ width, height = 1024, 1024
225
+ image_input = None
226
+ print(f"T2I {width}x{height} steps={steps} cfg={guidance_scale}")
227
+
228
+ try:
229
+ kwargs = dict(
230
+ prompt=str(prompt).strip(),
231
+ height=height,
232
+ width=width,
233
+ num_inference_steps=steps,
234
+ guidance_scale=guidance_scale,
235
+ generator=generator,
236
+ )
237
+ if image_input is not None:
238
+ kwargs["image"] = image_input
239
+ result = pipe(**kwargs).images[0]
240
+ return {
241
+ "image": pil_to_b64_png(result),
242
+ "seed": seed,
243
+ "status": "success",
244
+ "width": width,
245
+ "height": height,
246
+ "steps": steps,
247
+ "guidance_scale": guidance_scale,
248
+ "lora": lora_adapter,
249
+ "lora_strength": strength,
250
+ }
251
+ except Exception as e:
252
+ raise gr.Error(f"Inference failed: {type(e).__name__}: {e}")
253
+ finally:
254
+ gc.collect()
255
+ if torch.cuda.is_available():
256
+ torch.cuda.empty_cache()
257
+
258
+
259
+ @app.get("/api/config")
260
+ def client_config():
261
+ return {
262
+ "model": MODEL_ID,
263
+ "uncensored_te": USE_UNCENSORED_TE,
264
+ "loras": ADAPTER_NAMES,
265
+ "defaults": {
266
+ "steps": 4,
267
+ "guidance_scale": 1.0,
268
+ "lora": "NSFW-Solo",
269
+ "lora_strength": 0.75,
270
+ },
271
+ "notes": [
272
+ "Distilled Klein 9B — typically 4 steps",
273
+ "TE uncensored is best-effort; LoRA teaches NSFW concepts to DiT",
274
+ "Accept FLUX license on black-forest-labs/FLUX.2-klein-9B",
275
+ "HF_TOKEN if gated models need auth",
276
+ ],
277
+ }
278
+
279
+
280
+ @app.get("/", response_class=HTMLResponse)
281
+ async def homepage():
282
+ html_path = Path(__file__).resolve().parent / "index.html"
283
+ if html_path.exists():
284
+ return html_path.read_text(encoding="utf-8")
285
+ return """
286
+ <!DOCTYPE html>
287
+ <html lang="en"><head>
288
+ <meta charset="utf-8"/><meta name="viewport" content="width=device-width,initial-scale=1"/>
289
+ <title>FLUX.2 Klein 9B · NSFW</title>
290
+ <style>
291
+ :root{--bg:#0b0d12;--card:#141821;--fg:#e8eaed;--muted:#9aa0a6;--a:#a78bfa}
292
+ *{box-sizing:border-box}body{margin:0;font-family:Inter,system-ui,sans-serif;background:var(--bg);color:var(--fg)}
293
+ .wrap{max-width:920px;margin:0 auto;padding:24px 16px 48px}h1{margin:0 0 4px;font-size:1.45rem}
294
+ .sub{color:var(--muted);margin-bottom:16px;font-size:.9rem}
295
+ .card{background:var(--card);border:1px solid #252a36;border-radius:14px;padding:14px;margin-bottom:12px}
296
+ label{display:block;font-size:.8rem;color:var(--muted);margin-bottom:6px}
297
+ textarea,input,select{width:100%;background:#0a0c10;color:var(--fg);border:1px solid #2a2f3a;border-radius:10px;padding:10px}
298
+ textarea{min-height:88px}.row{display:grid;grid-template-columns:1fr 1fr;gap:10px}
299
+ button{width:100%;border:0;border-radius:12px;padding:12px;font-weight:700;cursor:pointer;
300
+ background:linear-gradient(135deg,#7c3aed,#db2777);color:#fff;font-size:1rem}
301
+ button:disabled{opacity:.5}#out{max-width:100%;border-radius:12px;margin-top:10px}
302
+ #status{white-space:pre-wrap;font-size:.85rem;color:var(--muted)}.hint{font-size:.75rem;color:var(--muted);margin-top:6px}
303
+ </style></head><body>
304
+ <div class="wrap">
305
+ <h1>FLUX.2 Klein 9B</h1>
306
+ <div class="sub">Distilled · TE uncensored (best-effort) · NSFW LoRA</div>
307
+ <div class="card">
308
+ <label>Image (optional = edit / I2I)</label>
309
+ <input id="file" type="file" accept="image/*"/>
310
+ </div>
311
+ <div class="card">
312
+ <label>Prompt</label>
313
+ <textarea id="prompt">nude woman standing by a window, natural light, photorealistic skin, detailed</textarea>
314
+ <div class="hint">TE uncensored passa o prompt; LoRA ensina anatomia ao DiT.</div>
315
+ </div>
316
+ <div class="card row">
317
+ <div>
318
+ <label>LoRA</label>
319
+ <select id="lora"><option>NSFW-Solo</option><option>Consistency</option><option>None</option></select>
320
+ </div>
321
+ <div>
322
+ <label>LoRA strength</label>
323
+ <input id="str" type="number" value="0.75" min="0" max="1.5" step="0.05"/>
324
+ </div>
325
+ <div>
326
+ <label>Steps (distilled ≈ 4)</label>
327
+ <input id="steps" type="number" value="4" min="1" max="28"/>
328
+ </div>
329
+ <div>
330
+ <label>Guidance</label>
331
+ <input id="cfg" type="number" value="1.0" min="0" max="8" step="0.1"/>
332
+ </div>
333
+ </div>
334
+ <button id="run">Generate</button>
335
+ <div class="card"><div id="status">Ready.</div><img id="out"/></div>
336
+ </div>
337
+ <script>
338
+ const $=id=>document.getElementById(id);
339
+ function fileToB64(f){return new Promise((res,rej)=>{const r=new FileReader();r.onload=()=>res(r.result);r.onerror=rej;r.readAsDataURL(f);});}
340
+ $("run").onclick=async()=>{
341
+ const btn=$("run");btn.disabled=true;$("status").textContent="Queuing…";$("out").removeAttribute("src");
342
+ try{
343
+ const images=[];const f=$("file").files[0];if(f)images.push(await fileToB64(f));
344
+ const body={data:[JSON.stringify(images),$("prompt").value,$("lora").value,0,true,
345
+ parseFloat($("cfg").value),parseInt($("steps").value,10),parseFloat($("str").value)]};
346
+ const res=await fetch("/gradio_api/call/generate",{method:"POST",headers:{"Content-Type":"application/json"},body:JSON.stringify(body)});
347
+ const ev=await res.json();if(!ev.event_id)throw new Error(JSON.stringify(ev));
348
+ const tr=await fetch(`/gradio_api/call/generate/${ev.event_id}`);
349
+ const text=await tr.text();let payload=null;
350
+ for(const line of text.trim().split("\\n")){if(line.startsWith("data:")){try{payload=JSON.parse(line.slice(5).trim());}catch{}}}
351
+ const out=Array.isArray(payload)?payload[0]:payload;
352
+ if(out&&out.image){$("out").src=out.image;$("status").textContent=`OK · seed=${out.seed} · steps=${out.steps} · lora=${out.lora}@${out.lora_strength}`;}
353
+ else $("status").textContent="Done: "+JSON.stringify(out).slice(0,400);
354
+ }catch(e){$("status").textContent="Error: "+e.message;}finally{btn.disabled=false;}
355
+ };
356
+ </script></body></html>
357
+ """
358
+
359
+
360
+ if __name__ == "__main__":
361
+ print("Klein 9B distilled + NSFW LoRA — accept BFL license on the model page")
362
+ app.launch(show_error=True, mcp_server=True)