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
@@ -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()