174 lines
5.2 KiB
Python
174 lines
5.2 KiB
Python
#!/usr/bin/env python3
|
||
"""FLUX.2 [klein] 4B – Bild-Worker with reference-image editing.
|
||
|
||
Protokoll: zeilenbasiertes JSON über stdin/stdout.
|
||
|
||
Start: Worker gibt {"status": "ready"} aus (Modell noch NICHT geladen).
|
||
Request: {"cmd": "generate", "prompt": ..., "width": ..., "height": ...,
|
||
"steps": ..., "guidance": ..., "seed": ..., "output": ...}
|
||
Antwort: {"status": "ok", "path": ..., "seconds": ..., "load_seconds": ...}
|
||
oder {"status": "error", "message": ...}
|
||
Request: {"cmd": "unload"} -> {"status": "ok"}
|
||
Request: {"cmd": "status"} -> {"status": "ok", "model_loaded": bool}
|
||
|
||
Das Modell wird beim ersten generate geladen (bf16, cpu_offload) und auf
|
||
Anforderung wieder entladen (VRAM freigeben). Der Prozess bleibt danach
|
||
laufen – ohne geladenes Modell belegt er kaum Ressourcen.
|
||
|
||
Alle torch-/diffusers-Logs gehen nach stderr, stdout ist reines Protokoll.
|
||
"""
|
||
|
||
import gc
|
||
import json
|
||
import os
|
||
import signal
|
||
import sys
|
||
import time
|
||
|
||
# stderr-Logs von torch & Co. unterdrücken, bevor importiert wird
|
||
os.environ.setdefault("DIFFUSERS_VERBOSITY", "error")
|
||
os.environ.setdefault("TRANSFORMERS_VERBOSITY", "error")
|
||
os.environ.setdefault("HF_HUB_DISABLE_PROGRESS_BARS", "1")
|
||
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
|
||
|
||
MODEL_DIR = os.environ.get(
|
||
"FLUX_MODEL_DIR", "/opt/mike-ai/models/FLUX.2-klein-4B")
|
||
|
||
_pipe = None # geladene Pipeline (None = entladen)
|
||
_load_seconds = 0.0 # Dauer des letzten Ladens
|
||
|
||
|
||
def _emit(payload: dict) -> None:
|
||
sys.stdout.write(json.dumps(payload) + "\n")
|
||
sys.stdout.flush()
|
||
|
||
|
||
def _log(msg: str) -> None:
|
||
print(f"[image-worker] {msg}", file=sys.stderr, flush=True)
|
||
|
||
|
||
def _load() -> None:
|
||
"""Pipeline laden (bf16, CPU-Offload)."""
|
||
global _pipe, _load_seconds
|
||
if _pipe is not None:
|
||
return
|
||
import torch
|
||
from diffusers import Flux2KleinPipeline
|
||
|
||
t0 = time.monotonic()
|
||
_log(f"lade Modell aus {MODEL_DIR} ...")
|
||
_pipe = Flux2KleinPipeline.from_pretrained(
|
||
MODEL_DIR, torch_dtype=torch.bfloat16)
|
||
_pipe.enable_model_cpu_offload()
|
||
_load_seconds = time.monotonic() - t0
|
||
_log(f"Modell geladen in {_load_seconds:.1f} s")
|
||
|
||
|
||
def _unload() -> None:
|
||
"""Pipeline entladen und VRAM freigeben."""
|
||
global _pipe
|
||
if _pipe is None:
|
||
return
|
||
t0 = time.monotonic()
|
||
del _pipe
|
||
_pipe = None
|
||
gc.collect()
|
||
try:
|
||
import torch
|
||
torch.cuda.empty_cache()
|
||
except Exception:
|
||
pass
|
||
_log(f"Modell entladen in {time.monotonic() - t0:.1f} s")
|
||
|
||
|
||
def _generate(req: dict) -> dict:
|
||
import torch
|
||
from PIL import Image
|
||
|
||
prompt = req["prompt"]
|
||
width = int(req.get("width", 1024))
|
||
height = int(req.get("height", 1024))
|
||
steps = int(req.get("steps", 4))
|
||
guidance = float(req.get("guidance", 1.0))
|
||
seed = req.get("seed")
|
||
output = req["output"]
|
||
|
||
_load()
|
||
|
||
t0 = time.monotonic()
|
||
generator = None
|
||
if seed is not None:
|
||
generator = torch.Generator(device="cuda").manual_seed(int(seed))
|
||
kwargs = dict(
|
||
prompt=prompt,
|
||
height=height,
|
||
width=width,
|
||
guidance_scale=guidance,
|
||
num_inference_steps=steps,
|
||
generator=generator,
|
||
)
|
||
source_files = req.get("source_files") or []
|
||
if not isinstance(source_files, list) or len(source_files) > 4:
|
||
raise ValueError("invalid source image list")
|
||
sources = []
|
||
for source in source_files:
|
||
if not isinstance(source, str):
|
||
raise ValueError("invalid source image filename")
|
||
path = os.path.join(os.path.dirname(output), source)
|
||
with Image.open(path) as opened:
|
||
sources.append(opened.convert("RGB"))
|
||
if sources:
|
||
kwargs["image"] = sources[0] if len(sources) == 1 else sources
|
||
image = _pipe(**kwargs).images[0]
|
||
|
||
os.makedirs(os.path.dirname(output) or ".", exist_ok=True)
|
||
image.save(output)
|
||
seconds = time.monotonic() - t0
|
||
_log(f"generiert {output} in {seconds:.1f} s "
|
||
f"({width}x{height}, {steps} steps, seed={seed})")
|
||
return {
|
||
"status": "ok",
|
||
"path": output,
|
||
"seconds": round(seconds, 2),
|
||
"load_seconds": round(_load_seconds, 2),
|
||
}
|
||
|
||
|
||
def _handle(line: str) -> None:
|
||
try:
|
||
req = json.loads(line)
|
||
except ValueError:
|
||
_emit({"status": "error", "message": "ungültiges JSON"})
|
||
return
|
||
|
||
cmd = req.get("cmd")
|
||
try:
|
||
if cmd == "generate":
|
||
_emit(_generate(req))
|
||
elif cmd == "unload":
|
||
_unload()
|
||
_emit({"status": "ok"})
|
||
elif cmd == "status":
|
||
_emit({"status": "ok", "model_loaded": _pipe is not None})
|
||
else:
|
||
_emit({"status": "error", "message": f"unbekanntes Kommando: {cmd}"})
|
||
except Exception as e: # noqa: BLE001 – Fehler ans Router-Protokoll
|
||
_log(f"Fehler bei {cmd}: {e!r}")
|
||
_emit({"status": "error", "message": str(e)})
|
||
|
||
|
||
def main() -> None:
|
||
signal.signal(signal.SIGTERM, lambda *_: sys.exit(0))
|
||
_emit({"status": "ready"})
|
||
for line in sys.stdin:
|
||
line = line.strip()
|
||
if not line:
|
||
continue
|
||
_handle(line)
|
||
if _pipe is None and line.startswith('{"cmd": "unload"'):
|
||
pass # Worker bleibt laufen, Modell ist entladen
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|