Remove ineffective photo restoration pipeline
This commit is contained in:
1 parent
f82dc081c9
commit
e82e0340e4
9 files changed
+17
-552
No files matched your search
@@ -1,20 +0,0 @@
|
||||
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"]
|
||||
@@ -1,199 +0,0 @@
|
||||
#!/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")
|
||||
PROMPT_MODE = os.environ.get("HYPIR_PROMPT_MODE", "empty").strip().lower()
|
||||
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:
|
||||
# HYPIR inherits text-guided texture synthesis from Stable Diffusion. For
|
||||
# the explicit restoration profile fidelity is more important than
|
||||
# creativity, so the production default disables text conditioning. A
|
||||
# future experimental profile can opt back into it without changing code.
|
||||
if PROMPT_MODE == "empty":
|
||||
return ""
|
||||
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)})
|
||||
|
||||
|
||||
def main() -> None:
|
||||
global MODEL
|
||||
try:
|
||||
ThreadingHTTPServer((HOST, PORT), Handler).serve_forever()
|
||||
finally:
|
||||
MODEL = None
|
||||
gc.collect()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in new issue
Block a user