mantrakp's picture
Isolate reference inference in a dedicated ZeroGPU worker
5c331a4 verified
Raw History Blame Contribute Delete
8.18 kB
"""CPU orchestration of independent ZeroGPU Spaces with per-call artifact isolation."""
from __future__ import annotations
import json
import os
from pathlib import Path
import shutil
from tempfile import TemporaryDirectory
import zipfile
import numpy as np
from PIL import Image
STAGES = ("reference", "geometry", "segment", "texture", "motion")
DEFAULT_WORKERS = {stage: f"mantrakp/component-studio-{stage}" for stage in STAGES}
def workers_from_environment():
configured = json.loads(os.environ.get("STUDIO_WORKERS", "{}"))
if not isinstance(configured, dict) or set(configured) - set(STAGES):
raise ValueError("STUDIO_WORKERS must map known stage names to Space IDs.")
workers = DEFAULT_WORKERS | configured
if any(not isinstance(value, str) or not value.strip() for value in workers.values()):
raise ValueError("Each worker must have a nonempty Space ID or URL.")
return workers
def artifact_path(value):
if isinstance(value, dict):
value = value.get("path")
if not isinstance(value, (str, Path)):
raise RuntimeError("Worker did not return a downloaded artifact.")
path = Path(value)
if not path.is_file() or path.stat().st_size == 0:
raise RuntimeError("Worker returned a missing or empty artifact.")
return path
def load_image(value):
with Image.open(artifact_path(value)) as image:
image.load()
return image.copy()
class RemoteModels:
def __init__(self, workers=None, client_factory=None, token=None):
self.workers = workers or workers_from_environment()
self.client_factory = client_factory
self.token = token
def _call(self, stage, directory, *args):
from gradio_client import Client
from huggingface_hub import get_token
# Construct a fresh client for each stage: GPU proxy authorization must not
# survive the CPU review/rigging interval between generation and motion.
factory = self.client_factory or Client
client = factory(self.workers[stage], hf_token=self.token or os.environ.get("HF_TOKEN") or get_token(),
download_files=str(directory), verbose=False)
return client.predict(*args, api_name=f"/{stage}")
def reference(self, image, prompt, seed):
from gradio_client import handle_file
with TemporaryDirectory(prefix="studio-reference-") as directory:
root = Path(directory)
source = None
if image is not None:
image.save(root / "input.png")
source = handle_file(str(root / "input.png"))
return load_image(self._call("reference", root, source, prompt, seed))
def geometry(self, image, seed, output):
from gradio_client import handle_file
from .mesh import load_mesh
with TemporaryDirectory(prefix="studio-geometry-") as directory:
root = Path(directory)
image.save(root / "input.png")
result = artifact_path(self._call("geometry", root, handle_file(str(root / "input.png")), seed))
load_mesh(result)
output = Path(output)
output.parent.mkdir(parents=True, exist_ok=True)
shutil.copyfile(result, output)
return output
def segment(self, mesh, count, seed):
from gradio_client import handle_file
with TemporaryDirectory(prefix="studio-segment-") as directory:
root = Path(directory)
np.savez(root / "mesh.npz", vertices=mesh.vertices, faces=mesh.faces)
result = self._call("segment", root, handle_file(str(root / "mesh.npz")), count, seed)
labels = np.load(artifact_path(result), allow_pickle=False)
if (labels.shape != (len(mesh.faces),) or labels.dtype.kind not in "iu"
or np.any(labels < 0) or np.any(labels >= count)):
raise RuntimeError("Segmentation worker returned invalid per-face labels.")
return labels.copy()
def upscale(self, parts, log):
from gradio_client import handle_file
if not parts:
raise ValueError("Texture upscaling requires at least one component.")
with TemporaryDirectory(prefix="studio-texture-") as directory:
root = Path(directory)
manifest = []
with zipfile.ZipFile(root / "input.zip", "w", zipfile.ZIP_DEFLATED) as bundle:
for index, (key, part) in enumerate(parts):
name = f"{index}.png"
part.visual.material.baseColorTexture.save(root / name)
bundle.write(root / name, name)
manifest.append({"key": key, "file": name})
bundle.writestr("manifest.json", json.dumps(manifest))
result, stats, events = self._call("texture", root, handle_file(str(root / "input.zip")))
images = []
with zipfile.ZipFile(artifact_path(result)) as bundle:
returned = json.loads(bundle.read("manifest.json"))
if returned != manifest:
raise RuntimeError("Texture worker changed the component manifest.")
for record in manifest:
with bundle.open(record["file"]) as stream, Image.open(stream) as image:
image.load()
images.append(image.copy())
# Validate everything before mutating the meshes. Geometry, UVs and all
# other material channels remain owned by the caller.
if not isinstance(stats, dict) or not isinstance(events, list):
raise RuntimeError("Texture worker returned invalid metadata.")
for (_, part), image in zip(parts, images, strict=True):
original = part.visual.material.baseColorTexture
if image.size != (original.width * 4, original.height * 4):
raise RuntimeError("Texture worker did not return a four-times upscaled image.")
for (_, part), image in zip(parts, images, strict=True):
part.visual.material.baseColorTexture = image
for event in events:
log(str(event))
return {**stats, "worker_space": self.workers["texture"]}
def motion(self, prompt, duration, seed, output):
with TemporaryDirectory(prefix="studio-motion-") as directory:
root = Path(directory)
npz, bvh, provenance = self._call("motion", root, prompt, duration, seed)
npz, bvh = artifact_path(npz), artifact_path(bvh)
with np.load(npz, allow_pickle=False) as data:
frames = round(duration * 30)
shapes = {"posed_joints": (frames, 77, 3), "global_rot_mats": (frames, 77, 3, 3),
"local_rot_mats": (frames, 77, 3, 3), "root_positions": (frames, 3)}
for key, shape in shapes.items():
if key not in data or data[key].shape != shape or not np.isfinite(data[key]).all():
raise RuntimeError(f"Motion worker returned invalid {key}.")
if ("foot_contacts" not in data or data["foot_contacts"].shape[0] != frames
or not np.isfinite(data["foot_contacts"]).all()):
raise RuntimeError("Motion worker returned invalid foot contacts.")
if not bvh.read_text().lstrip().startswith("HIERARCHY"):
raise RuntimeError("Motion worker returned an invalid BVH file.")
if not isinstance(provenance, dict) or provenance.get("frames") != frames:
raise RuntimeError("Motion worker returned invalid provenance.")
output = Path(output)
output.mkdir(parents=True, exist_ok=True)
shutil.copyfile(npz, output / "motion.npz")
shutil.copyfile(bvh, output / "motion.bvh")
provenance = {**provenance, "worker_space": self.workers["motion"]}
(output / "provenance.json").write_text(json.dumps(provenance, indent=2))
return {"npz": output / "motion.npz", "bvh": output / "motion.bvh", "provenance": provenance}