Add explicit HYPIR restoration profile
This commit is contained in:
@@ -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()
|
||||
Reference in New Issue
Block a user