SkinTokens / demo.py
Egor
runtime glibc 2.39 install instead of Dockerfile
343c755
Raw
History Blame Contribute Delete
32.9 kB
import argparse
import importlib
import os
import sys
import tempfile
from pathlib import Path
from typing import List, Optional, Tuple
import gradio as gr
from torch import Tensor
from tqdm import tqdm
# ---------------------------------------------------------------------------
# ZeroGPU compatibility shim. The hosted HF Space provides the `spaces`
# package; running locally we substitute a no-op.
# ---------------------------------------------------------------------------
try:
spaces = importlib.import_module("spaces")
except Exception:
class _SpacesCompat:
@staticmethod
def GPU(*args, **kwargs):
if len(args) == 1 and callable(args[0]) and not kwargs:
return args[0]
def _decorator(fn):
return fn
return _decorator
spaces = _SpacesCompat()
os.environ.setdefault("XFORMERS_IGNORE_FLASH_VERSION_CHECK", "1")
gr.TEMP_DIR = "tmp_gradio"
# ---------------------------------------------------------------------------
# Install the bundled `bpy` wheel at runtime if it isn't already importable.
#
# Why this is non-trivial:
# - Putting the wheel in requirements.txt fails: HF Spaces' Docker build
# mounts only requirements.txt BEFORE the repo COPY, so the wheel path
# doesn't exist at pip-install time.
# - PyPI doesn't ship a bpy wheel matching this exact build (rc0 / cp312 /
# manylinux_2_39).
# - The `bpy-*.whl` committed in this repo gets auto-tracked by HF's LFS
# layer (Hub auto-LFS for blobs > ~10 MB even when .gitattributes doesn't
# list `*.whl`). The container's COPY-from-repo only carries the LFS
# *pointer* file โ€” a ~150-byte text stub โ€” not the actual wheel binary.
# So `pip install <wheel>` and `zipfile.ZipFile(<wheel>)` both fail with
# "is not a zip file" / "Wheel is invalid".
#
# So: we detect the LFS-pointer case and re-fetch the real wheel from the
# HF Hub at runtime (where the API resolves LFS server-side), then extract
# it directly into site-packages.
# ---------------------------------------------------------------------------
def _ensure_bpy_installed():
try:
import bpy # noqa: F401
return
except Exception:
pass
import glob
import sysconfig
import zipfile
here = os.path.dirname(os.path.abspath(__file__))
wheels = sorted(glob.glob(os.path.join(here, "bpy-*.whl")))
if not wheels:
print("[demo] WARNING: bpy not importable and no bundled wheel found", flush=True)
return
wheel = wheels[-1]
wheel_name = os.path.basename(wheel)
# Detect LFS pointer (text stub starting with "version https://git-lfs...").
is_real_zip = False
try:
with open(wheel, "rb") as f:
is_real_zip = f.read(4).startswith(b"PK")
except Exception:
pass
if not is_real_zip:
print(
f"[demo] {wheel_name} on disk is an LFS pointer ({os.path.getsize(wheel)} B); "
f"fetching real wheel from HF Hub...",
flush=True,
)
from huggingface_hub import hf_hub_download
space_id = os.environ.get("SPACE_ID", "VAST-AI/SkinTokens")
token = os.environ.get("HF_TOKEN") # set as a Space secret for private repos
wheel = hf_hub_download(
repo_id=space_id,
repo_type="space",
filename=wheel_name,
token=token,
)
print(f"[demo] fetched -> {wheel} ({os.path.getsize(wheel)} B)", flush=True)
site = sysconfig.get_paths()["purelib"]
print(f"[demo] Extracting {wheel_name} into {site}", flush=True)
with zipfile.ZipFile(wheel) as z:
z.extractall(site)
print("[demo] bpy wheel extracted.", flush=True)
_ensure_bpy_installed()
# ---------------------------------------------------------------------------
# Download model checkpoints (TokenRig + SkinTokens FSQ-CVAE) and the Qwen3
# tokenizer/config on first cold-start.
#
# These live in the *model* repo `VAST-AI/SkinTokens` (private), separate
# from this Space repo, so they aren't COPYed into the container. Re-uses
# `HF_TOKEN` from the Space secrets.
# ---------------------------------------------------------------------------
def _ensure_models_downloaded()
# ===== Install newer glibc for bpy compatibility =====
import os, subprocess, sys, tarfile, urllib.request, tempfile, shutil
def _ensure_glibc():
"""Install glibc 2.39 if current version is too old for bpy."""
try:
# Check current glibc version
result = subprocess.run(["ldd", "--version"], capture_output=True, text=True)
ver_line = result.stdout.split(chr(10))[0] if result.returncode == 0 else ""
if "2.39" in ver_line or "2.40" in ver_line or "2.41" in ver_line:
print(f"[glibc] Current version is sufficient: {ver_line.strip()}", flush=True)
return
print(f"[glibc] Current version too old: {ver_line.strip()}", flush=True)
except Exception:
pass
# Download and extract glibc 2.39 from Ubuntu 24.04 packages
GLIBC_DIR = os.path.join(tempfile.gettempdir(), "glibc239")
if os.path.exists(os.path.join(GLIBC_DIR, "lib", "libc.so.6")):
print("[glibc] Already installed", flush=True)
os.environ["LD_LIBRARY_PATH"] = os.path.join(GLIBC_DIR, "lib") + ":" + os.environ.get("LD_LIBRARY_PATH", "")
return
os.makedirs(GLIBC_DIR, exist_ok=True)
base_url = "http://archive.ubuntu.com/ubuntu/pool/main/g/glibc/"
packages = [
"libc6_2.39-0ubuntu8.4_amd64.deb",
]
for pkg in packages:
url = base_url + pkg
deb_path = os.path.join(GLIBC_DIR, pkg)
print(f"[glibc] Downloading {pkg}...", flush=True)
try:
urllib.request.urlretrieve(url, deb_path)
except Exception as e:
print(f"[glibc] Failed to download: {e}", flush=True)
# Try alternative URL
alt_url = f"http://archive.ubuntu.com/ubuntu/pool/main/g/glibc/{pkg}"
try:
urllib.request.urlretrieve(alt_url, deb_path)
except Exception:
print(f"[glibc] Alternative URL also failed", flush=True)
shutil.rmtree(GLIBC_DIR)
return
# Extract .deb
print(f"[glibc] Extracting {pkg}...", flush=True)
subprocess.run(["dpkg-deb", "-x", deb_path, GLIBC_DIR], check=True)
os.remove(deb_path)
# Verify installation
libc = os.path.join(GLIBC_DIR, "lib", "x86_64-linux-gnu", "libc.so.6")
if not os.path.exists(libc):
# libc might be in lib/ not lib/x86_64-linux-gnu/
alt_libc = os.path.join(GLIBC_DIR, "lib", "libc.so.6")
if os.path.exists(alt_libc):
libc = alt_libc
if os.path.exists(libc):
lib_dir = os.path.dirname(libc)
os.environ["LD_LIBRARY_PATH"] = lib_dir + ":" + os.environ.get("LD_LIBRARY_PATH", "")
print(f"[glibc] Installed glibc 2.39 at {lib_dir}", flush=True)
else:
print("[glibc] Installation failed โ€” running without glibc upgrade", flush=True)
shutil.rmtree(GLIBC_DIR)
_ensure_glibc()
:
here = os.path.dirname(os.path.abspath(__file__))
needed_ckpts = [
"experiments/skin_vae_2_10_32768/last.ckpt",
"experiments/articulation_xl_quantization_256_token_4/grpo_1400.ckpt",
]
qwen_dir = os.path.join(here, "models", "Qwen3-0.6B")
all_present = (
all(os.path.exists(os.path.join(here, p)) for p in needed_ckpts)
and os.path.exists(os.path.join(qwen_dir, "tokenizer.json"))
)
if all_present:
return
from huggingface_hub import hf_hub_download, snapshot_download
token = os.environ.get("HF_TOKEN")
for rel in needed_ckpts:
target = os.path.join(here, rel)
if os.path.exists(target):
continue
print(f"[demo] Downloading checkpoint: {rel}", flush=True)
hf_hub_download(
repo_id="VAST-AI/SkinTokens",
filename=rel,
local_dir=here,
token=token,
)
if not os.path.exists(os.path.join(qwen_dir, "tokenizer.json")):
print("[demo] Downloading Qwen3-0.6B tokenizer/config", flush=True)
snapshot_download(
repo_id="Qwen/Qwen3-0.6B",
local_dir=qwen_dir,
ignore_patterns=["*.bin", "*.safetensors"],
)
print("[demo] All checkpoints ready.", flush=True)
_ensure_models_downloaded()
# ===== DEBUG: test bpy import and store result =====
_DEBUG_BPY = {"status": "unknown", "stdout": "", "stderr": ""}
import subprocess, sys, os as _os
_here = _os.path.dirname(_os.path.abspath(__file__))
try:
_proc = subprocess.Popen(
[sys.executable, _os.path.join(_here, "test_import.py")],
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
)
_out, _err = _proc.communicate(timeout=60)
_DEBUG_BPY = {
"status": f"rc={_proc.returncode}",
"stdout": _out.decode()[:500],
"stderr": _err.decode()[:500],
}
except Exception as _e:
_DEBUG_BPY = {"status": f"exception: {_e}", "stdout": "", "stderr": ""}
print(f"[DEBUG] bpy test result: {_DEBUG_BPY['status']}", flush=True)
from src.data.dataset import DatasetConfig, RigDatasetModule
from src.data.transform import Transform
from src.model.tokenrig import TokenRigResult
from src.tokenizer.parse import get_tokenizer
from src.server.spec import (
BpyClient,
bpy_server_request,
is_bpy_server_alive,
get_model,
)
from src.data.vertex_group import voxel_skin
# ---------------------------------------------------------------------------
# Pre-warm `bpy_server` in the main (Gradio) process at module load.
#
# Why this is necessary on ZeroGPU: each user request runs inside a fresh
# `@spaces.GPU` worker process with a hard time budget (โ‰ˆ60 s on free tier).
# Importing the Blender shared object inside that budget burns 30โ€“60 s, so
# the worker is killed *during* bpy import โ€” manifesting as
# "GPU task aborted" before any model code runs.
#
# We start `bpy_server.py` here, in the always-running main process, so the
# slow bpy import happens exactly once at Space boot.
#
# Communication between the main process / ZeroGPU workers and the bpy
# worker happens via **file-based IPC** under ``/tmp/bpy_ipc/``, avoiding
# TCP listening ports (which trigger HF Space proxy issues) and HTTP body-
# size limits.
# ---------------------------------------------------------------------------
MODEL_CKPTS = [
"experiments/articulation_xl_quantization_256_token_4/grpo_1400.ckpt",
]
HF_PATHS = [
"None",
]
def get_dataloader_workers() -> int:
# Always use 0 workers โ€” the bpy IPC server is single-threaded and
# must process requests sequentially, so parallel loading provides
# no benefit. More importantly, DataLoader worker processes wrap
# exceptions from asset loading in a way that can leave the
# DataLoader iterator in an inconsistent state, blocking Gradio's
# event loop on subsequent requests.
return 0
# ---------------------------------------------------------------------------
# bpy_server lifecycle โ€” lazy start so the heavy import doesn't fight ZeroGPU
# during module load. Uses file-based IPC (no TCP listening).
# ---------------------------------------------------------------------------
_BPY_CLIENT: "Optional[BpyClient]" = None
def ensure_bpy_server_started(force: bool = False):
"""Start the bpy worker (if not already running) and wait until ready.
Uses a fast PID-file check to detect an already-running bpy_server
(works across processes, including ZeroGPU workers). Falls back to
launching a fresh subprocess when no server is found.
When *force* is True the existing worker (if any) is killed and a
fresh one is started unconditionally โ€” use after a timeout to
recover from a hung server.
"""
global _BPY_CLIENT
if not force:
# Fast path: we hold a live client handle (same-process, e.g. main).
if _BPY_CLIENT is not None and _BPY_CLIENT.is_alive():
return
# Cross-process fast path: another process already started the worker.
if is_bpy_server_alive():
return
# If we had a dead client handle, clean it up.
if _BPY_CLIENT is not None:
try:
_BPY_CLIENT.shutdown()
except Exception:
pass
# Launch a fresh bpy worker (this blocks until bpy is imported, ~30โ€“60 s
# on a cold container, but the @spaces.GPU duration budget accommodates).
_BPY_CLIENT = BpyClient.launch()
# ---------------------------------------------------------------------------
# Lazy model loading.
# ---------------------------------------------------------------------------
model = None
tokenizer = None
transform = None
CURRENT_MODEL_CKPT: Optional[str] = None
CURRENT_HF_PATH: Optional[str] = None
def load_model(model_ckpt: str, hf_path: Optional[str]) -> Tuple[str, str]:
global model, tokenizer, transform, CURRENT_MODEL_CKPT, CURRENT_HF_PATH
if hf_path == "None":
hf_path = None
if model is not None and model_ckpt == CURRENT_MODEL_CKPT and hf_path == CURRENT_HF_PATH:
return ("Model already loaded.", model_ckpt)
if not model_ckpt:
raise RuntimeError("model_ckpt is empty. Please select a checkpoint.")
print(f"Loading model: {model_ckpt}, hf_path={hf_path}")
model = get_model(model_ckpt, hf_path=hf_path)
assert model.tokenizer_config is not None
tokenizer = get_tokenizer(**model.tokenizer_config)
transform = Transform.parse(**model.transform_config["predict_transform"])
CURRENT_MODEL_CKPT = model_ckpt
CURRENT_HF_PATH = hf_path
return ("Model loaded.", model_ckpt)
# ---------------------------------------------------------------------------
# File utilities (CLI-side).
# ---------------------------------------------------------------------------
SUPPORTED_EXT = {".obj", ".fbx", ".glb"}
def collect_files(input_path: Path) -> List[Path]:
if input_path.is_file():
return [input_path]
files = []
for p in input_path.rglob("*"):
if p.suffix.lower() in SUPPORTED_EXT:
files.append(p)
return files
def map_output_path(in_path: Path, input_root: Path, output_root: Path) -> Path:
rel = in_path.relative_to(input_root)
return (output_root / rel).with_suffix(".glb")
# ---------------------------------------------------------------------------
# Core inference (shared by CLI and Gradio).
# ---------------------------------------------------------------------------
def run_rig(
filepaths: List[Path],
top_k: int,
top_p: float,
temperature: float,
repetition_penalty: float,
num_beams: int,
use_skeleton: bool,
use_transfer: bool,
use_postprocess: bool,
output_paths: List[Path],
model_ckpt: str,
hf_path: Optional[str],
):
assert len(filepaths) == len(output_paths)
ensure_bpy_server_started()
load_model(model_ckpt, hf_path)
datapath = {
"data_name": None,
"loader": "bpy_server",
"filepaths": {"articulation": [str(p) for p in filepaths]},
}
dataset_config = DatasetConfig.parse(
shuffle=False,
batch_size=1,
num_workers=get_dataloader_workers(),
pin_memory=get_dataloader_workers() > 0,
persistent_workers=False,
datapath=datapath,
).split_by_cls()
module = RigDatasetModule(
predict_dataset_config=dataset_config,
predict_transform=transform,
tokenizer=tokenizer,
process_fn=model._process_fn,
)
dataloader = module.predict_dataloader()["articulation"]
results_out = []
infer_device = model.device if model is not None else "cuda"
# Use explicit next() so we can catch errors from asset loading
# (which happens inside __getitem__ โ†’ next(dataloader)) per-file.
dataloader_iter = iter(dataloader)
for i in tqdm(range(len(filepaths)), desc="Processing"):
filepath = filepaths[i]
out_path = output_paths[i]
try:
batch = next(dataloader_iter)
except StopIteration:
break
except Exception as exc:
err_msg = str(exc)
print(f"[Error] Failed to load '{filepath}': {err_msg}", flush=True)
import traceback
traceback.print_exc()
# Force-restart if the server timed out (hung) so the next file
# gets a clean interpreter instead of hitting the same hang.
_restart_bpy_if_dead(force=("timeout" in err_msg.lower()))
continue
try:
batch = {
k: v.to(infer_device) if isinstance(v, Tensor) else v
for k, v in batch.items()
}
if not use_skeleton:
batch.pop("skeleton_tokens", None)
batch.pop("skeleton_mask", None)
batch["generate_kwargs"] = dict(
max_length=2048,
top_k=int(top_k),
top_p=float(top_p),
temperature=float(temperature),
repetition_penalty=float(repetition_penalty),
num_return_sequences=1,
num_beams=int(num_beams),
do_sample=True,
)
if "skeleton_tokens" in batch and "skeleton_mask" in batch:
mask = batch["skeleton_mask"][0] == 1
skeleton_tokens = batch["skeleton_tokens"][0][mask].cpu().numpy()
else:
skeleton_tokens = None
preds: List[TokenRigResult] = model.predict_step(
batch,
skeleton_tokens=[skeleton_tokens] if skeleton_tokens is not None else None,
make_asset=True,
)["results"]
asset = preds[0].asset
assert asset is not None
if use_postprocess:
voxel = asset.voxel(resolution=196)
asset.skin *= voxel_skin(
grid=0,
grid_coords=voxel.coords,
joints=asset.joints,
vertices=asset.vertices,
faces=asset.faces,
mode="square",
voxel_size=voxel.voxel_size,
)
asset.normalize_skin()
out_path.parent.mkdir(parents=True, exist_ok=True)
if use_transfer:
payload = dict(
source_asset=asset,
target_path=asset.path,
export_path=str(out_path),
group_per_vertex=4,
)
res = bpy_server_request("transfer", payload)
else:
payload = dict(
asset=asset,
filepath=str(out_path),
group_per_vertex=4,
)
res = bpy_server_request("export", payload)
if res != "ok":
print(f"[Error] {filepath}: {res}")
else:
print(f"[OK] Exported: {out_path}")
results_out.append(out_path)
except Exception as exc:
err_msg = str(exc)
print(f"[Error] Failed to process '{filepath}': {err_msg}", flush=True)
import traceback
traceback.print_exc()
# Force-restart if the server timed out (hung).
_restart_bpy_if_dead(force=("timeout" in err_msg.lower()))
continue
return results_out
def _restart_bpy_if_dead(force: bool = False):
"""If the bpy worker process died or hung, start a new one.
When *force* is True, the existing worker is killed and replaced
unconditionally โ€” use after a request timeout to guarantee a clean
bpy interpreter for the next file.
"""
if force:
print("[run_rig] bpy worker timed out โ€” force-restarting...", flush=True)
try:
ensure_bpy_server_started(force=True)
except Exception as exc:
print(f"[run_rig] Failed to force-restart bpy worker: {exc}", flush=True)
elif not is_bpy_server_alive():
print("[run_rig] bpy worker is dead โ€” restarting...", flush=True)
try:
ensure_bpy_server_started()
except Exception as exc:
print(f"[run_rig] Failed to restart bpy worker: {exc}", flush=True)
# ---------------------------------------------------------------------------
# CLI entry point.
# ---------------------------------------------------------------------------
def run_cli(args):
input_path = Path(args.input).resolve()
output_path = Path(args.output).resolve()
files = collect_files(input_path)
if not files:
raise RuntimeError("No valid 3D files found.")
if len(files) == 1 and output_path.suffix:
outputs = [output_path]
else:
outputs = [map_output_path(f, input_path, output_path) for f in files]
run_rig(
files,
args.top_k,
args.top_p,
args.temperature,
args.repetition_penalty,
args.num_beams,
args.use_skeleton,
args.use_transfer,
args.use_postprocess,
outputs,
args.model_ckpt,
args.hf_path,
)
# ---------------------------------------------------------------------------
# Gradio wrapper (with ZeroGPU duration estimator).
# ---------------------------------------------------------------------------
TOT = 0
def _gpu_duration(
files,
top_k,
top_p,
temperature,
repetition_penalty,
num_beams,
use_skeleton,
use_transfer,
use_postprocess,
model_ckpt,
hf_path,
):
# Cold workers spend ~30โ€“60 s importing bpy + loading the model before
# any GPU work. Give every request a generous 240 s floor.
file_count = len(files) if files is not None else 1
return min(900, max(10, 10 + 10 * file_count))
@spaces.GPU(duration=_gpu_duration)
def run_gradio(
files,
top_k,
top_p,
temperature,
repetition_penalty,
num_beams,
use_skeleton,
use_transfer,
use_postprocess,
model_ckpt,
hf_path,
):
try:
if not files:
return "Please upload at least one 3D model.", None
tmp_out = Path(tempfile.mkdtemp(prefix="tokenrig_"))
filepaths = [Path(f.name) for f in files]
global TOT
outputs = []
for filepath in filepaths:
TOT += 1
outputs.append(tmp_out / f"res_{TOT}.glb")
run_rig(
filepaths,
top_k,
top_p,
temperature,
repetition_penalty,
num_beams,
use_skeleton,
use_transfer,
use_postprocess,
outputs,
model_ckpt,
hf_path,
)
# Collect successfully exported files
succeeded = [str(p) for p in outputs if p.exists()]
if not succeeded:
return "โŒ No output files were produced โ€” check the server logs for details.", None
return f"โœ… Processed {len(succeeded)} model(s).", succeeded
except Exception as exc:
import traceback
traceback.print_exc()
return f"โŒ Error: {exc}", None
# ---------------------------------------------------------------------------
# Gradio UI.
# ---------------------------------------------------------------------------
def build_gradio_app():
model_ckpts = MODEL_CKPTS
hf_paths = HF_PATHS
default_ckpt = model_ckpts[0] if model_ckpts else ""
default_hf = hf_paths[0] if hf_paths else "None"
with gr.Blocks(title="SkinTokens ยท TokenRig Demo") as app:
gr.Markdown(
"""
## ๐Ÿฆด Mesh to Rig with [SkinTokens](https://zjp-shadow.github.io/works/SkinTokens/) ยท TokenRig
Automated **skeleton generation + skinning weight prediction** for any 3D mesh, via a unified
autoregressive model over learned *SkinTokens*. Successor to
[UniRig](https://github.com/VAST-AI-Research/UniRig) (SIGGRAPH&nbsp;'25).
* Upload one or more meshes โ†’ click **Run** โ†’ download a rigged `.glb`.
* **Paper**: [arXiv&nbsp;2602.04805](https://arxiv.org/abs/2602.04805) &nbsp;ยท&nbsp;
**Code**: [VAST-AI-Research/SkinTokens](https://github.com/VAST-AI-Research/SkinTokens) &nbsp;ยท&nbsp;
**Weights**: [๐Ÿค—&nbsp;VAST-AI/SkinTokens](https://huggingface.co/VAST-AI/SkinTokens)
* Looking for **image โ†’ rigged 3D** instead? Try our sibling Space
[๐Ÿค—&nbsp;VAST-AI/AniGen](https://huggingface.co/spaces/VAST-AI/AniGen).
* Want a full AI-powered 3D workspace? โ†’ [Tripo](https://www.tripo3d.ai)
"""
)
gr.HTML(
"""
<style>
@keyframes gentle-pulse {
0%, 100% { opacity: 1; }
50% { opacity: 0.35; }
}
</style>
<div style="text-align:left; color:#888; font-size:1em; line-height:1.6; margin: 4px 0 -4px 0;">
<span style="animation: gentle-pulse 3s ease-in-out infinite; display:inline-block;">&#128161; <b>Tips</b></span>&ensp;
Defaults work well for most meshes.
&nbsp;โ€ข If your mesh already has a skeleton and you only want skinning, enable
<b>Use existing skeleton</b> below.
&nbsp;โ€ข To keep your original textures and world scale, enable <b>Preserve original texture &amp; scale</b>.
</div>
"""
)
with gr.Row():
with gr.Column(scale=1):
files = gr.File(
label="3D Models ( .obj / .fbx / .glb, up to a few at a time )",
file_count="multiple",
file_types=[".obj", ".fbx", ".glb"],
)
with gr.Accordion("โš™๏ธ Generation Settings", open=False):
model_ckpt = gr.Dropdown(
choices=model_ckpts,
value=default_ckpt,
label="Model checkpoint",
info="TokenRig autoregressive rigging model. The default is the GRPO-refined checkpoint recommended for most assets.",
interactive=True,
)
# Keep the hf_path component for callback compatibility, but hide it
# from the UI since it currently only exposes the default ("None") option.
hf_path = gr.Dropdown(
choices=hf_paths,
value=default_hf,
label="HF path (advanced)",
visible=False,
)
gr.Markdown("**Sampling parameters** โ€” control autoregressive decoding of the rig.")
top_k = gr.Slider(
1, 200, value=5, step=1,
label="top_k",
info="Sample from the K most likely next tokens at each step. Lower = more deterministic output.",
)
top_p = gr.Slider(
0.1, 1.0, value=0.95, step=0.01,
label="top_p (nucleus)",
info="Sample from the smallest set of tokens whose cumulative probability โ‰ฅ p.",
)
temperature = gr.Slider(
0.1, 2.0, value=1.0, step=0.1,
label="temperature",
info="Softmax temperature. <1 sharpens the distribution (more conservative), >1 makes it flatter (more diverse).",
)
repetition_penalty = gr.Slider(
0.5, 3.0, value=2.0, step=0.1,
label="repetition_penalty",
info="Multiplicative penalty on tokens that have already been generated. 1.0 = no penalty.",
)
num_beams = gr.Slider(
1, 20, value=10, step=1,
label="num_beams",
info="Beam-search width. Larger = higher quality but slower; 1 disables beam search.",
)
gr.Markdown("**Pipeline toggles**")
use_skeleton = gr.Checkbox(
False,
label="Use existing skeleton (predict skinning only)",
info="If the uploaded file already contains a skeleton, keep it and only predict per-vertex skinning weights.",
)
use_transfer = gr.Checkbox(
False,
label="Preserve original texture & scale",
info="Transfer the predicted rig back onto the original (unprocessed) mesh, so textures and world units are preserved.",
)
use_postprocess = gr.Checkbox(
False,
label="Voxel skin post-processing",
info="Apply a voxel-based mask to the predicted skin weights before normalization. Slower.",
)
run_btn = gr.Button("๐Ÿš€ Run", variant="primary")
with gr.Column(scale=1):
log = gr.Textbox(label="Status", lines=2, interactive=False)
output = gr.File(label="Rigged GLB output", interactive=False)
gr.Markdown(
"""
**Notes**
- The output `.glb` contains the predicted **skeleton + skinning weights**. Import it in Blender (File โ†’ Import โ†’ glTF&nbsp;2.0) or any DCC tool that reads glTF.
- In Blender, if you see a `glTF_not_exported` placeholder node, you can safely remove it.
- On busy moments Zero-GPU may queue your request for ~10โ€“30&nbsp;s before inference starts โ€” the status box will update once the GPU is attached.
- Please do **not** upload confidential or NSFW content. See the
[project page](https://zjp-shadow.github.io/works/SkinTokens/) for paper-accurate results and the
[code repo](https://github.com/VAST-AI-Research/SkinTokens) for local / batch inference.
"""
)
debug_btn = gr.Button("๐Ÿ” Debug", variant="secondary", size="sm")
debug_out = gr.Textbox(label="Debug Info", lines=10, interactive=False)
debug_btn.click(lambda: _DEBUG_BPY, outputs=[debug_out])
run_btn.click(
run_gradio,
inputs=[
files,
top_k,
top_p,
temperature,
repetition_penalty,
num_beams,
use_skeleton,
use_transfer,
use_postprocess,
model_ckpt,
hf_path,
],
outputs=[log, output],
)
return app
demo = build_gradio_app()
# Note: we do NOT pre-warm `bpy_server` in the main process. `bpy_server.py`
# transitively imports `src.model.michelangelo.utils.misc`, whose
# module-level `use_flash3 = FLASH3()` calls `torch.cuda.get_device_name(0)`
# at import time. That call fails ("RuntimeError: No CUDA GPUs are
# available") in the main Gradio process on ZeroGPU, where the GPU is only
# attached inside `@spaces.GPU`-decorated workers. So the bpy_server boot
# happens on first request, inside the worker.
# ---------------------------------------------------------------------------
# Entry point.
# ---------------------------------------------------------------------------
if __name__ == "__main__":
parser = argparse.ArgumentParser("TokenRig Demo")
parser.add_argument("--input", help="Input file or directory")
parser.add_argument("--output", help="Output file or directory")
parser.add_argument("--top_k", type=int, default=5)
parser.add_argument("--top_p", type=float, default=0.95)
parser.add_argument("--temperature", type=float, default=1.0)
parser.add_argument("--repetition_penalty", type=float, default=2.0)
parser.add_argument("--num_beams", type=int, default=10)
parser.add_argument("--use_skeleton", action="store_true")
parser.add_argument("--use_transfer", action="store_true")
parser.add_argument("--use_postprocess", action="store_true")
parser.add_argument("--model_ckpt", default=MODEL_CKPTS[0] if MODEL_CKPTS else "")
parser.add_argument("--hf_path", default=None)
parser.add_argument("--gradio", action="store_true")
args = parser.parse_args()
if args.gradio or not args.input:
demo.queue()
demo.launch(ssr_mode=False)
else:
ensure_bpy_server_started()
run_cli(args)