From 118e32005eebee526f9bcc66c0cac94e8fc961ac Mon Sep 17 00:00:00 2001 From: Mikei386 <44135113+Mikei386@users.noreply.github.com> Date: Mon, 7 Sep 2026 22:52:59 +0200 Subject: [PATCH] Add explicit HYPIR restoration profile --- compose.yaml | 41 ++++ config/install.env.example | 3 + dev/test_profile_controller.py | 36 +++- dev/test_router_support.py | 25 +++ docs/IMAGE_RESTORATION.md | 71 +++++++ integrations/hermes-athena-image/README.md | 7 + integrations/hermes-athena-image/__init__.py | 43 +++- .../profile-controller/profile_controller.py | 40 +++- platform/docker/restoration-worker/Dockerfile | 20 ++ .../restoration-worker/restoration_worker.py | 186 ++++++++++++++++++ router/ai_profile_router.py | 153 ++++++++++++-- 11 files changed, 593 insertions(+), 32 deletions(-) create mode 100644 docs/IMAGE_RESTORATION.md create mode 100644 platform/docker/restoration-worker/Dockerfile create mode 100644 platform/docker/restoration-worker/restoration_worker.py diff --git a/compose.yaml b/compose.yaml index 2f3f787..47bcab6 100644 --- a/compose.yaml +++ b/compose.yaml @@ -571,6 +571,7 @@ services: CONTROLLER_TOKEN: "${CONTROLLER_TOKEN:?CONTROLLER_TOKEN is required}" ALLOWED_PROFILES: fast,medium,beta1,large,ultra,uncensored IMAGE_WORKER: image + RESTORE_WORKER: restore TTS_WORKER: qwen3 networks: [control] security_opt: ["no-new-privileges:true"] @@ -617,6 +618,11 @@ services: IMAGE_WORKER_URL: http://image-worker:8086 IMAGE_WORKER_TOKEN: "${CONTROLLER_TOKEN:?CONTROLLER_TOKEN is required}" IMAGE_MODEL_NAME: FLUX.2-klein-9B-fp8-beta + RESTORATION_WORKER_URL: http://restoration-worker:8087 + RESTORATION_WORKER_TOKEN: "${CONTROLLER_TOKEN:?CONTROLLER_TOKEN is required}" + RESTORATION_MODEL_NAME: HYPIR-SD2 + RESTORATION_CHAT_MODEL: restauration + RESTORATION_CHAT_PROFILE: fast CHAT_IMAGE_ALLOW_REMOTE_URLS: "false" ENABLE_IMAGE_GENERATION: "true" ENABLE_TTS: "true" @@ -693,6 +699,41 @@ services: timeout: 3s retries: 12 + restoration-worker: + build: + context: platform/docker/restoration-worker + args: + HYPIR_COMMIT: b61d107c6cef38f01a93c7833558869731cfa8c1 + image: mike-ai/restoration-worker:hypir-sd2 + container_name: mike-ai-restoration-worker + restart: "no" + profiles: [image] + labels: + com.mike-ai.image-worker: restore + gpus: all + read_only: true + tmpfs: ["/tmp:size=1g,mode=1777"] + volumes: + - "${HYPIR_MODEL_DIR:-/data/models/HYPIR}:/models/hypir:ro" + - "${HYPIR_SD2_DIR:-/data/models/stable-diffusion-2-1-base}:/models/sd2:ro" + - router-images:/data/images + environment: + NVIDIA_VISIBLE_DEVICES: "${RESTORATION_GPU_DEVICE:-1}" + NVIDIA_DRIVER_CAPABILITIES: compute,utility + WORKER_TOKEN: "${CONTROLLER_TOKEN:?CONTROLLER_TOKEN is required}" + HYPIR_BASE_MODEL: /models/sd2 + HYPIR_WEIGHT_FILE: /models/hypir/HYPIR_sd2.pth + HYPIR_DEVICE: cuda:0 + IMAGE_DIR: /data/images + networks: [inference] + security_opt: ["no-new-privileges:true"] + cap_drop: [ALL] + healthcheck: + test: [CMD, python, -c, "import urllib.request; urllib.request.urlopen('http://127.0.0.1:8087/health', timeout=2)"] + interval: 5s + timeout: 3s + retries: 12 + piper: build: context: platform/docker/piper diff --git a/config/install.env.example b/config/install.env.example index 2397581..4491f39 100644 --- a/config/install.env.example +++ b/config/install.env.example @@ -23,6 +23,9 @@ IMAGE_GPU_DEVICES=GPU-8ad38c6c-5a01-9d8e-1dfa-ed662ad78fbe HF_TOKEN_FILE=/root/.cache/huggingface/token FLUX_COMPONENT_DIR=/data/models/FLUX.2-klein-9B-components FLUX_TRANSFORMER_DIR=/data/models/FLUX.2-klein-9B-fp8 +HYPIR_MODEL_DIR=/data/models/HYPIR +HYPIR_SD2_DIR=/data/models/stable-diffusion-2-1-base +RESTORATION_GPU_DEVICE=1 # Headless remote reachability. Firmware power-loss recovery is configured # separately once at the physical machine. diff --git a/dev/test_profile_controller.py b/dev/test_profile_controller.py index f8592ba..96d9ab8 100644 --- a/dev/test_profile_controller.py +++ b/dev/test_profile_controller.py @@ -26,6 +26,11 @@ def image_item(state="exited"): "Labels": {controller.IMAGE_LABEL_KEY: controller.IMAGE_WORKER}} +def restore_item(state="exited"): + return {"Id": "id-restore", "State": state, + "Labels": {controller.IMAGE_LABEL_KEY: controller.RESTORE_WORKER}} + + def tts_item(state="running"): return {"Id": "id-tts", "State": state, "Labels": {controller.TTS_LABEL_KEY: controller.TTS_WORKER}} @@ -48,7 +53,7 @@ class ProfileControllerTests(unittest.TestCase): return 204, b"" with patch.object(controller, "containers", return_value=profiles), \ - patch.object(controller, "image_container", return_value=image_item()), \ + patch.object(controller, "image_containers", return_value=[image_item()]), \ patch.object(controller, "tts_container", return_value=tts_item()), \ patch.object(controller, "docker_request", side_effect=request): result = controller.activate("medium") @@ -62,7 +67,7 @@ class ProfileControllerTests(unittest.TestCase): def test_fails_if_profile_container_is_missing(self): profiles = {name: item(name) for name in controller.ALLOWED[:-1]} with patch.object(controller, "containers", return_value=profiles), \ - patch.object(controller, "image_container", return_value=image_item()), \ + patch.object(controller, "image_containers", return_value=[image_item()]), \ patch.object(controller, "tts_container", return_value=tts_item()): with self.assertRaisesRegex(RuntimeError, "missing"): controller.activate("fast") @@ -78,6 +83,8 @@ class ProfileControllerTests(unittest.TestCase): with patch.object(controller, "containers", return_value=profiles), \ patch.object(controller, "image_container", return_value=image_item()), \ + patch.object(controller, "image_containers", + return_value=[image_item(), restore_item()]), \ patch.object(controller, "tts_container", return_value=tts_item()), \ patch.object(controller, "docker_request", side_effect=request): controller.set_image_worker(True) @@ -87,6 +94,27 @@ class ProfileControllerTests(unittest.TestCase): ("POST", "/containers/id-flux/start"), ]) + def test_restore_start_stops_flux_and_starts_restore(self): + profiles = {name: item(name) for name in controller.ALLOWED} + calls = [] + + def request(method, path): + calls.append((method, path)) + return 204, b"" + + with patch.object(controller, "containers", return_value=profiles), \ + patch.object(controller, "image_container", return_value=restore_item()), \ + patch.object(controller, "image_containers", + return_value=[image_item("running"), restore_item()]), \ + patch.object(controller, "tts_container", return_value=tts_item()), \ + patch.object(controller, "docker_request", side_effect=request): + controller.set_image_worker(True, controller.RESTORE_WORKER) + self.assertEqual(calls, [ + ("POST", "/containers/id-tts/stop?t=30"), + ("POST", "/containers/id-flux/stop?t=20"), + ("POST", "/containers/id-restore/start"), + ]) + def test_profile_activation_stops_image_worker_first(self): profiles = {name: item(name) for name in controller.ALLOWED} calls = [] @@ -96,8 +124,8 @@ class ProfileControllerTests(unittest.TestCase): return 204, b"" with patch.object(controller, "containers", return_value=profiles), \ - patch.object(controller, "image_container", - return_value=image_item("running")), \ + patch.object(controller, "image_containers", + return_value=[image_item("running"), restore_item()]), \ patch.object(controller, "tts_container", return_value=tts_item()), \ patch.object(controller, "docker_request", side_effect=request): controller.activate("fast") diff --git a/dev/test_router_support.py b/dev/test_router_support.py index 6029576..872b54e 100644 --- a/dev/test_router_support.py +++ b/dev/test_router_support.py @@ -24,10 +24,13 @@ from router_support import ( # noqa: E402 load_profile_registry, ) from ai_profile_router import ( # noqa: E402 + RESTORATION_CHAT_MODEL, + VIRTUAL_MODELS, STATE, _cap_chat_generation, _context_matches, _inject_global_system_policy, + _inject_restoration_system_policy, _normalize_chat_image, _normalize_chat_images, _normalize_llamacpp_reasoning, @@ -80,6 +83,10 @@ class RuntimeStoreTests(unittest.TestCase): class ProfileRegistryTests(unittest.TestCase): + def test_restoration_model_maps_to_fast_instruction_profile(self) -> None: + self.assertEqual(RESTORATION_CHAT_MODEL, "restauration") + self.assertEqual(VIRTUAL_MODELS[RESTORATION_CHAT_MODEL], "fast") + def test_fallback_contains_uncensored_profile(self) -> None: registry = load_profile_registry(None) self.assertEqual(registry["uncensored"]["context"], 80000) @@ -270,6 +277,24 @@ class GlobalSystemPolicyTests(unittest.TestCase): self._inject(request, "/v1/images/generations"), request) +class RestorationSystemPolicyTests(unittest.TestCase): + def test_chat_policy_requires_image_tool_and_preservation(self) -> None: + request = {"messages": [{"role": "user", "content": "Mach schöner"}]} + normalized = _inject_restoration_system_policy( + request, "/v1/chat/completions") + policy = normalized["messages"][0]["content"] + self.assertIn("image generation/editing tool", policy) + self.assertIn("Preserve identity", policy) + self.assertIn("Do not attempt restoration with Python", policy) + + def test_responses_policy_keeps_client_instructions(self) -> None: + request = {"instructions": "Client policy.", "input": "Mach schöner"} + normalized = _inject_restoration_system_policy( + request, "/v1/responses") + self.assertIn("Photo-restoration mode", normalized["instructions"]) + self.assertTrue(normalized["instructions"].endswith("Client policy.")) + + class RetentionTests(unittest.TestCase): def test_oldest_pairs_are_removed(self) -> None: with tempfile.TemporaryDirectory() as temp: diff --git a/docs/IMAGE_RESTORATION.md b/docs/IMAGE_RESTORATION.md new file mode 100644 index 0000000..1835e15 --- /dev/null +++ b/docs/IMAGE_RESTORATION.md @@ -0,0 +1,71 @@ +# GPU-Bildrestaurierung auf Athena + +Stand: 7. September 2026 + +## Zweck + +Athena stellt zusätzlich zur kreativen FLUX-Bildgenerierung eine konservative +Fotorestaurierung mit `HYPIR-SD2` bereit. HYPIR wird ausschließlich für +Referenzbilder verwendet und soll Rauschen, Unschärfe und Kompressionsschäden +reduzieren, ohne die Szene wie ein Text-zu-Bild-Modell vollständig neu zu +erfinden. + +## Aufruf und Routing + +Der OpenAI-kompatible Endpunkt bleibt `/v1/images/edits`. Das gewünschte +Backend wird über `model` gewählt: + +```json +{ + "model": "HYPIR-SD2", + "prompt": "Restauriere dieses Foto möglichst originalgetreu.", + "image_b64": "...", + "upscale": 1, + "patch_size": 512, + "stride": 256 +} +``` + +Hermes zeigt dafür das zusätzliche Chatmodell `restauration` neben +`qwen-fast`, `qwen-medium`, `qwen-large` und `qwen-ultra`. Es verwendet +`qwen-fast` als kurzes Anweisungsmodell und bindet das Bildwerkzeug fest an +HYPIR. Die Formulierung des Prompts entscheidet nicht über das Backend. +In jedem normalen Chatmodell bleiben Referenzbildänderungen bei FLUX. + +## Lebenszyklus + +1. Der Router merkt sich das aktive Qwen-Profil und wartet laufende Chats ab. +2. Profile Controller stoppt Qwen und Qwen3-TTS. +3. Der disposable `restoration-worker` lädt HYPIR auf der RTX 5080. +4. Nach der Ausgabe wird der Worker beendet und sein CUDA-Kontext freigegeben. +5. Qwen3-TTS und das zuvor aktive Qwen-Profil werden wiederhergestellt. + +FLUX- und HYPIR-Worker können durch dieselbe Docker-Label-Allowlist nie +gleichzeitig mit einem Qwen-Profil laufen. Das Originalbild wird nicht +überschrieben; Eingaben werden nur als temporäre Dateien im privaten +`router-images`-Volume abgelegt und nach dem Auftrag entfernt. + +## Modelle und Lizenz + +- Code: `XPixelGroup/HYPIR`, Commit + `b61d107c6cef38f01a93c7833558869731cfa8c1` +- Restaurationsgewicht: `lxq007/HYPIR/HYPIR_sd2.pth` +- Basis: `LanguageMachines/stable-diffusion-2-1-base`, nur die benötigten + Diffusers-Komponenten +- HYPIR ist ausschließlich für nichtkommerzielle Nutzung freigegeben. + +## Test + +Für die Produktionsprobe in Hermes zuerst `restauration` auswählen, dann ein +Foto hochladen und beispielsweise `Mach das bitte schöner und schärfer` +schreiben. Danach sind zu prüfen: + +```bash +docker ps -a --filter name=mike-ai-restoration-worker \ + --format '{{.Names}} {{.Status}}' +nvidia-smi +curl -fsS http://127.0.0.1:8081/status +``` + +Erwartet: Restaurations-Worker beendet, vorheriges Qwen-Profil und TTS gesund, +Routerphase `idle`. diff --git a/integrations/hermes-athena-image/README.md b/integrations/hermes-athena-image/README.md index aed406b..4bf273f 100644 --- a/integrations/hermes-athena-image/README.md +++ b/integrations/hermes-athena-image/README.md @@ -4,6 +4,12 @@ Hermes backend plugin for the OpenAI-compatible image API exposed by the Athena profile router. The router starts the local FLUX worker on demand, unloads the active LLM and Qwen3-TTS, and restores both after generation. +When the explicit Hermes model `restauration` is selected, reference-image +requests are routed to Athena's `HYPIR-SD2` worker. Creative edits in every +normal chat model continue to use FLUX. Prompt text is deliberately never used +for routing, so words such as "improve" cannot accidentally switch pipelines. +No second Hermes provider or desktop installation is required. + ## Gateway installation Install this directory on the Hermes gateway, not on each Desktop client: @@ -40,6 +46,7 @@ second computer connected to the same gateway needs no plugin installation. Optional overrides: - `ATHENA_IMAGE_MODEL` defaults to `FLUX.2-klein-9B-fp8-beta`. +- `ATHENA_RESTORATION_MODEL` defaults to `HYPIR-SD2`. - `ROUTER_API_KEY` is accepted as a migration fallback. - An existing `HERMES_CUSTOM_192_168_1_212_8081_API_KEY` is accepted as the final fallback, so an existing Athena chat-provider setup needs no duplicate diff --git a/integrations/hermes-athena-image/__init__.py b/integrations/hermes-athena-image/__init__.py index 54d9415..d84b61b 100644 --- a/integrations/hermes-athena-image/__init__.py +++ b/integrations/hermes-athena-image/__init__.py @@ -34,6 +34,7 @@ _SIZES = { } _DEFAULT_BASE_URL = "http://192.168.1.212:8081/v1" _DEFAULT_MODEL = "FLUX.2-klein-9B-fp8-beta" +_DEFAULT_RESTORATION_MODEL = "HYPIR-SD2" _MAX_IMAGE_BYTES = 20 * 1024 * 1024 @@ -47,6 +48,34 @@ def _model() -> str: return os.environ.get("ATHENA_IMAGE_MODEL", "").strip() or _DEFAULT_MODEL +def _restoration_model() -> str: + return (os.environ.get("ATHENA_RESTORATION_MODEL", "").strip() + or _DEFAULT_RESTORATION_MODEL) + + +def _restoration_profile_selected() -> bool: + """Use HYPIR only when Hermes explicitly selected the restoration model.""" + try: + from hermes_cli.config import load_config + config = load_config() + model = config.get("model") if isinstance(config, dict) else None + selected = model.get("default") if isinstance(model, dict) else None + return isinstance(selected, str) and selected.casefold() in { + "restauration", "restoration", "qwen-restoration", + } + except Exception: + return False + + +def _select_model(requested: object) -> str: + """Resolve an explicit image/profile selection; never inspect prompt text.""" + if _restoration_profile_selected(): + return _restoration_model() + if isinstance(requested, str) and requested.strip() == _restoration_model(): + return _restoration_model() + return _model() + + def _api_key() -> str: """Prefer a scoped key; accept the existing router key for migration.""" return ( @@ -116,6 +145,12 @@ class AthenaLocalImageProvider(ImageGenProvider): "speed": "local", "strengths": "Private local generation and multi-reference editing", "price": "local / no cloud", + }, { + "id": _restoration_model(), + "display": "HYPIR-SD2 Restoration on Athena", + "speed": "local", + "strengths": "Faithful denoise, deblur and photo restoration", + "price": "local / no cloud", }] def default_model(self) -> Optional[str]: @@ -178,6 +213,7 @@ class AthenaLocalImageProvider(ImageGenProvider): sources.append(image_url.strip()) sources.extend(normalize_reference_images(reference_image_urls) or []) sources = sources[:4] + model = _select_model(kwargs.get("model")) try: encoded_sources = [ base64.b64encode(_load_private_image(source)).decode("ascii") @@ -204,6 +240,9 @@ class AthenaLocalImageProvider(ImageGenProvider): endpoint = "edits" request_data["image_b64"] = encoded_sources[0] request_data["reference_images_b64"] = encoded_sources[1:] + if model == _restoration_model(): + request_data.update({"upscale": 1, "patch_size": 512, + "stride": 256}) request = urllib.request.Request( f"{base_url}/images/{endpoint}", data=json.dumps(request_data).encode("utf-8"), method="POST", @@ -241,7 +280,9 @@ class AthenaLocalImageProvider(ImageGenProvider): prompt=clean_prompt, aspect_ratio=aspect) try: - saved = save_b64_image(b64_data, prefix="athena_flux2") + prefix = ("athena_hypir" if model == _restoration_model() + else "athena_flux2") + saved = save_b64_image(b64_data, prefix=prefix) except Exception as exc: return error_response( error=f"Generated image could not be saved: {exc}", diff --git a/platform/docker/profile-controller/profile_controller.py b/platform/docker/profile-controller/profile_controller.py index 7fa8e46..057d571 100644 --- a/platform/docker/profile-controller/profile_controller.py +++ b/platform/docker/profile-controller/profile_controller.py @@ -25,6 +25,7 @@ ALLOWED = tuple(x.strip() for x in os.environ.get( LABEL_KEY = "com.mike-ai.llama-profile" IMAGE_LABEL_KEY = "com.mike-ai.image-worker" IMAGE_WORKER = os.environ.get("IMAGE_WORKER", "image") +RESTORE_WORKER = os.environ.get("RESTORE_WORKER", "restore") TTS_LABEL_KEY = "com.mike-ai.tts-worker" TTS_WORKER = os.environ.get("TTS_WORKER", "qwen3") LOCK = threading.Lock() @@ -70,15 +71,22 @@ def labelled_containers(label: str) -> list[dict]: return json.loads(body) -def image_container() -> dict: +def image_container(kind: str = IMAGE_WORKER) -> dict: matches = [item for item in labelled_containers(IMAGE_LABEL_KEY) - if item.get("Labels", {}).get(IMAGE_LABEL_KEY) == IMAGE_WORKER] + if item.get("Labels", {}).get(IMAGE_LABEL_KEY) == kind] if len(matches) != 1: raise RuntimeError( - f"expected exactly one image worker {IMAGE_WORKER!r}, found {len(matches)}") + f"expected exactly one image worker {kind!r}, found {len(matches)}") return matches[0] +def image_containers() -> list[dict]: + """All allowlisted GPU workers that must never overlap an LLM.""" + allowed = {IMAGE_WORKER, RESTORE_WORKER} + return [item for item in labelled_containers(IMAGE_LABEL_KEY) + if item.get("Labels", {}).get(IMAGE_LABEL_KEY) in allowed] + + def tts_container() -> dict: matches = [item for item in labelled_containers(TTS_LABEL_KEY) if item.get("Labels", {}).get(TTS_LABEL_KEY) == TTS_WORKER] @@ -113,9 +121,11 @@ def stop_inference() -> dict: return {"active_profile": None, "previous_profile": previous} -def set_image_worker(running: bool) -> dict: +def set_image_worker(running: bool, kind: str = IMAGE_WORKER) -> dict: + if kind not in {IMAGE_WORKER, RESTORE_WORKER}: + raise ValueError("worker is not allowlisted") with LOCK: - item = image_container() + item = image_container(kind) if running: # The image worker may never overlap a llama profile on the 5080. for profile_item in containers().values(): @@ -123,6 +133,9 @@ def set_image_worker(running: bool) -> dict: # The 9B beta text encoder temporarily borrows the RTX 3060 from # Qwen3-TTS. The gateway retains Piper as a fallback meanwhile. stop_container(tts_container(), timeout=30) + for other in image_containers(): + if other["Id"] != item["Id"]: + stop_container(other, timeout=20) start_container(item) else: # CUDA/PyTorch may not react promptly to SIGTERM after an OOM. @@ -131,7 +144,8 @@ def set_image_worker(running: bool) -> dict: # TTS is restored by the following profile activation. Keeping it # stopped here lets the router verify that both GPUs really # released the image model before Qwen and TTS are reloaded. - return {"image_worker": "running" if running else "stopped"} + return {"image_worker": kind, + "state": "running" if running else "stopped"} def active_profile(items: dict[str, dict] | None = None) -> str | None: @@ -147,7 +161,8 @@ def activate(profile: str) -> dict: raise ValueError("profile is not allowlisted") with LOCK: # Defensive mutual exclusion even if a caller bypasses the router. - stop_container(image_container()) + for worker in image_containers(): + stop_container(worker) start_container(tts_container()) items = containers() missing = [name for name in ALLOWED if name not in items] @@ -230,9 +245,16 @@ class Handler(BaseHTTPRequestHandler): log.exception("stopping inference failed") self.reply(503, {"error": str(exc)}) return - if self.path in ("/workers/image/start", "/workers/image/stop"): + worker_paths = { + "/workers/image/start": (IMAGE_WORKER, True), + "/workers/image/stop": (IMAGE_WORKER, False), + "/workers/restore/start": (RESTORE_WORKER, True), + "/workers/restore/stop": (RESTORE_WORKER, False), + } + if self.path in worker_paths: try: - self.reply(200, set_image_worker(self.path.endswith("/start"))) + kind, running = worker_paths[self.path] + self.reply(200, set_image_worker(running, kind)) except Exception as exc: log.exception("image worker transition failed") self.reply(503, {"error": str(exc)}) diff --git a/platform/docker/restoration-worker/Dockerfile b/platform/docker/restoration-worker/Dockerfile new file mode 100644 index 0000000..87693a3 --- /dev/null +++ b/platform/docker/restoration-worker/Dockerfile @@ -0,0 +1,20 @@ +FROM pytorch/pytorch:2.11.0-cuda12.8-cudnn9-runtime + +ARG HYPIR_COMMIT=b61d107c6cef38f01a93c7833558869731cfa8c1 + +RUN apt-get update && apt-get install -y --no-install-recommends git ca-certificates python3.12-venv && \ + git clone https://github.com/XPixelGroup/HYPIR.git /opt/HYPIR && \ + cd /opt/HYPIR && git checkout "${HYPIR_COMMIT}" && \ + rm -rf /opt/HYPIR/.git /var/lib/apt/lists/* && \ + python -m venv --system-site-packages /opt/restore-venv && \ + /opt/restore-venv/bin/pip install --no-cache-dir \ + accelerate==1.4.0 diffusers==0.32.2 peft==0.14.0 \ + transformers==4.49.0 einops==0.8.1 omegaconf==2.3.0 \ + opencv-python-headless==4.11.0.86 safetensors pillow && \ + useradd --system --uid 10002 --home /nonexistent --shell /usr/sbin/nologin restoration-worker + +COPY restoration_worker.py /app/restoration_worker.py +USER 10002:10002 +WORKDIR /opt/HYPIR +ENV HF_HOME=/tmp/huggingface +ENTRYPOINT ["/opt/restore-venv/bin/python", "/app/restoration_worker.py"] diff --git a/platform/docker/restoration-worker/restoration_worker.py b/platform/docker/restoration-worker/restoration_worker.py new file mode 100644 index 0000000..bb83f46 --- /dev/null +++ b/platform/docker/restoration-worker/restoration_worker.py @@ -0,0 +1,186 @@ +#!/usr/bin/env python3 +"""Private HYPIR-SD2 still-image restoration worker for Athena. + +The container is intentionally disposable. The profile controller starts it +only for a restoration request and stops it before restoring the previous LLM +profile, which guarantees that CUDA allocations cannot leak into text mode. +""" + +from __future__ import annotations + +import gc +import json +import os +import signal +import sys +import time +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from pathlib import Path + +HOST = os.environ.get("WORKER_HOST", "0.0.0.0") +PORT = int(os.environ.get("WORKER_PORT", "8087")) +TOKEN = os.environ.get("WORKER_TOKEN", "").strip() +BASE_MODEL = os.environ.get("HYPIR_BASE_MODEL", "/models/sd2") +WEIGHT_FILE = os.environ.get("HYPIR_WEIGHT_FILE", "/models/hypir/HYPIR_sd2.pth") +OUTPUT_DIR = Path(os.environ.get("IMAGE_DIR", "/data/images")).resolve() +DEVICE = os.environ.get("HYPIR_DEVICE", "cuda:0") +ACTIVE = False +MODEL = None +LOAD_SECONDS = 0.0 + +# The entrypoint lives in /app, so Python would otherwise omit the cloned +# upstream repository from sys.path even though WORKDIR is /opt/HYPIR. +sys.path.insert(0, "/opt/HYPIR") + +LORA_MODULES = [ + "to_k", "to_q", "to_v", "to_out.0", "conv", "conv1", "conv2", + "conv_shortcut", "conv_out", "proj_in", "proj_out", "ff.net.2", + "ff.net.0.proj", +] + +if len(TOKEN) < 32: + raise RuntimeError("WORKER_TOKEN is missing or too short") + +signal.signal(signal.SIGTERM, lambda *_: os._exit(0)) + + +def _load_model(): + global MODEL, LOAD_SECONDS + if MODEL is not None: + return MODEL + from HYPIR.enhancer.sd2 import SD2Enhancer + + started = time.monotonic() + MODEL = SD2Enhancer( + base_model_path=BASE_MODEL, + weight_path=WEIGHT_FILE, + lora_modules=LORA_MODULES, + lora_rank=256, + model_t=200, + coeff_t=200, + device=DEVICE, + ) + MODEL.init_models() + LOAD_SECONDS = time.monotonic() - started + return MODEL + + +def _safe_source(name: object) -> Path: + if not isinstance(name, str) or Path(name).name != name: + raise ValueError("invalid source image filename") + source = (OUTPUT_DIR / name).resolve() + if source.parent != OUTPUT_DIR or not source.is_file(): + raise ValueError("source image not found") + return source + + +def _prompt(value: object) -> str: + text = value.strip() if isinstance(value, str) else "" + operation_words = { + "restauriere", "restaurieren", "restore", "verbessere", "verbessern", + "enhance", "repariere", "reparieren", "dieses", "das", "bild", "foto", + "photo", "image", "bitte", "schärfer", "schaerfer", "aufarbeiten", + } + words = {part.strip(".,:;!?-_").lower() for part in text.split()} + if not text or words.issubset(operation_words): + return "high quality natural photograph, faithful details, realistic textures" + return text[:1000] + + +def restore(data: dict) -> dict: + global ACTIVE + import torch + import numpy as np + from PIL import Image + + filename = data.get("filename") + if (not isinstance(filename, str) or Path(filename).name != filename + or not filename.endswith(".png")): + raise ValueError("invalid filename") + source_files = data.get("source_files") or [] + if not isinstance(source_files, list) or len(source_files) != 1: + raise ValueError("HYPIR restoration requires exactly one source image") + source = _safe_source(source_files[0]) + upscale = int(data.get("upscale", 1)) + if upscale not in (1, 2, 4): + raise ValueError("upscale must be 1, 2, or 4") + patch_size = int(data.get("patch_size", 512)) + stride = int(data.get("stride", 256)) + if patch_size not in (512, 768, 1024) or stride <= 0 or stride > patch_size: + raise ValueError("invalid patch_size/stride") + + ACTIVE = True + started = time.monotonic() + try: + model = _load_model() + with Image.open(source) as opened: + image = opened.convert("RGB") + array = np.asarray(image, dtype=np.float32) / 255.0 + tensor = torch.from_numpy(array).permute(2, 0, 1).unsqueeze(0) + result = model.enhance( + lq=tensor, + prompt=_prompt(data.get("prompt")), + scale_by="factor", + upscale=upscale, + patch_size=patch_size, + stride=stride, + return_type="pil", + )[0] + OUTPUT_DIR.mkdir(parents=True, exist_ok=True) + result.save(OUTPUT_DIR / filename) + return { + "status": "ok", + "filename": filename, + "seconds": round(time.monotonic() - started, 3), + "load_seconds": round(LOAD_SECONDS, 3), + "model": "HYPIR-SD2", + "upscale": upscale, + } + finally: + gc.collect() + torch.cuda.empty_cache() + ACTIVE = False + + +class Handler(BaseHTTPRequestHandler): + def log_message(self, fmt: str, *args: object) -> None: + print(f"[hypir-restore] {self.client_address[0]} {fmt % args}", flush=True) + + def reply(self, status: int, payload: dict) -> None: + body = json.dumps(payload, separators=(",", ":")).encode() + self.send_response(status) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def do_GET(self) -> None: # noqa: N802 + if self.path == "/health": + self.reply(200, {"status": "ok", "model_loaded": MODEL is not None, + "active": ACTIVE, "model": "HYPIR-SD2"}) + else: + self.reply(404, {"error": "not found"}) + + def do_POST(self) -> None: # noqa: N802 + if self.headers.get("Authorization", "") != f"Bearer {TOKEN}": + self.reply(401, {"error": "unauthorized"}) + return + if self.path != "/restore": + self.reply(404, {"error": "not found"}) + return + try: + length = int(self.headers.get("Content-Length", "0")) + if length < 2 or length > 16384: + raise ValueError("invalid request size") + self.reply(200, restore(json.loads(self.rfile.read(length)))) + except Exception as exc: + print(f"[hypir-restore] restoration failed: " + f"{type(exc).__name__}: {str(exc)[:1000]}", flush=True) + self.reply(400, {"status": "error", "message": str(exc)}) + + +try: + ThreadingHTTPServer((HOST, PORT), Handler).serve_forever() +finally: + MODEL = None + gc.collect() diff --git a/router/ai_profile_router.py b/router/ai_profile_router.py index 2593a11..095b8e1 100755 --- a/router/ai_profile_router.py +++ b/router/ai_profile_router.py @@ -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")