mantrakp's picture
Isolate reference inference in a dedicated ZeroGPU worker
5c331a4 verified
Raw History Blame Contribute Delete
6.77 kB
"""Learn humanoid skinning remotely and retain the production geometry/materials."""
from dataclasses import dataclass
import json
import os
from pathlib import Path
import shutil
import subprocess
import time
from .artifacts import digest
@dataclass(frozen=True)
class RigResult:
mesh: Path
source_fbx: Path
provenance: dict
def rig_character(mesh: Path, output_dir: Path, notify=None) -> RigResult:
from gradio_client import Client, handle_file
from huggingface_hub import HfApi, get_token
output_dir.mkdir(parents=True, exist_ok=True)
space = os.environ.get("STUDIO_RIG_SPACE", "jasongzy/Make-It-Animatable")
token = os.environ.get("HF_TOKEN") or get_token()
revision = HfApi(token=token).space_info(space).sha
client = Client(
space,
token=token,
download_files=output_dir / "provider",
verbose=False,
)
endpoint = client.view_api(return_format="dict", print_info=False)["named_endpoints"]["/pipeline"]
restore_global = any(p["label"] == "Restore Global Transform" for p in endpoint["parameters"])
arguments = [handle_file(str(mesh.resolve())), False, "No", [], False, 0.01]
if restore_global:
arguments.append(True)
arguments.extend([restore_global, True, "LeftArm", False, None, True, True])
job = client.submit(*arguments, api_name="/pipeline")
receipt = {
"provider": space,
"provider_revision": revision,
"endpoint": "/pipeline",
"source_sha256": digest(mesh),
"started": time.time(),
"status": "running",
"restore_global": restore_global,
"reset_to_rest": False,
"animation_file": None,
}
receipt_path = output_dir / "rigging.json"
receipt_path.write_text(json.dumps(receipt, indent=2))
previous = None
try:
while not job.done():
status = job.status().code.name
if status != previous and notify:
notify(f"Learned character rig: {status.lower()}")
previous = status
time.sleep(2)
results = job.result()
# The generator's final yield can contain only Gradio skip updates. Keep the
# last concrete downloadable FBX among all yielded outputs.
candidates = [results, *job.outputs()]
fbx = None
normalized = None
for result in candidates:
for value in result if isinstance(result, (tuple, list)) else [result]:
if isinstance(value, dict):
value = value.get("value", value.get("path"))
if isinstance(value, dict):
value = value.get("path")
if isinstance(value, str) and value.lower().endswith(".fbx") and Path(value).is_file():
fbx = Path(value)
if isinstance(value, str) and Path(value).name == "normed.glb" and Path(value).is_file():
normalized = Path(value)
if fbx is None:
raise RuntimeError("Learned rig service did not return an FBX skin")
source_fbx = output_dir / "learned-rig.fbx"
shutil.copy2(fbx, source_fbx)
blender = os.environ.get("BLENDER_BIN") or shutil.which("blender")
if not blender:
mac = Path("/Applications/Blender.app/Contents/MacOS/Blender")
blender = str(mac) if mac.is_file() else None
if not blender:
raise RuntimeError("Blender is required to preserve production textures during rig transfer")
rigged = output_dir / "rigged.glb"
script = Path(__file__).resolve().parents[1] / "scripts/transfer_learned_rig.py"
normalization_args = []
if not restore_global:
if normalized is None:
raise RuntimeError("Rig service omitted the normalization mesh needed to restore coordinates")
normalization = output_dir / "normalization.json"
normalization.write_text(json.dumps(recover_normalization(mesh, normalized), indent=2))
normalization_args = ["--normalization", str(normalization.resolve())]
completed = subprocess.run(
[
blender,
"--background",
"--python-exit-code",
"1",
"--python",
str(script),
"--",
"--source",
str(mesh.resolve()),
"--rig",
str(source_fbx.resolve()),
"--output",
str(rigged.resolve()),
*normalization_args,
],
capture_output=True,
text=True,
timeout=600,
)
(output_dir / "transfer.log").write_text(completed.stdout + completed.stderr)
if completed.returncode or not rigged.is_file():
raise RuntimeError("Learned rig transfer failed; inspect rigging/transfer.log")
receipt.update(
status="complete",
finished=time.time(),
rig_sha256=digest(rigged),
source_fbx_sha256=digest(source_fbx),
transfer=json.loads(rigged.with_suffix(".transfer.json").read_text()),
)
return RigResult(rigged, source_fbx, receipt)
except Exception as exc:
receipt.update(status="failed", error=str(exc), finished=time.time())
raise
finally:
receipt_path.write_text(json.dumps(receipt, indent=2))
def recover_normalization(source: Path, normalized: Path) -> dict:
"""Recover the provider normalization using preserved indexed mesh topology."""
import numpy as np
import trimesh
original = trimesh.load(source, force="mesh")
canonical = trimesh.load(normalized, force="mesh")
if original.vertices.shape != canonical.vertices.shape or not np.array_equal(
original.faces, canonical.faces
):
raise RuntimeError("Rig normalization changed topology; cannot safely restore original coordinates")
homogeneous = np.column_stack((original.vertices, np.ones(len(original.vertices))))
transform = np.linalg.lstsq(homogeneous, canonical.vertices, rcond=None)[0]
residual = float(np.max(np.linalg.norm(homogeneous @ transform - canonical.vertices, axis=1)))
if residual > float(np.max(canonical.extents)) * 1e-5:
raise RuntimeError("Rig normalization is not an affine transform of the production mesh")
matrix = np.eye(4)
matrix[:3] = transform.T
conversion = np.array([[1, 0, 0, 0], [0, 0, -1, 0], [0, 1, 0, 0], [0, 0, 0, 1]])
return {
"inverse_blender": (conversion @ np.linalg.inv(matrix) @ np.linalg.inv(conversion)).tolist(),
"normalized_height": float(np.ptp(canonical.vertices[:, 1])),
"fit_residual": residual,
}