"""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}