""" FLUX.2 Klein 9B (distilled) + optional uncensored TE + NSFW LoRAs Lighter than SDXL UnrealVision / Klein-base-50-steps. """ import os import gc import json import random import base64 from io import BytesIO from pathlib import Path import gradio as gr from gradio import Server from fastapi.responses import HTMLResponse import numpy as np import spaces import torch from PIL import Image MAX_SEED = np.iinfo(np.int32).max LANCZOS = getattr(Image, "Resampling", Image).LANCZOS device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print("device:", device, "cuda:", torch.cuda.is_available()) from diffusers import Flux2KleinPipeline dtype = torch.bfloat16 # Distilled 9B = 4 steps, much lighter than base-9B (28–50 steps) MODEL_ID = os.getenv("KLEIN_MODEL", "black-forest-labs/FLUX.2-klein-9B") USE_UNCENSORED_TE = os.getenv("KLEIN_UNCENSORED_TE", "1").strip() not in ("0", "false", "no") UNCENSORED_TE_REPO = os.getenv( "KLEIN_TE_REPO", "ponpoke/flux2-klein-9b-uncensored-text-encoder", ) ADAPTER_SPECS = { "None": { "repo": None, "weights": None, "adapter_name": None, "default_strength": 0.0, "hint": "No LoRA", }, "NSFW-Solo": { "repo": "diroverflo/FLux_Klein_9B_NSFW", "weights": None, # auto from repo "adapter_name": "nsfw-solo", "default_strength": 0.75, "hint": "Solo NSFW focus (no faces in train set)", }, "Consistency": { "repo": "dx8152/Flux2-Klein-9B-Consistency", "weights": None, "adapter_name": "klein-consistency", "default_strength": 0.8, "hint": "Identity / edit consistency (not NSFW unlock)", }, } LOADED_ADAPTERS: set = set() ADAPTER_NAMES = list(ADAPTER_SPECS.keys()) print(f"Loading Flux2KleinPipeline: {MODEL_ID}") pipe = Flux2KleinPipeline.from_pretrained( MODEL_ID, torch_dtype=dtype, ) if torch.cuda.is_available(): try: pipe.enable_model_cpu_offload() print("cpu_offload OK") except Exception as e: print("offload fail:", e) pipe = pipe.to(device) else: pipe = pipe.to(device) # Best-effort uncensored text encoder swap if USE_UNCENSORED_TE: try: from transformers import AutoModel, AutoTokenizer print(f"[TE] loading uncensored encoder: {UNCENSORED_TE_REPO}") tok = AutoTokenizer.from_pretrained( UNCENSORED_TE_REPO, token=os.getenv("HF_TOKEN") or None, ) te = AutoModel.from_pretrained( UNCENSORED_TE_REPO, torch_dtype=dtype, token=os.getenv("HF_TOKEN") or None, ) # Flux2Klein may expose text_encoder / tokenizer attributes if hasattr(pipe, "text_encoder"): pipe.text_encoder = te if hasattr(pipe, "tokenizer"): pipe.tokenizer = tok print("[TE] uncensored text encoder attached (best-effort)") except Exception as e: print(f"[TE] skip uncensored TE ({type(e).__name__}: {e})") print("[TE] continuing with official encoder + NSFW LoRAs") print("Pipeline ready.") def b64_to_pil_list(b64_json_str): if not b64_json_str or str(b64_json_str).strip() in ("", "[]"): return [] try: b64_list = json.loads(b64_json_str) except Exception: return [] out = [] for b64_str in b64_list: if not b64_str or not isinstance(b64_str, str): continue try: if b64_str.startswith("data:image"): _, data = b64_str.split(",", 1) else: data = b64_str out.append(Image.open(BytesIO(base64.b64decode(data))).convert("RGB")) except Exception as e: print("decode error:", e) return out def pil_to_b64_png(image: Image.Image) -> str: buf = BytesIO() image.save(buf, format="PNG") return f"data:image/png;base64,{base64.b64encode(buf.getvalue()).decode()}" def update_dimensions(image: Image.Image, max_side: int = 1024): if image is None: return 1024, 1024 w, h = image.size if w >= h: nw = max_side nh = int(nw * h / w) else: nh = max_side nw = int(nh * w / h) return max(16, (nw // 16) * 16), max(16, (nh // 16) * 16) app = Server(title="FLUX2-Klein-9B-NSFW") @app.api(name="check_safety", queue=False) def check_safety(prompt: str) -> dict: return {"status": "ok"} @app.api(name="generate") @spaces.GPU(duration=30) def ping(): return "gpu-ok" gr.Interface(ping, None, "text").launch() if not prompt or not str(prompt).strip(): raise gr.Error("Prompt is empty.") lora_adapter = (lora_adapter or "NSFW-Solo").strip() spec = ADAPTER_SPECS.get(lora_adapter) or ADAPTER_SPECS["None"] strength = float(lora_strength if lora_strength is not None else spec["default_strength"]) strength = max(0.0, min(strength, 1.5)) if spec["adapter_name"] is None or strength <= 0: try: pipe.disable_lora() except Exception: pass print("--- LoRA off ---") else: name = spec["adapter_name"] if name not in LOADED_ADAPTERS: print(f"--- Loading LoRA {lora_adapter} from {spec['repo']} ---") try: kwargs = {"adapter_name": name} if spec.get("weights"): kwargs["weight_name"] = spec["weights"] pipe.load_lora_weights(spec["repo"], **kwargs) LOADED_ADAPTERS.add(name) except Exception as e: raise gr.Error(f"LoRA load failed ({lora_adapter}): {e}") try: pipe.set_adapters([name], adapter_weights=[strength]) except Exception: pipe.set_adapters([name]) print(f"--- LoRA {name} @ {strength} ---") if randomize_seed: seed = random.randint(0, MAX_SEED) seed = int(seed) generator = torch.Generator( device="cuda" if torch.cuda.is_available() else "cpu" ).manual_seed(seed) # Distilled Klein defaults: 4 steps, guidance ~1.0–4.0 steps = int(max(1, min(int(steps or 4), 28))) guidance_scale = float(max(0.0, min(float(guidance_scale or 1.0), 8.0))) pil_images = b64_to_pil_list(images_b64_json) if pil_images: width, height = update_dimensions(pil_images[0]) processed = [im.resize((width, height), LANCZOS) for im in pil_images] image_input = processed if len(processed) > 1 else processed[0] print(f"I2I {width}x{height} steps={steps} cfg={guidance_scale}") else: width, height = 1024, 1024 image_input = None print(f"T2I {width}x{height} steps={steps} cfg={guidance_scale}") try: kwargs = dict( prompt=str(prompt).strip(), height=height, width=width, num_inference_steps=steps, guidance_scale=guidance_scale, generator=generator, ) if image_input is not None: kwargs["image"] = image_input result = pipe(**kwargs).images[0] return { "image": pil_to_b64_png(result), "seed": seed, "status": "success", "width": width, "height": height, "steps": steps, "guidance_scale": guidance_scale, "lora": lora_adapter, "lora_strength": strength, } except Exception as e: raise gr.Error(f"Inference failed: {type(e).__name__}: {e}") finally: gc.collect() if torch.cuda.is_available(): torch.cuda.empty_cache() @app.get("/api/config") def client_config(): return { "model": MODEL_ID, "uncensored_te": USE_UNCENSORED_TE, "loras": ADAPTER_NAMES, "defaults": { "steps": 4, "guidance_scale": 1.0, "lora": "NSFW-Solo", "lora_strength": 0.75, }, "notes": [ "Distilled Klein 9B — typically 4 steps", "TE uncensored is best-effort; LoRA teaches NSFW concepts to DiT", "Accept FLUX license on black-forest-labs/FLUX.2-klein-9B", "HF_TOKEN if gated models need auth", ], } @app.get("/", response_class=HTMLResponse) async def homepage(): html_path = Path(__file__).resolve().parent / "index.html" if html_path.exists(): return html_path.read_text(encoding="utf-8") return """