Add explicit HYPIR restoration profile

This commit is contained in:
Mikei386
2026-09-07 22:52:59 +02:00
parent 2ae61baec7
commit 118e32005e
11 changed files with 593 additions and 32 deletions
+41
View File
@@ -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
+3
View File
@@ -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.
+32 -4
View File
@@ -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")
+25
View File
@@ -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:
+71
View File
@@ -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`.
@@ -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
+42 -1
View File
@@ -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}",
@@ -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)})
@@ -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"]
@@ -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()
+135 -18
View File
@@ -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")