Meridian / service /app.py
yycc's picture
Two-adapter release: docs and code
1f487cc verified
Raw History Blame Contribute Delete
27.9 kB
# Copyright 2026 Viggle AI. Licensed under the Apache License, Version 2.0 (see LICENSE-CODE).
# SPDX-License-Identifier: Apache-2.0
"""The demo service: upload a clip -> geometry is reconstructed at once -> author a keyframe camera path in
3D -> truthful warp preview -> render.
One process, one GPU, everything resident (VGGT-Omega + VAE + teacher + the DMD student LoRA, ~77 GiB at boot).
A warp of the whole trajectory takes 0.24 s, so every judgement the visitor makes is made on the real thing
the model will be conditioned on -- downsampled to `cond_canvas`, the 480 class, not the output canvas.
CUDA_VISIBLE_DEVICES=0 python service/app.py --port 8412 # or service/run.sh
See README.md, "Self-hosting the demo", for the endpoints and the measured budget.
"""
import argparse
import hashlib
import io
import itertools
import json
import math
import os
import subprocess
import sys
import threading
import time
import av
import numpy as np
import torch
import uvicorn
from diffusers import AutoencoderKLMiniMaxH3, MiniMaxH3Transformer3DModel
from diffusers.utils.export_utils import encode_video as write_mp4
from fastapi import FastAPI, File, HTTPException, Request, UploadFile
from fastapi.responses import FileResponse, HTMLResponse, JSONResponse, Response
from PIL import Image
HERE = os.path.dirname(os.path.abspath(__file__))
ROOT = os.path.dirname(HERE)
sys.path.insert(0, ROOT)
from recam.geometry import FULL, RES, reconstruct, resize_u8, to_input, vggt, warp # noqa: E402
from recam.h3 import FPS, bucket, decode_video, denoise, encode_video, pack # noqa: E402
from recam.path import describe, look_at, plan_path # noqa: E402
p = argparse.ArgumentParser()
p.add_argument("--ckpt", default=None, help="a local diffusers `transformer/` dir; default is `transformer/` inside --model-dir (stock MiniMax-H3)")
p.add_argument("--lora", nargs="+", default=[f"{ROOT}/teacher_lora", f"{ROOT}/turbo_lora"], help="adapter dirs applied together: the recam teacher, then the DMD turbo")
p.add_argument("--model-dir", default="MiniMaxAI/MiniMax-H3", help="the base MiniMax-H3 repo or a local copy, for `vae/`")
p.add_argument("--vggt", default=None, help="vggt_omega_1b_512.pt; default $VGGT_OMEGA_CKPT (see README)")
p.add_argument("--vggt-repo", default=None, help="a checkout of facebookresearch/vggt-omega; default $VGGT_OMEGA_DIR")
p.add_argument("--steps", type=int, default=4, help="the student's grid: 4 timesteps = 3 forwards")
p.add_argument("--flow-shift", type=float, default=3.0)
p.add_argument("--work", default=f"{ROOT}/work", help="uploads, warps and takes land here")
p.add_argument("--samples", default=f"{ROOT}/examples/media", help="mp4s offered on the first screen")
p.add_argument("--host", default="0.0.0.0")
p.add_argument("--port", type=int, default=8412)
p.add_argument("--max-clips", type=int, default=8)
p.add_argument("--max-prep", type=int, default=8)
args = p.parse_args()
ASSETS = f"{ROOT}/assets"
LIVE_SPEED = 0.33
dev = torch.device("cuda")
torch.set_grad_enabled(False) # main thread only -- grad mode is thread-local, and uvicorn runs sync
# endpoints in a threadpool while /render runs in its own thread, so every GPU entry point below
# carries its own @torch.no_grad(). Without it the VAE encode retains activations: 78 -> 177 GiB.
os.makedirs(args.work, exist_ok=True)
CLIPS, PREP, JOBS = {}, {}, {}
GPU = threading.Lock() # one card: a render and a warp cannot overlap
app = FastAPI()
print("loading VGGT ...", flush=True)
geometry = vggt(args.vggt, args.vggt_repo, dev)
print("loading VAE ...", flush=True)
vae = AutoencoderKLMiniMaxH3.from_pretrained(args.model_dir, subfolder="vae").to(dev).eval()
print(f"loading {args.ckpt or args.model_dir} + {[os.path.basename(d) for d in args.lora]} ...", flush=True)
ckpt, sfold = (args.ckpt, None) if args.ckpt else (args.model_dir, "transformer")
transformer = MiniMaxH3Transformer3DModel.from_pretrained(ckpt, subfolder=sfold, torch_dtype=torch.bfloat16)
names = [f"lora{i}" for i in range(len(args.lora))]
for name, ldir in zip(names, args.lora): # set_adapters, not the load alone, is what sums them at weight 1.0
transformer.load_lora_adapter(ldir, weight_name="pytorch_lora_weights.safetensors", prefix=None, adapter_name=name)
transformer.set_adapters(names, [1.0] * len(names))
assert any("lora_" in n for n, _ in transformer.named_parameters()), "no LoRA weights landed"
transformer.set_attention_backend("_native_cudnn")
transformer.to(dev).eval()
print(f"ready, {torch.cuda.memory_allocated() / 2**30:.1f} GiB resident", flush=True)
# --- clip ingest ------------------------------------------------------------------------------------
def normalise(src, dst):
"""H.264 / CFR 24 / yuv420p / long edge <= 1280, rotation baked in. `:129` sends frames to the GPU at
their original resolution, so an un-normalised 4K upload is 1.8 GB on-card for 73 frames."""
subprocess.run(["ffmpeg", "-v", "error", "-y", "-autorotate", "1", "-i", src, "-vf",
"scale=1280:1280:force_original_aspect_ratio=decrease:force_divisible_by=2",
"-r", "24", "-c:v", "libx264", "-crf", "18", "-pix_fmt", "yuv420p",
"-c:a", "aac", dst], check=True)
def cuts_of(fr):
"""Hard-cut frames. VGGT reconstructs the span jointly, so a cut inside it poisons every camera
downstream -- this is a gate, not a warning. Cheap: 64x64 greyscale frame difference."""
x = torch.from_numpy(fr[:, ::max(1, fr.shape[1] // 64), ::max(1, fr.shape[2] // 64)]).float().mean(-1)
d = (x[1:] - x[:-1]).abs().mean(dim=(1, 2))
return [int(i) + 1 for i in torch.nonzero(d > torch.maximum(3 * d.median(), torch.tensor(14.0)))[:, 0]]
def ingest(raw, name):
sha = hashlib.sha256(raw).hexdigest()[:16]
d = f"{args.work}/clips/{sha}"
os.makedirs(d, exist_ok=True)
mp4 = f"{d}/clip.mp4"
if sha not in CLIPS:
if not os.path.exists(mp4):
open(f"{d}/raw", "wb").write(raw)
normalise(f"{d}/raw", mp4)
os.remove(f"{d}/raw")
c = av.open(mp4)
fr = np.stack([f.to_ndarray(format="rgb24") for f in c.decode(video=0)])
c.close()
CLIPS[sha] = dict(path=mp4, fr=fr, n=len(fr), h=fr.shape[1], w=fr.shape[2], name=name, cuts=cuts_of(fr))
for k in list(CLIPS)[: max(0, len(CLIPS) - args.max_clips)]:
CLIPS.pop(k)
c = CLIPS[sha]
return dict(clip=sha, frames=c["n"], w=c["w"], h=c["h"], name=name, cuts=c["cuts"],
seconds=round(c["n"] / FPS, 2), lengths=[73, 124, 175, 243] if c["n"] >= 73 else []) # the take's length is free of the source's: a path can hold or slow the clip
# --- geometry ---------------------------------------------------------------------------------------
@torch.no_grad()
def prepare(clip, start, span_end):
"""Letterbox + one VGGT pass over source frames `start..span_end`, cached.
The key is the DECODED span, not `(clip, start, frames)`: `tmap[-1] = start + frames - n`, so the
freeze length changes which frames VGGT sees, and VGGT is a joint pass -- 50 frames and 73 frames do
not give the same tensors for the frames they share. Moving the freeze *frame* at a fixed `n` keeps
the span, so it is free; changing `n` costs one 2.0 s pass."""
key = (clip, start, span_end)
if key in PREP:
return PREP[key]
c = CLIPS[clip]
assert 0 <= start < span_end < c["n"], f"the window runs to source frame {span_end} but the clip has only {c['n']} frames"
t0 = time.time()
frames = torch.from_numpy(c["fr"][start : span_end + 1]).to(dev)
h, w = frames.shape[1:3]
s = FULL / max(h, w)
ch, cw = round(h * s), round(w * s)
ox, oy = (FULL - cw) // 2, (FULL - ch) // 2
full = torch.zeros(len(frames), FULL, FULL, 3, dtype=torch.uint8, device=dev)
full[:, oy : oy + ch, ox : ox + cw] = resize_u8(frames, (ch, cw))
del frames
canvas, cond_canvas = bucket(w, h)
a = canvas[0] / canvas[1]
bw, bh = (cw, round(cw / a)) if cw / ch <= a else (round(ch * a), ch)
box = (ox + (cw - bw) // 2, oy + (ch - bh) // 2, bw, bh, canvas[0] / bw)
S0 = reconstruct(geometry, to_input(full))
picture = torch.zeros(RES, RES, dtype=torch.bool, device=dev)
picture[round(oy * RES / FULL) : round((oy + ch) * RES / FULL), round(ox * RES / FULL) : round((ox + cw) * RES / FULL)] = True
P = dict(S0=S0, full0=full, box=box, canvas=canvas, cond_canvas=cond_canvas, start=start, picture=picture,
scale=s, ox=ox, oy=oy, ch=ch, cw=cw, ms=int(1000 * (time.time() - t0)))
PREP[key] = P
for k in list(PREP)[: max(0, len(PREP) - args.max_prep)]:
PREP.pop(k)
return P
def plan(q):
"""The payload -> (tmap, ramp, pf). One enum, not flags: `--ease` on a flat ramp is a static camera."""
start, n_out, shape = int(q["start"]), int(q["frames"]), q["shape"]
if shape.startswith("freeze"):
fz, n = int(q["freeze_frame"]), int(q["freeze_n"])
tail = n_out - n - (fz - start)
assert fz >= start and tail >= 0, f"freeze {fz}x{n} does not fit a {n_out}-frame window from {start}"
tmap = list(range(start, fz)) + [fz] * n + list(range(fz + 1, fz + 1 + tail))
ramp = torch.cat([torch.zeros(fz - start), torch.linspace(0, 1, n), torch.ones(tail)])
if shape == "freeze_sweep":
wt = torch.tensor([LIVE_SPEED] * (fz - start) + [1.0] * n + [LIVE_SPEED] * tail)
ramp = torch.cumsum(wt, 0) - wt[0]
ramp = ramp / ramp[-1]
pf = tmap.index(fz)
else:
tmap = list(range(start, start + n_out))
ramp = torch.ones(n_out) if shape == "hold" else torch.linspace(0, 1, n_out)
if shape == "sweep_ease":
ramp = (1 - torch.cos(math.pi * ramp)) / 2
if shape == "bounce":
ramp = (1 - torch.cos(2 * math.pi * ramp)) / 2
if shape == "swing":
ramp = torch.sin(2 * math.pi * ramp)
pf = 0
return tmap, ramp, pf
def pivot_frac(P, u, v):
"""A click at (u,v) in source-frame fractions -> `--pivot fx,fy`, which is a fraction of the CROP BOX
inside the 1280 letterbox, not of the source frame. Doing this conversion in the browser is what killed
110 render-farm jobs, so it lives here. The clamp is the silent-NaN edge: an out-of-range window gives
an empty median, i.e. `zm = nan`, and 25 s of grey mud with no error."""
x0, y0, bw, bh, _ = P["box"]
fx = (P["ox"] + u * P["cw"] - x0) / bw
fy = (P["oy"] + v * P["ch"] - y0) / bh
return min(max(fx, 0.06), 0.94), min(max(fy, 0.06), 0.94)
def unit(P, q):
"""The path's frame and unit. `zm` = median depth in a 10% window around the pivot (default: the centre of
the picture) at `pivot_frame` (default: `start`), falling back to the whole picture when the window has no
depth (sky); `piv` = that point in the pivot frame's camera; `W` = the camera of frame `start`, the frame
every key is expressed in. Returns zi, zm, piv, W, piv_world, (fx, fy)."""
Z, start = P["S0"], P["start"]
zi = int(q.get("pivot_frame", start)) - start
assert 0 <= zi < len(P["full0"]), f"pivot frame {zi + start} outside the prepared span"
x0, y0, bw, bh, _ = P["box"]
fx, fy = pivot_frac(P, *[float(v) for v in q.get("pivot", [0.5, 0.5])])
r = RES / FULL
px, py, rw, rh = (x0 + fx * bw) * r, (y0 + fy * bh) * r, 0.05 * bw * r, 0.05 * bh * r
win = (slice(round(py - rh), round(py + rh)), slice(round(px - rw), round(px + rw)))
sub = Z["depth"][zi][win][Z["keep"][zi][win]]
if sub.numel() < 20:
sub = Z["depth"][zi][Z["keep"][zi] & P["picture"]]
zm = float(sub.median()) if sub.numel() else float("nan")
assert math.isfinite(zm), "no usable depth in this frame"
K = Z["intr"][zi]
piv = torch.tensor([(px - float(K[0, 2])) / float(K[0, 0]) * zm,
(py - float(K[1, 2])) / float(K[1, 1]) * zm, zm], device=dev)
W = Z["c2w"][0]
piv_world = Z["c2w"][zi][:3, :3] @ piv + Z["c2w"][zi][:3, 3]
return zi, zm, piv, W, piv_world, (fx, fy)
def poses(P, q):
"""Every prepared source frame's camera as a key (pos, look, roll) in the path's frame, plus the pivot."""
_, zm, _, W, piv_world, _ = unit(P, q)
Wi = torch.linalg.inv(W)
piv_P = (Wi[:3, :3] @ piv_world + Wi[:3, 3]).cpu().numpy()
out = describe((Wi[None] @ P["S0"]["c2w"]).cpu().numpy(), piv_P, zm)
x0, y0, bw, bh, _ = P["box"]
r = FULL / RES # each frame's lens as canvas fractions (fx/w, fy/h, cx/w, cy/h), exactly what /warp1 renders with, so the editor can look through a key
for p, K in zip(out, P["S0"]["intr"].cpu().numpy()):
p["k"] = [round(float(v), 5) for v in (K[0, 0] * r / bw, K[1, 1] * r / bh, (K[0, 2] * r - x0) / bw, (K[1, 2] * r - y0) / bh)]
return dict(src_poses=out, piv=(piv_P / zm).round(4).tolist(), zm=zm)
@torch.no_grad()
def geo(q):
"""Everything up to and including the warp. Returns the folded session, the cameras and the gauges."""
start = int(q["start"])
P = PREP[(q["clip"], start, int(q["span_end"]))]
Z = P["S0"]
if q.get("path"): # keyframe path (recam/path.py): the pivot only sets the unit `zm` and the gauges
n_out, pf_src = int(q["frames"]), int(q.get("pivot_frame", q.get("freeze_frame", q["path"][0]["src"])))
else: # the parametric family: yaw / dolly / truck / boom about the pivot, one ramp
tmap, ramp, pf = plan(q)
n_out, pf_src = len(tmap), tmap[pf]
zi, zm, piv, W, piv_world, (fx, fy) = unit(P, {**q, "pivot_frame": pf_src}) # piv: --pivot-lock, always
if q.get("path"):
c2w_P, tmap, focal, speed = plan_path(q["path"], n_out, zm)
pf = tmap.index(pf_src) if pf_src in tmap else 0
idx = torch.tensor(tmap) - start
assert 0 <= int(idx.min()) and int(idx.max()) < len(P["full0"]), f"tmap {tmap[0]}..{tmap[-1]} outside the prepared span"
idx = idx.to(dev) # an out-of-range gather is a device-side assert, which kills the CUDA context for good
S = {k: v[idx] for k, v in Z.items()}
full = P["full0"][idx]
c2w, intr_t = S["c2w"].clone(), S["intr"].clone()
if q.get("path"):
c2w = W[None] @ torch.tensor(c2w_P, dtype=c2w.dtype, device=dev)
f = torch.tensor(focal, dtype=intr_t.dtype, device=dev)
intr_t[:, 0, 0] *= f
intr_t[:, 1, 1] *= f
yaw, dolly, truck, boom, aim = (float(q.get("yaw", 0)), float(q.get("dolly", 1)), float(q.get("truck", 0)),
float(q.get("boom", 0)), bool(q.get("aim", False)))
for ti in ([] if q.get("path") else range(n_out)):
th = math.radians(yaw * float(ramp[ti]))
rr = 1 + (dolly - 1) * float(ramp[ti])
R = torch.tensor([[math.cos(th), 0, math.sin(th)], [0, 1, 0], [-math.sin(th), 0, math.cos(th)]], device=dev)
delta = torch.eye(4, device=dev)
delta[:3, :3] = R
delta[:3, 3] = piv - R @ (rr * piv) - R @ torch.tensor(
[-truck * zm * float(ramp[ti]), boom * zm * float(ramp[ti]), 0.0], device=dev)
if aim:
a = piv / piv.norm()
b = piv - delta[:3, 3]
b = b / b.norm()
v, c = torch.cross(a, b, dim=0), float(a @ b)
sn = float(v.norm())
Kx = torch.zeros(3, 3, device=dev)
Kx[0, 1], Kx[0, 2], Kx[1, 0], Kx[1, 2], Kx[2, 0], Kx[2, 1] = -v[2], v[1], v[2], -v[0], -v[1], v[0]
delta[:3, :3] = torch.eye(3, device=dev) + Kx + Kx @ Kx * ((1 - c) / sn ** 2) if sn > 1e-8 else torch.eye(3, device=dev)
c2w[ti] = S["c2w"][ti] @ delta
intr_t[ti, 0, 0] *= rr
intr_t[ti, 1, 1] *= rr
w2c = torch.linalg.inv(c2w)
render, cov = warp(S, w2c, intr_t, full, P["box"], P["canvas"])
yy, xx = torch.meshgrid(torch.arange(RES, device=dev), torch.arange(RES, device=dev), indexing="ij")
g1 = torch.stack([xx, yy, torch.ones_like(xx)], -1).float().reshape(-1, 3)
near, behind, coll, ahead = [], [], [], []
for ti in range(n_out):
X = ((g1 @ torch.linalg.inv(S["intr"][ti]).T).reshape(RES, RES, 3) * S["depth"][ti][..., None]) \
@ S["c2w"][ti][:3, :3].T + S["c2w"][ti][:3, 3]
zd = (X - c2w[ti][:3, 3]) @ c2w[ti][:3, 2]
near.append(float((zd[S["keep"][ti]] < 0.1 * zm).float().mean()))
behind.append(float((zd[S["keep"][ti]] < 0).float().mean()))
coll.append(float(((X - c2w[ti][:3, 3]).norm(dim=-1)[S["keep"][ti]] < 0.05 * zm).float().mean()))
cb = S["keep"][ti].clone()
cb[: RES // 4] = cb[-(RES // 4):] = False
cb[:, : RES // 4] = cb[:, -(RES // 4):] = False
ahead.append(float(torch.quantile(zd[cb], 0.05)) / zm if cb.any() else 9.0)
g = dict(zm=zm, ahead=min(ahead), near=max(near), behind=max(behind), coll=max(coll),
coverage=float(cov.float().mean()),
moved=float((c2w[:, :3, 3] - S["c2w"][:, :3, 3]).norm(dim=-1).max()) / zm, # largest departure from the source camera anywhere in the take (a freeze-orbit returns to it at the end)
pivot_frac=[fx, fy], frames=n_out, tmap=tmap)
# the path view of whatever was authored: per-frame (pos, look) in W / zm, so the studio can turn a
# parametric move (the templates) into keys
Wi = torch.linalg.inv(W)
piv_P = (Wi[:3, :3] @ piv_world + Wi[:3, 3]).cpu().numpy()
g["cams"] = describe((Wi[None] @ c2w).cpu().numpy(), piv_P, zm)
g["piv"] = (piv_P / zm).round(4).tolist()
if q.get("path"):
g["speed"] = [round(x, 3) for x in speed]
return P, S, full, tmap, pf, c2w, intr_t, render, cov, g
# --- endpoints --------------------------------------------------------------------------------------
@app.get("/")
def index():
return HTMLResponse(open(f"{HERE}/index.html").read()) # re-read per request, so frontend edits need no restart
@app.get("/samples")
def samples():
if not os.path.isdir(args.samples):
return []
return sorted(f for f in os.listdir(args.samples) if f.endswith(".mp4"))
@app.post("/upload")
async def upload(file: UploadFile = File(...)):
return ingest(await file.read(), file.filename)
@app.post("/sample")
async def sample(req: Request):
q = await req.json()
path = os.path.join(args.samples, os.path.basename(q["name"]))
return ingest(open(path, "rb").read(), os.path.basename(path))
@app.get("/frame/{clip}/{i}.jpg")
def frame(clip: str, i: int):
if clip not in CLIPS:
raise HTTPException(404, "unknown clip")
c = CLIPS[clip]
d = f"{args.work}/clips/{clip}/f"
os.makedirs(d, exist_ok=True)
fp = f"{d}/{i}.jpg"
if not os.path.exists(fp):
Image.fromarray(c["fr"][min(max(i, 0), c["n"] - 1)]).save(fp, "JPEG", quality=88)
return FileResponse(fp)
@app.post("/prepare")
async def prepare_ep(req: Request):
q = await req.json()
with GPU:
P = prepare(q["clip"], int(q["start"]), int(q["span_end"]))
r = poses(P, q)
return dict(box=list(P["box"]), canvas=list(P["canvas"]), cond_canvas=list(P["cond_canvas"]), ms=P["ms"], **r)
@app.post("/cloud")
async def cloud_ep(req: Request):
"""The 3D editor's scene: one source frame's coloured point cloud in the path's frame (W / zm, see
recam/path.py), one point every `stride` pixels, nothing beyond 10 pivot depths."""
q = await req.json()
start, stride = int(q["start"]), int(q.get("stride", 5))
with GPU:
P = prepare(q["clip"], start, int(q["span_end"]))
Z, i = P["S0"], int(q["frame"]) - start
assert 0 <= i < len(P["full0"]), f"frame {q['frame']} outside the prepared span"
_, zm, _, W, _, _ = unit(P, q)
Wi = torch.linalg.inv(W)
yy, xx = torch.meshgrid(torch.arange(RES, device=dev), torch.arange(RES, device=dev), indexing="ij")
g1 = torch.stack([xx, yy, torch.ones_like(xx)], -1).float().reshape(-1, 3)
rgb = resize_u8(P["full0"][i : i + 1], (RES, RES))[0]
X = ((g1 @ torch.linalg.inv(Z["intr"][i]).T).reshape(RES, RES, 3) * Z["depth"][i][..., None]) \
@ Z["c2w"][i][:3, :3].T + Z["c2w"][i][:3, 3]
X = (X @ Wi[:3, :3].T + Wi[:3, 3]) / zm
m = (Z["keep"][i] & P["picture"] & (Z["depth"][i] < 10 * zm))[::stride, ::stride]
pts, cols = X[::stride, ::stride][m].cpu(), rgb[::stride, ::stride][m].cpu()
return dict(n=len(pts), zm=zm, pts=pts.reshape(-1).numpy().round(3).tolist(), rgb=cols.reshape(-1).tolist())
@app.post("/warp1")
async def warp1_ep(req: Request):
"""One key's still: source frame `src` warped to the camera (`pos`, `look`, in the path's frame and unit),
at `cond_canvas`, as JPEG. What the model would be given at that instant."""
q = await req.json()
with GPU:
P = prepare(q["clip"], int(q["start"]), int(q["span_end"]))
Z, i = P["S0"], int(q["src"]) - int(q["start"])
assert 0 <= i < len(P["full0"]), f"source frame {q['src']} outside the prepared span"
_, zm, _, W, _, _ = unit(P, q)
pos, look = np.asarray(q["pos"], float) * zm, np.asarray(q["look"], float) * zm
c2w_P = np.eye(4)
c2w_P[:3, :3], c2w_P[:3, 3] = look_at(pos, look), pos
c2w = (W @ torch.tensor(c2w_P, dtype=W.dtype, device=dev))[None]
S = {k: v[i : i + 1] for k, v in Z.items()}
intr = S["intr"].clone()
intr[:, 0, 0] *= float(q.get("focal", 1))
intr[:, 1, 1] *= float(q.get("focal", 1))
render, _ = warp(S, torch.linalg.inv(c2w), intr, P["full0"][i : i + 1], P["box"], P["canvas"])
img = resize_u8(render, P["cond_canvas"][::-1])[0].cpu().numpy()
buf = io.BytesIO()
Image.fromarray(img).save(buf, "JPEG", quality=82)
return Response(buf.getvalue(), media_type="image/jpeg")
@app.post("/warp")
async def warp_ep(req: Request):
"""Tier 1. The whole trajectory, 0.24 s. The pane that decides whether a take is worth 25 s of B200."""
q = await req.json()
t0 = time.time()
with GPU:
P, S, full, tmap, pf, c2w, intr_t, render, cov, g = geo(q)
cc = P["cond_canvas"]
truth = resize_u8(render, cc[::-1]).cpu() # what the model is actually given
holes = None if q.get("lite") else truth.clone()
m = None if q.get("lite") else ~resize_u8(cov[..., None].to(torch.uint8) * 255, cc[::-1])[..., 0].bool().cpu()
if holes is not None:
holes[m] = torch.tensor([255, 0, 200], dtype=torch.uint8)
sketch = None if q.get("lite") else render.cpu()
d = f"{args.work}/warp/{q['clip']}"
os.makedirs(d, exist_ok=True)
tag = str(int(time.time() * 1000))
write_mp4(truth, fps=int(FPS), output_path=f"{d}/truth_{tag}.mp4")
g.update(truth=f"/warpfile/{q['clip']}/truth_{tag}.mp4",
cond_canvas=list(cc), canvas=list(P["canvas"]))
if sketch is not None:
write_mp4(holes, fps=int(FPS), output_path=f"{d}/holes_{tag}.mp4")
write_mp4(sketch, fps=int(FPS), output_path=f"{d}/sketch_{tag}.mp4")
g.update(holes=f"/warpfile/{q['clip']}/holes_{tag}.mp4", sketch=f"/warpfile/{q['clip']}/sketch_{tag}.mp4")
g["ms"] = int(1000 * (time.time() - t0))
print(f"[warp] {torch.cuda.memory_allocated() / 2**30:.1f} GiB live, {torch.cuda.max_memory_allocated() / 2**30:.1f} peak", flush=True)
return g
@app.get("/warpfile/{clip}/{name}")
def warpfile(clip: str, name: str):
if clip not in CLIPS:
raise HTTPException(404, "unknown clip")
return FileResponse(f"{args.work}/warp/{clip}/{os.path.basename(name)}")
@torch.no_grad()
def do_render(job, q):
"""`inference/sample.py`'s render path, resident. The seed lives one line before the encodes on purpose:
`pack` consumes global CPU RNG and the VAE posterior consumes global CUDA RNG, and the noise-augmented
condition rows are fed to the model and never overwritten -- so `torch.manual_seed` is what makes a take
reproducible."""
d = f"{args.work}/takes/{job}"
os.makedirs(d, exist_ok=True)
def stage(name, pct):
JOBS[job].update(stage=name, pct=pct, t=round(time.time() - JOBS[job]["t0"], 1))
print(f"[{job}] {name} {torch.cuda.memory_allocated() / 2**30:.1f} GiB live, "
f"{torch.cuda.max_memory_allocated() / 2**30:.1f} peak", flush=True)
json.dump(JOBS[job], open(f"{d}/status.tmp", "w"))
os.replace(f"{d}/status.tmp", f"{d}/status.json")
try:
with GPU:
stage("warp", 5)
P, S, full, tmap, pf, c2w, intr_t, render, cov, g = geo(q)
canvas, cc, box = P["canvas"], P["cond_canvas"], P["box"]
x0, y0, bw, bh, _ = box
n_out = len(tmap)
stage("encode", 15)
torch.manual_seed(int(q.get("seed", 1234))) # the one-line seed fix
render_c = resize_u8(render, cc[::-1])
source = resize_u8(full[:, y0 : y0 + bh, x0 : x0 + bw], canvas[::-1])
cond = resize_u8(full[:, y0 : y0 + bh, x0 : x0 + bw], cc[::-1])
dd = {"cond": encode_video(vae, cond)[0], "render": encode_video(vae, render_c)[0],
"target": encode_video(vae, source)[1]}
embed = torch.load(f"{ASSETS}/fixed_embed_{n_out}.pt", weights_only=False)
audio_x0 = torch.load(f"{ASSETS}/silence_audio_{n_out}.pt", weights_only=True)["audio_x0"].float()
batch = pack(dd["cond"], dd["render"], dd["target"], embed["prompt_embeds"][0], embed["text_token_tags"], audio_x0)
rows = denoise(transformer, batch, args.steps, args.flow_shift, dev,
on_step=lambda i, n: stage(f"forward {i + 1}/{n}", 25 + int(45 * i / n)))
stage("decode", 75)
out = decode_video(vae, rows, dd["target"].shape[1:])
src_c, ren_c = source.cpu(), render_c.cpu()
stage("write", 90)
for name, fr in ("out", out), ("render", ren_c), ("source", src_c):
write_mp4(fr, fps=int(FPS), output_path=f"{d}/{name}.mp4")
write_mp4(torch.cat([src_c, resize_u8(ren_c, canvas[::-1]), out], 2), fps=int(FPS), output_path=f"{d}/grid.mp4")
Image.fromarray(out[-1].numpy()).save(f"{d}/last.png")
JOBS[job].update(gauges=g, done=True)
stage("done", 100)
except Exception as e:
JOBS[job].update(error=f"{type(e).__name__}: {e}", done=True)
stage("error", 100)
if isinstance(e, torch.AcceleratorError):
die(e)
raise
def die(e):
"""A sticky CUDA error (device-side assert, illegal address) cannot be cleared in-process: every later
kernel fails too and the page goes dead. Exit 3; run.sh relaunches on that code only (~2 min reload)."""
print(f"CUDA context poisoned, exiting for restart: {e}", flush=True)
threading.Timer(0.5, os._exit, [3]).start()
@app.exception_handler(torch.AcceleratorError)
async def cuda_dead(req, e):
die(e)
return JSONResponse({"error": f"GPU fault: {e}"[:300] + " -- service restarting, ~2 min"}, status_code=503)
NJOB = itertools.count()
@app.post("/render")
async def render_ep(req: Request):
q = await req.json()
job = f"{int(time.time())}_{q['clip'][:6]}_{next(NJOB)}" # two clicks in the same second must not share a job dir
JOBS[job] = dict(job=job, t0=time.time(), stage="queued", pct=0, done=False, payload=q)
threading.Thread(target=do_render, args=(job, q), daemon=True).start()
return dict(job=job)
@app.get("/job/{job}")
def job_ep(job: str):
if job not in JOBS:
raise HTTPException(404, "unknown job")
j = dict(JOBS[job])
j.pop("t0", None)
return j
@app.get("/take/{job}/{name}")
def take(job: str, name: str):
if job not in JOBS:
raise HTTPException(404, "unknown job")
return FileResponse(f"{args.work}/takes/{job}/{os.path.basename(name)}")
uvicorn.run(app, host=args.host, port=args.port, log_level="warning")