Add explicit HYPIR restoration profile
This commit is contained in:
+135
-18
@@ -138,6 +138,16 @@ IMAGE_WORKER_URL = os.environ.get("IMAGE_WORKER_URL", "").rstrip("/")
|
||||
IMAGE_WORKER_TOKEN = os.environ.get("IMAGE_WORKER_TOKEN", "").strip()
|
||||
IMAGE_MODEL_NAME = os.environ.get(
|
||||
"IMAGE_MODEL_NAME", "FLUX.2-klein-9B-fp8-beta")
|
||||
RESTORATION_WORKER_URL = os.environ.get(
|
||||
"RESTORATION_WORKER_URL", "").rstrip("/")
|
||||
RESTORATION_WORKER_TOKEN = os.environ.get(
|
||||
"RESTORATION_WORKER_TOKEN", IMAGE_WORKER_TOKEN).strip()
|
||||
RESTORATION_MODEL_NAME = os.environ.get(
|
||||
"RESTORATION_MODEL_NAME", "HYPIR-SD2")
|
||||
RESTORATION_CHAT_MODEL = os.environ.get(
|
||||
"RESTORATION_CHAT_MODEL", "restauration").strip()
|
||||
RESTORATION_CHAT_PROFILE = os.environ.get(
|
||||
"RESTORATION_CHAT_PROFILE", "fast").strip()
|
||||
IMAGE_DIR = os.environ.get(
|
||||
"IMAGE_DIR", "/opt/mike-ai/ai-profile-router/images")
|
||||
IMAGE_WORKER_LOG = os.environ.get(
|
||||
@@ -215,6 +225,12 @@ VIRTUAL_MODELS = {
|
||||
(EXPECTED_MODELS.get(name) or f"qwen-{name}"): name
|
||||
for name in PROFILES
|
||||
}
|
||||
if RESTORATION_CHAT_PROFILE not in PROFILES:
|
||||
raise ConfigurationError(
|
||||
f"RESTORATION_CHAT_PROFILE ist unbekannt: {RESTORATION_CHAT_PROFILE!r}")
|
||||
if not RESTORATION_CHAT_MODEL or RESTORATION_CHAT_MODEL in VIRTUAL_MODELS:
|
||||
raise ConfigurationError("RESTORATION_CHAT_MODEL fehlt oder kollidiert")
|
||||
VIRTUAL_MODELS[RESTORATION_CHAT_MODEL] = RESTORATION_CHAT_PROFILE
|
||||
|
||||
log = logging.getLogger("ai-profile-router")
|
||||
AUTH: AuthPolicy | None = None
|
||||
@@ -256,6 +272,7 @@ class _ImageState:
|
||||
self.last_error: str | None = None
|
||||
self.last_image: str | None = None
|
||||
self.last_seconds: float | None = None
|
||||
self.current_model: str | None = None
|
||||
|
||||
|
||||
IMAGE_PHASES = (
|
||||
@@ -857,13 +874,23 @@ class _Worker:
|
||||
RUNTIME.clear_worker("image")
|
||||
|
||||
|
||||
def _worker() -> _Worker:
|
||||
def _worker(model: str = IMAGE_MODEL_NAME) -> _Worker:
|
||||
"""Worker-Instanz liefern (startet bei Bedarf)."""
|
||||
img = STATE.image
|
||||
if not img.worker or not img.worker.alive():
|
||||
if img.worker:
|
||||
img.worker.stop()
|
||||
img.worker = _RemoteWorker() if IMAGE_WORKER_URL else _Worker()
|
||||
if model == RESTORATION_MODEL_NAME:
|
||||
if not RESTORATION_WORKER_URL:
|
||||
raise RuntimeError("Restaurations-Worker ist nicht konfiguriert")
|
||||
img.worker = _RemoteWorker(
|
||||
kind="restore", url=RESTORATION_WORKER_URL,
|
||||
token=RESTORATION_WORKER_TOKEN, endpoint="/restore")
|
||||
else:
|
||||
img.worker = (_RemoteWorker(kind="image", url=IMAGE_WORKER_URL,
|
||||
token=IMAGE_WORKER_TOKEN,
|
||||
endpoint="/generate")
|
||||
if IMAGE_WORKER_URL else _Worker())
|
||||
img.worker.start()
|
||||
return img.worker
|
||||
|
||||
@@ -873,7 +900,12 @@ class _RemoteWorker:
|
||||
|
||||
model_loaded = False
|
||||
|
||||
def __init__(self) -> None:
|
||||
def __init__(self, *, kind: str, url: str, token: str,
|
||||
endpoint: str) -> None:
|
||||
self.kind = kind
|
||||
self.url = url
|
||||
self.token = token
|
||||
self.endpoint = endpoint
|
||||
self.running = False
|
||||
|
||||
def alive(self) -> bool:
|
||||
@@ -882,10 +914,10 @@ class _RemoteWorker:
|
||||
def _request(self, method: str, path: str, payload: dict | None = None,
|
||||
timeout: float = 120) -> dict:
|
||||
body = None if payload is None else json.dumps(payload).encode()
|
||||
headers = {"Authorization": f"Bearer {IMAGE_WORKER_TOKEN}"}
|
||||
headers = {"Authorization": f"Bearer {self.token}"}
|
||||
if body is not None:
|
||||
headers["Content-Type"] = "application/json"
|
||||
req = urllib.request.Request(IMAGE_WORKER_URL + path, data=body,
|
||||
req = urllib.request.Request(self.url + path, data=body,
|
||||
method=method, headers=headers)
|
||||
try:
|
||||
with urllib.request.urlopen(req, timeout=timeout) as response:
|
||||
@@ -900,9 +932,9 @@ class _RemoteWorker:
|
||||
raise RuntimeError(f"Bild-Worker nicht erreichbar: {exc}") from exc
|
||||
|
||||
def start(self) -> None:
|
||||
if not IMAGE_WORKER_TOKEN or len(IMAGE_WORKER_TOKEN) < 32:
|
||||
if not self.token or len(self.token) < 32:
|
||||
raise RuntimeError("Bild-Worker-Token fehlt oder ist zu kurz")
|
||||
_profile_controller_request("POST", "/workers/image/start")
|
||||
_profile_controller_request("POST", f"/workers/{self.kind}/start")
|
||||
deadline = time.monotonic() + IMAGE_START_TIMEOUT
|
||||
while time.monotonic() < deadline:
|
||||
try:
|
||||
@@ -922,11 +954,11 @@ class _RemoteWorker:
|
||||
clean.pop("cmd", None)
|
||||
output = clean.pop("output", "")
|
||||
clean["filename"] = os.path.basename(output)
|
||||
return self._request("POST", "/generate", clean, timeout)
|
||||
return self._request("POST", self.endpoint, clean, timeout)
|
||||
|
||||
def stop(self) -> None:
|
||||
try:
|
||||
_profile_controller_request("POST", "/workers/image/stop")
|
||||
_profile_controller_request("POST", f"/workers/{self.kind}/stop")
|
||||
finally:
|
||||
self.running = False
|
||||
self.model_loaded = False
|
||||
@@ -1013,6 +1045,8 @@ def generate_image(prompt: str, width: int, height: int, steps: int,
|
||||
guidance: float, seed: int | None, n: int,
|
||||
quality: str = "standard",
|
||||
source_files: list[str] | None = None,
|
||||
model: str = IMAGE_MODEL_NAME,
|
||||
restore_options: dict | None = None,
|
||||
) -> tuple[list[str], str | None]:
|
||||
"""Orchestriert die Bildgenerierung inkl. Qwen-Hotswap.
|
||||
|
||||
@@ -1032,6 +1066,7 @@ def generate_image(prompt: str, width: int, height: int, steps: int,
|
||||
results: list[str] = []
|
||||
warning: str | None = None
|
||||
img.last_error = None
|
||||
img.current_model = model
|
||||
# Qwen wird gestoppt → für Chats nicht verfügbar (die warten).
|
||||
_set_qwen_unavailable(True)
|
||||
try:
|
||||
@@ -1056,7 +1091,7 @@ def generate_image(prompt: str, width: int, height: int, steps: int,
|
||||
|
||||
# 2) Worker starten (Modell wird beim ersten generate geladen).
|
||||
img.phase = "loading-image"
|
||||
worker = _worker()
|
||||
worker = _worker(model)
|
||||
|
||||
# 3) Generieren.
|
||||
for i in range(n):
|
||||
@@ -1064,7 +1099,7 @@ def generate_image(prompt: str, width: int, height: int, steps: int,
|
||||
filename = time.strftime("%Y%m%d-%H%M%S") + \
|
||||
f"-{os.urandom(2).hex()}.png"
|
||||
output = os.path.join(IMAGE_DIR, filename)
|
||||
resp = worker.request({
|
||||
worker_payload = {
|
||||
"cmd": "generate",
|
||||
"prompt": prompt,
|
||||
"width": width,
|
||||
@@ -1074,7 +1109,9 @@ def generate_image(prompt: str, width: int, height: int, steps: int,
|
||||
"seed": seed,
|
||||
"output": output,
|
||||
"source_files": source_files or [],
|
||||
}, timeout=IMAGE_GEN_TIMEOUT)
|
||||
}
|
||||
worker_payload.update(restore_options or {})
|
||||
resp = worker.request(worker_payload, timeout=IMAGE_GEN_TIMEOUT)
|
||||
if resp.get("status") != "ok":
|
||||
raise RuntimeError(
|
||||
resp.get("message", "Bildgenerierung fehlgeschlagen"))
|
||||
@@ -1092,10 +1129,11 @@ def generate_image(prompt: str, width: int, height: int, steps: int,
|
||||
"steps": steps,
|
||||
"guidance": guidance,
|
||||
"quality": quality,
|
||||
"mode": "image-edit" if source_files else "text-to-image",
|
||||
"mode": ("image-restoration" if model == RESTORATION_MODEL_NAME
|
||||
else "image-edit" if source_files else "text-to-image"),
|
||||
"reference_images": len(source_files or []),
|
||||
"seconds": resp.get("seconds"),
|
||||
"model": IMAGE_MODEL_NAME,
|
||||
"model": model,
|
||||
"created": time.strftime("%Y-%m-%dT%H:%M:%S"),
|
||||
}
|
||||
meta_path = os.path.join(IMAGE_DIR, filename[:-4] + ".json")
|
||||
@@ -1141,6 +1179,7 @@ def generate_image(prompt: str, width: int, height: int, steps: int,
|
||||
log.error(warning)
|
||||
# Qwen ist down → qwen_unavailable bleibt True.
|
||||
img.phase = "idle"
|
||||
img.current_model = None
|
||||
return results, warning
|
||||
|
||||
|
||||
@@ -1344,6 +1383,28 @@ def _inject_global_system_policy(data: dict, path: str) -> dict:
|
||||
return data
|
||||
|
||||
|
||||
def _inject_restoration_system_policy(data: dict, path: str) -> dict:
|
||||
"""Make the explicitly selected restoration model use the image tool."""
|
||||
policy = (
|
||||
"Photo-restoration mode is selected. When the user supplies an image, "
|
||||
"use the image generation/editing tool exactly once with that source "
|
||||
"image and the user's requested restoration. Preserve identity, "
|
||||
"anatomy, pose, composition and objects unless the user explicitly "
|
||||
"asks for a creative change. Do not attempt restoration with Python, "
|
||||
"PIL, OpenCV or shell tools."
|
||||
)
|
||||
if path == "/v1/chat/completions":
|
||||
messages = data.get("messages")
|
||||
if isinstance(messages, list):
|
||||
messages.insert(0, {"role": "system", "content": policy})
|
||||
elif path == "/v1/responses":
|
||||
instructions = data.get("instructions")
|
||||
data["instructions"] = (
|
||||
f"{policy}\n\n{instructions}"
|
||||
if isinstance(instructions, str) and instructions else policy)
|
||||
return data
|
||||
|
||||
|
||||
def _normalize_llamacpp_reasoning(data: dict) -> dict:
|
||||
"""Mappt OpenAI/Hermes-Reasoning auf llama.cpp-Template-Parameter.
|
||||
|
||||
@@ -1696,6 +1757,15 @@ class Handler(BaseHTTPRequestHandler):
|
||||
}
|
||||
for name, ctx in PROFILES.items()
|
||||
]
|
||||
models.append({
|
||||
"id": RESTORATION_CHAT_MODEL,
|
||||
"object": "model",
|
||||
"created": 0,
|
||||
"owned_by": "ai-profile-router",
|
||||
"context_length": PROFILES[RESTORATION_CHAT_PROFILE],
|
||||
"context_window": PROFILES[RESTORATION_CHAT_PROFILE],
|
||||
"purpose": "image-restoration",
|
||||
})
|
||||
if REVIEW_UPSTREAM_URL:
|
||||
models.append({
|
||||
"id": REVIEW_MODEL_NAME,
|
||||
@@ -1743,7 +1813,7 @@ class Handler(BaseHTTPRequestHandler):
|
||||
"phase": img.phase,
|
||||
"worker": "running" if (img.worker and img.worker.alive())
|
||||
else "stopped",
|
||||
"model": IMAGE_MODEL_NAME if img.phase != "idle" else None,
|
||||
"model": img.current_model if img.phase != "idle" else None,
|
||||
"model_loaded": bool(img.worker and img.worker.model_loaded),
|
||||
"last_image": img.last_image,
|
||||
"last_seconds": img.last_seconds,
|
||||
@@ -1850,6 +1920,19 @@ class Handler(BaseHTTPRequestHandler):
|
||||
"invalid_request_error", "prompt_too_long")
|
||||
return
|
||||
|
||||
model = data.get("model", IMAGE_MODEL_NAME)
|
||||
if model not in {IMAGE_MODEL_NAME, RESTORATION_MODEL_NAME}:
|
||||
self._send_error(
|
||||
400, f"unbekanntes Bildmodell: {model!r}",
|
||||
"invalid_request_error", "invalid_model")
|
||||
return
|
||||
restoring = model == RESTORATION_MODEL_NAME
|
||||
if restoring and len(source_files) != 1:
|
||||
self._send_error(
|
||||
400, f"{RESTORATION_MODEL_NAME} benötigt genau ein Referenzbild",
|
||||
"invalid_request_error", "missing_image")
|
||||
return
|
||||
|
||||
# Größe
|
||||
size = data.get("size", "1024x1024")
|
||||
if size not in IMAGE_SIZES:
|
||||
@@ -1866,6 +1949,10 @@ class Handler(BaseHTTPRequestHandler):
|
||||
self._send_error(400, f"'n' muss eine Ganzzahl 1..{IMAGE_MAX_N} sein",
|
||||
"invalid_request_error", "invalid_n")
|
||||
return
|
||||
if restoring and n != 1:
|
||||
self._send_error(400, "Bildrestaurierung unterstützt nur 'n'=1",
|
||||
"invalid_request_error", "invalid_n")
|
||||
return
|
||||
|
||||
# Qualität / Schritte / Guidance
|
||||
quality = data.get("quality", IMAGE_DEFAULT_QUALITY)
|
||||
@@ -1875,7 +1962,9 @@ class Handler(BaseHTTPRequestHandler):
|
||||
"invalid_request_error", "invalid_quality")
|
||||
return
|
||||
steps = data.get("steps", IMAGE_QUALITY[quality])
|
||||
if not isinstance(steps, int) or isinstance(steps, bool) or steps != 4:
|
||||
if (not restoring and
|
||||
(not isinstance(steps, int) or isinstance(steps, bool)
|
||||
or steps != 4)):
|
||||
self._send_error(400, f"{IMAGE_MODEL_NAME} erfordert 'steps'=4",
|
||||
"invalid_request_error", "invalid_steps")
|
||||
return
|
||||
@@ -1886,11 +1975,35 @@ class Handler(BaseHTTPRequestHandler):
|
||||
self._send_error(400, "'guidance' muss eine Zahl sein",
|
||||
"invalid_request_error", "invalid_guidance")
|
||||
return
|
||||
if guidance != 1.0:
|
||||
if not restoring and guidance != 1.0:
|
||||
self._send_error(400, f"{IMAGE_MODEL_NAME} erfordert 'guidance'=1.0",
|
||||
"invalid_request_error", "invalid_guidance")
|
||||
return
|
||||
|
||||
restore_options: dict = {}
|
||||
if restoring:
|
||||
try:
|
||||
upscale = int(data.get("upscale", 1))
|
||||
patch_size = int(data.get("patch_size", 512))
|
||||
stride = int(data.get("stride", 256))
|
||||
except (TypeError, ValueError):
|
||||
self._send_error(400, "ungültige Restaurationsparameter",
|
||||
"invalid_request_error", "invalid_restore_options")
|
||||
return
|
||||
if upscale not in (1, 2, 4):
|
||||
self._send_error(400, "'upscale' muss 1, 2 oder 4 sein",
|
||||
"invalid_request_error", "invalid_upscale")
|
||||
return
|
||||
if patch_size not in (512, 768, 1024) or not 0 < stride <= patch_size:
|
||||
self._send_error(400, "ungültige patch_size/stride-Kombination",
|
||||
"invalid_request_error", "invalid_tiling")
|
||||
return
|
||||
restore_options = {
|
||||
"upscale": upscale,
|
||||
"patch_size": patch_size,
|
||||
"stride": stride,
|
||||
}
|
||||
|
||||
seed = data.get("seed")
|
||||
if seed is not None:
|
||||
try:
|
||||
@@ -1915,7 +2028,7 @@ class Handler(BaseHTTPRequestHandler):
|
||||
try:
|
||||
results, warning = generate_image(
|
||||
prompt.strip(), width, height, steps, guidance, seed, n,
|
||||
quality, source_files)
|
||||
quality, source_files, model, restore_options)
|
||||
except (ValueError, RuntimeError) as e:
|
||||
self._send_error(503, str(e), "server_error", "image_generation_failed")
|
||||
return
|
||||
@@ -2339,6 +2452,7 @@ class Handler(BaseHTTPRequestHandler):
|
||||
data = None
|
||||
requested_profile: str | None = None
|
||||
requested_review = False
|
||||
requested_restoration = False
|
||||
# Virtuelles Modell erkennen. Umschalten und Chat-Lease werden weiter
|
||||
# unten atomar unter dem zentralen Orchestrierungs-Lock ausgeführt.
|
||||
if body is not None and self.path.startswith("/v1/"):
|
||||
@@ -2352,6 +2466,7 @@ class Handler(BaseHTTPRequestHandler):
|
||||
requested_review = True
|
||||
elif isinstance(model, str) and model in VIRTUAL_MODELS:
|
||||
requested_profile = VIRTUAL_MODELS[model]
|
||||
requested_restoration = model == RESTORATION_CHAT_MODEL
|
||||
elif isinstance(model, str) and model.startswith("qwen-"):
|
||||
# qwen-* ist der Namensraum des Routers
|
||||
self._send_error(400, f"unbekanntes virtuelles Modell: {model}",
|
||||
@@ -2361,6 +2476,8 @@ class Handler(BaseHTTPRequestHandler):
|
||||
if path in {"/v1/chat/completions", "/v1/responses"}:
|
||||
try:
|
||||
data = _inject_global_system_policy(data, path)
|
||||
if requested_restoration:
|
||||
data = _inject_restoration_system_policy(data, path)
|
||||
except ValueError as exc:
|
||||
self._send_error(500, str(exc), "server_error",
|
||||
"system_policy_unavailable")
|
||||
|
||||
Reference in New Issue
Block a user