200 lines
9 KiB
Python
200 lines
9 KiB
Python
"""Serve a resident prompt-only H3 runtime over a small JSON HTTP API."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import json
|
|
import threading
|
|
import time
|
|
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
|
from pathlib import Path
|
|
from urllib.parse import urlparse
|
|
|
|
import torch
|
|
from PIL import Image
|
|
|
|
from h3_blackwell_runtime.attention import AVAILABLE_BACKENDS, DEFAULT_ATTENTION_BACKEND, attention_backend_status
|
|
from h3_blackwell_runtime.runtime import H3HotRuntime, RuntimeConfig, TURBO_VARIANTS, normalize_upscale
|
|
|
|
|
|
def _load_image(value) -> torch.Tensor | None:
|
|
"""Accept a keyframe as on-disk path or base64 JPEG/PNG -> ``[1,3,H,W]`` float ``[0,1]``."""
|
|
if value in (None, ""):
|
|
return None
|
|
if isinstance(value, (list, tuple)):
|
|
value = value[0]
|
|
if isinstance(value, dict):
|
|
value = value.get("url") or value.get("path") or value.get("b64")
|
|
if isinstance(value, str) and "\n" not in value and len(value) < 2048 and not value.startswith("data:"):
|
|
path = Path(value)
|
|
if path.exists():
|
|
image = Image.open(path).convert("RGB")
|
|
import numpy as np
|
|
|
|
tensor = torch.from_numpy(np.array(image))[None].permute(0, 3, 1, 2).float() / 255.0
|
|
return tensor
|
|
data = value
|
|
if isinstance(value, str) and value.startswith("data:"):
|
|
data = value.split(",", 1)[1]
|
|
if isinstance(data, str) and len(data) >= 1024:
|
|
import base64
|
|
import io
|
|
|
|
raw = base64.b64decode(data)
|
|
image = Image.open(io.BytesIO(raw)).convert("RGB")
|
|
import numpy as np
|
|
|
|
return torch.from_numpy(np.array(image))[None].permute(0, 3, 1, 2).float() / 255.0
|
|
raise ValueError("first_frame/last_frame must be a path or a base64/data-URL image")
|
|
|
|
|
|
parser = argparse.ArgumentParser()
|
|
parser.add_argument("--host", default="0.0.0.0")
|
|
parser.add_argument("--port", type=int, default=8000)
|
|
parser.add_argument("--attention", choices=AVAILABLE_BACKENDS, default=DEFAULT_ATTENTION_BACKEND, help="Initial attention backend. Requests can switch with the JSON attention field.")
|
|
parser.add_argument("--vae-dtype", choices=("float32", "float16", "bfloat16"), default="float16")
|
|
parser.add_argument("--vae-tile-size", type=int, default=256)
|
|
parser.add_argument("--mlp-chunks", type=int, default=1)
|
|
parser.add_argument("--mlp-chunk-threshold", type=int, default=4096)
|
|
parser.add_argument("--turbo-lora", action="append", default=[], metavar="NAME=PATH", help="Load a resident 4step or 8step Turbo adapter.")
|
|
parser.add_argument("--latent-upscaler", help="Load the optional H3 3D latent upscaler for request-level spatial upscaling.")
|
|
parser.add_argument("--warmup", action="store_true", help="Run a tiny generation before accepting traffic.")
|
|
parser.add_argument("--warmup-output", type=Path, default=Path("/output/h3-blackwell-runtime/hot-runtime-warmup.mp4"))
|
|
args = parser.parse_args()
|
|
|
|
turbo_loras = []
|
|
for value in args.turbo_lora:
|
|
if "=" not in value:
|
|
parser.error("--turbo-lora must use NAME=PATH")
|
|
name, path = value.split("=", 1)
|
|
turbo_loras.append((name, path))
|
|
|
|
runtime = H3HotRuntime(RuntimeConfig(attention=args.attention, vae_dtype=args.vae_dtype, vae_tile_size=args.vae_tile_size, mlp_chunks=args.mlp_chunks, mlp_chunk_threshold=args.mlp_chunk_threshold, turbo_loras=tuple(turbo_loras), latent_upscaler_path=args.latent_upscaler))
|
|
runtime_lock = threading.Lock()
|
|
warmup_result = None
|
|
if args.warmup:
|
|
warmup_result = runtime.generate(
|
|
prompt="A small warmup cat blinks in soft light.",
|
|
output=args.warmup_output,
|
|
width=320,
|
|
height=192,
|
|
frames=22,
|
|
steps=2,
|
|
seed=440501,
|
|
mux_audio=True,
|
|
)
|
|
|
|
|
|
def service_status() -> dict:
|
|
return {
|
|
"ready": True,
|
|
"attention_backends": list(AVAILABLE_BACKENDS),
|
|
"attention_backend_status": attention_backend_status(),
|
|
"runtime": runtime.status(),
|
|
}
|
|
|
|
|
|
def write_json(handler: BaseHTTPRequestHandler, status: int, payload: dict) -> None:
|
|
body = json.dumps(payload, indent=2).encode("utf-8")
|
|
handler.send_response(status)
|
|
handler.send_header("Content-Type", "application/json")
|
|
handler.send_header("Content-Length", str(len(body)))
|
|
handler.end_headers()
|
|
handler.wfile.write(body)
|
|
|
|
|
|
class Handler(BaseHTTPRequestHandler):
|
|
def log_message(self, format: str, *args) -> None:
|
|
return
|
|
|
|
def do_GET(self) -> None:
|
|
path = urlparse(self.path).path
|
|
if path in {"/health", "/ready"}:
|
|
write_json(self, 200, {**service_status(), "warmup_result": warmup_result})
|
|
return
|
|
write_json(self, 404, {"error": "not found"})
|
|
|
|
def do_POST(self) -> None:
|
|
path = urlparse(self.path).path
|
|
if path != "/generate":
|
|
write_json(self, 404, {"error": "not found"})
|
|
return
|
|
try:
|
|
length = int(self.headers.get("Content-Length", "0"))
|
|
payload = json.loads(self.rfile.read(length).decode("utf-8")) if length else {}
|
|
prompt = payload["prompt"]
|
|
output = payload["output"]
|
|
width = int(payload.get("width", 960))
|
|
height = int(payload.get("height", 544))
|
|
frames = int(payload.get("frames", 124))
|
|
turbo = payload.get("turbo")
|
|
if turbo in {"", "none"}:
|
|
turbo = None
|
|
if turbo is not None and turbo not in runtime.turbo_loras:
|
|
write_json(self, 400, {"error": "unsupported turbo", "turbo": turbo, "available": sorted(runtime.turbo_loras)})
|
|
return
|
|
steps = int(payload.get("steps", TURBO_VARIANTS[turbo]["steps"] if turbo else 12))
|
|
if turbo is not None and steps != TURBO_VARIANTS[turbo]["steps"]:
|
|
write_json(self, 400, {"error": "invalid turbo steps", "turbo": turbo, "required_steps": TURBO_VARIANTS[turbo]["steps"]})
|
|
return
|
|
seed = int(payload.get("seed", 440407))
|
|
attention = payload.get("attention")
|
|
if attention is not None and attention not in AVAILABLE_BACKENDS:
|
|
write_json(self, 400, {"error": "unsupported attention", "attention": attention, "available": list(AVAILABLE_BACKENDS)})
|
|
return
|
|
mux_audio = bool(payload.get("mux_audio", True))
|
|
keep_intermediates = bool(payload.get("keep_intermediates", False))
|
|
try:
|
|
upscale = normalize_upscale(payload.get("upscale"))
|
|
except ValueError as exc:
|
|
write_json(self, 400, {"error": "invalid upscale", "message": str(exc)})
|
|
return
|
|
if upscale is not None and runtime.latent_upscaler is None:
|
|
write_json(self, 400, {"error": "H3 latent upscaler is not loaded"})
|
|
return
|
|
ffmpeg_loglevel = payload.get("ffmpeg_loglevel", "error")
|
|
first_frame = _load_image(payload.get("first_frame"))
|
|
last_frame = _load_image(payload.get("last_frame"))
|
|
save_latent = payload.get("save_latent")
|
|
cache_mode = payload.get("cache_mode")
|
|
if turbo is not None and cache_mode not in {None, "", "disabled", "none"}:
|
|
write_json(self, 400, {"error": "turbo does not support denoiser caching", "turbo": turbo})
|
|
return
|
|
cache_threshold = float(payload.get("cache_threshold", 0.0))
|
|
cache_start_percent = float(payload.get("cache_start_percent", 0.0))
|
|
cache_end_percent = float(payload.get("cache_end_percent", 1.0))
|
|
cache_subsample_factor = int(payload.get("cache_subsample_factor", 2))
|
|
started = time.perf_counter()
|
|
with runtime_lock:
|
|
result = runtime.generate(
|
|
prompt=prompt,
|
|
output=output,
|
|
width=width,
|
|
height=height,
|
|
frames=frames,
|
|
steps=steps,
|
|
seed=seed,
|
|
attention=attention,
|
|
turbo=turbo,
|
|
upscale=upscale,
|
|
first_frame=first_frame,
|
|
last_frame=last_frame,
|
|
mux_audio=mux_audio,
|
|
keep_intermediates=keep_intermediates,
|
|
ffmpeg_loglevel=ffmpeg_loglevel,
|
|
save_latent=save_latent,
|
|
cache_mode=cache_mode,
|
|
cache_threshold=cache_threshold,
|
|
cache_start_percent=cache_start_percent,
|
|
cache_end_percent=cache_end_percent,
|
|
cache_subsample_factor=cache_subsample_factor,
|
|
)
|
|
result["wall_seconds"] = time.perf_counter() - started
|
|
write_json(self, 200, result)
|
|
except Exception as exc:
|
|
write_json(self, 500, {"error": type(exc).__name__, "message": str(exc)})
|
|
|
|
|
|
server = ThreadingHTTPServer((args.host, args.port), Handler)
|
|
print(json.dumps({"serving": True, "host": args.host, "port": args.port, **service_status(), "warmup_result": warmup_result}, indent=2), flush=True)
|
|
server.serve_forever()
|