mantrakp's picture
Isolate reference inference in a dedicated ZeroGPU worker
5c331a4 verified
Raw History Blame Contribute Delete
7.59 kB
"""Single-model ZeroGPU HTTP boundaries used by the CPU orchestrator."""
from __future__ import annotations
import json
from pathlib import Path
from tempfile import mkdtemp
from types import SimpleNamespace
import zipfile
import numpy as np
from PIL import Image
from .remote import STAGES
def read_texture_bundle(path):
"""Read PNG members directly without extracting any archive-provided paths."""
with zipfile.ZipFile(path) as bundle:
entries = bundle.infolist()
if len(entries) > 65 or sum(item.file_size for item in entries) > 256 * 1024**2:
raise ValueError("Texture bundle exceeds component or size limits.")
names = [item.filename for item in entries]
if len(names) != len(set(names)) or "manifest.json" not in names:
raise ValueError("Texture bundle requires unique members and a manifest.")
manifest = json.loads(bundle.read("manifest.json"))
if not isinstance(manifest, list) or not 1 <= len(manifest) <= 64:
raise ValueError("Texture manifest must contain 1–64 components.")
keys, files, parts = set(), set(), []
for record in manifest:
if not isinstance(record, dict) or set(record) != {"key", "file"}:
raise ValueError("Invalid texture manifest record.")
key, name = record["key"], record["file"]
if (not isinstance(key, str) or not key or key in keys or not isinstance(name, str)
or name != Path(name).name or "\\" in name or not name.endswith(".png") or name in files):
raise ValueError("Invalid or duplicate texture key or PNG filename.")
keys.add(key)
files.add(name)
with bundle.open(name) as stream, Image.open(stream) as source:
if source.format != "PNG" or source.width * source.height > 4096**2:
raise ValueError("Texture inputs must be PNG images of at most 4096 squared pixels.")
source.load()
image = source.convert("RGBA")
parts.append((key, SimpleNamespace(visual=SimpleNamespace(
material=SimpleNamespace(baseColorTexture=image)))))
if set(names) != files | {"manifest.json"}:
raise ValueError("Texture bundle contains unexpected members.")
return manifest, parts
def read_segment_mesh(path):
import trimesh
with np.load(path, allow_pickle=False) as arrays:
vertices, faces = arrays["vertices"], arrays["faces"]
if (vertices.ndim != 2 or vertices.shape[1] != 3 or not 3 <= len(vertices) <= 2_000_000
or not np.isfinite(vertices).all() or faces.ndim != 2 or faces.shape[1] != 3
or not 1 <= len(faces) <= 2_000_000 or faces.dtype.kind not in "iu"
or faces.min() < 0 or faces.max() >= len(vertices)):
raise ValueError("Segmentation requires finite vertices and valid triangle indices.")
return trimesh.Trimesh(vertices=vertices, faces=faces, process=False)
class WorkerEndpoints:
def __init__(self, models, output_root):
self.models = models
self.output_root = Path(output_root)
self.output_root.mkdir(parents=True, exist_ok=True)
def directory(self):
return Path(mkdtemp(prefix="request-", dir=self.output_root))
def reference(self, image, prompt, seed):
if not isinstance(prompt, str) or not prompt.strip() or len(prompt) > 4000:
raise ValueError("Reference prompt must contain 1–4000 characters.")
source = None
if image is not None:
with Image.open(image) as opened:
source = opened.convert("RGB")
result = self.models.reference(source, prompt, int(seed))
output = self.directory() / "reference.png"
result.save(output)
return str(output)
def geometry(self, image, seed):
with Image.open(image) as source:
image = source.convert("RGB")
output = self.directory() / "source.glb"
self.models.geometry(image, int(seed), output)
return str(output)
def segment(self, mesh, count, seed):
count = int(count)
if not 1 <= count <= 64:
raise ValueError("Component count must be between 1 and 64.")
mesh = read_segment_mesh(mesh)
labels = np.asarray(self.models.segment(mesh, count, int(seed)))
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("Model returned invalid segmentation labels.")
output = self.directory() / "labels.npy"
np.save(output, labels, allow_pickle=False)
return str(output)
def texture(self, textures):
manifest, parts = read_texture_bundle(textures)
events = []
stats = self.models.upscale(parts, events.append)
directory = self.directory()
output = directory / "textures.zip"
with zipfile.ZipFile(output, "w", zipfile.ZIP_DEFLATED) as bundle:
bundle.writestr("manifest.json", json.dumps(manifest))
for record, (_, part) in zip(manifest, parts, strict=True):
path = directory / record["file"]
part.visual.material.baseColorTexture.save(path)
bundle.write(path, record["file"])
return str(output), stats, events
def motion(self, prompt, duration, seed):
result = self.models.motion(prompt, float(duration), int(seed), self.directory())
return str(result["npz"]), str(result["bvh"]), result["provenance"]
def build_worker(role, models, output_root="outputs/worker"):
import gradio as gr
if role not in STAGES:
raise ValueError(f"Unknown worker role: {role}")
endpoints = WorkerEndpoints(models, output_root)
with gr.Blocks(title=f"Component Studio · {role}", delete_cache=(3600, 86400)) as app:
gr.Markdown(f"# Component Studio · {role}\nDedicated ZeroGPU worker for the character pipeline.")
if role == "reference":
inputs = [gr.File(label="Reference image (optional)", type="filepath"),
gr.Textbox(label="Prompt"), gr.Number(label="Seed", value=0, precision=0)]
outputs = [gr.File(label="Reference image")]
elif role == "geometry":
inputs = [gr.File(label="Reference image", type="filepath"),
gr.Number(label="Seed", value=0, precision=0)]
outputs = [gr.File(label="Textured GLB")]
elif role == "segment":
inputs = [gr.File(label="Mesh NPZ", type="filepath"),
gr.Number(label="Components", value=8, precision=0),
gr.Number(label="Seed", value=0, precision=0)]
outputs = [gr.File(label="Face labels NPY")]
elif role == "texture":
inputs = [gr.File(label="Texture bundle ZIP", type="filepath")]
outputs = [gr.File(label="Upscaled textures"), gr.JSON(label="Statistics"), gr.JSON(label="Events")]
else:
inputs = [gr.Textbox(label="Animation prompt"), gr.Number(label="Seconds", value=6),
gr.Number(label="Seed", value=0, precision=0)]
outputs = [gr.File(label="Motion NPZ"), gr.File(label="Motion BVH"), gr.JSON(label="Provenance")]
gr.Button("Generate").click(getattr(endpoints, role), inputs, outputs, api_name=role,
concurrency_limit=1)
return app.queue(default_concurrency_limit=1)