mantrakp's picture
Isolate reference inference in a dedicated ZeroGPU worker
5c331a4 verified
Raw History Blame Contribute Delete
7.87 kB
"""Gradio entry point for local development and Hugging Face ZeroGPU."""
import os
from pathlib import Path
import gradio as gr
from studio.contracts import Request
from studio.pipeline import STAGES, execute
ROOT = Path(__file__).resolve().parent
def load_models():
role = os.environ.get("STUDIO_WORKER_ROLE")
if role or os.environ.get("STUDIO_NATIVE") == "1":
from scripts.prepare_runtime import prepare
from studio.native import NativeModels
prepare(role=role)
return NativeModels(role=role)
if os.environ.get("SPACE_ID") or os.environ.get("STUDIO_WORKERS"):
from scripts.prepare_runtime import prepare_tools
from studio.remote import RemoteModels
prepare_tools()
return RemoteModels()
return None
def build_app(models=None):
def generate(image, prompt, parts, seed, animation_prompt="", animation_seconds=6,
progress=gr.Progress(), source_mesh=None):
try:
request = Request(image=Path(image) if image else None, prompt=prompt,
parts=int(parts), seed=int(seed), animation_prompt=animation_prompt,
animation_seconds=animation_seconds)
def notify(stage, detail):
progress((STAGES.index(stage), len(STAGES)), desc=detail)
run = execute(request, models, notify, source_mesh=source_mesh)
selected = "after" if run.data["refinement"]["verdict"]["accept"] else "before"
previews = sorted((run.path / selected).glob("*.png"))
previews = [str(path) for path in previews if path.name != "contact.png"]
report = run.data["refinement"]
status = f"{run.data['metrics']['components']} components ready. {report['summary']}"
if report["unresolved"]:
status += "\n\nRemaining issues: " + "; ".join(report["unresolved"])
animation = run.data.get("animation")
asset_path = run.path / "final.glb"
video_path = None
if not run.data.get("model_quality_passed", True):
status += "\n\nModel needs correction before animation: " + report["verdict"]["reason"]
if animation:
asset_path = run.path / animation["artifacts"]["animated_glb"]
video_path = str(run.path / animation["artifacts"]["video"])
status += "\n\n" + ("Animation accepted." if animation["accepted"]
else "Animation needs correction: " + animation["visual_review"]["reason"])
progress(1, desc="Ready" if run.data["status"] == "complete" else "Needs review")
return (str(asset_path), str(run.path / "reference.png"), previews,
str(run.path / "asset.zip"), status, run.data, video_path)
except Exception as error:
raise gr.Error(str(error)) from None
with gr.Blocks(title="Component Studio") as demo:
gr.Markdown("# Component Studio\nImage or idea → refined 3D → animated character.")
if models is None:
gr.Markdown("**Local interface preview.** GPU inference requires the ZeroGPU runtime described in README.md.")
with gr.Row():
with gr.Column(scale=2):
image = gr.Image(type="filepath", label="Reference image", image_mode="RGBA")
prompt = gr.Textbox(label="What do you want to make or change?", lines=3,
placeholder="A worn brass desk lamp with a green glass shade")
gr.Markdown("Use an image alone to reconstruct it. Use a prompt alone to create a reference. "
"Use both to edit the reference before reconstruction.")
animation_prompt = gr.Textbox(label="Animation", lines=2,
placeholder="Runs forward at a steady pace with natural arm swing")
gr.Markdown("Describe the motion for a humanoid character. Leave blank for a static asset.")
with gr.Accordion("Generation settings", open=False):
animation_seconds = gr.Slider(2, 12, value=6, step=1, label="Animation duration (seconds)")
parts = gr.Slider(2, 32, value=8, step=1, label="Target components")
seed = gr.Number(value=0, precision=0, label="Seed")
gr.Markdown("1024px initial textures · 4× base-color upscale · 10% memory reserve")
button = gr.Button("Create asset", variant="primary")
with gr.Accordion("Continue an existing textured mesh", open=False):
source_mesh = gr.File(label="Source GLB", file_types=[".glb"], type="filepath")
continue_button = gr.Button("Continue from geometry")
gr.Markdown("Supply the original reference above. Continue segmentation, upscaling and review without regenerating geometry.")
gr.Examples(examples=[
[str(ROOT / "assets/stool.png"), ""],
[None, "A worn brass desk lamp with a green glass shade, separate base, stem and shade"],
[str(ROOT / "assets/stool.png"), "Keep this stool's shape; make the seat dark walnut and the legs black metal"],
], inputs=[image, prompt], label="Image only · prompt only · image + prompt", cache_examples=False)
with gr.Column(scale=3):
model = gr.Model3D(label="Final asset", height=480)
video = gr.Video(label="Animation video", interactive=False)
status = gr.Markdown("Your assembled asset and individual components will appear here.")
download = gr.File(label="GLB, component GLBs, textures and review evidence")
with gr.Tab("Review views"):
gallery = gr.Gallery(label="Final review", columns=3)
with gr.Tab("Reference"):
reference = gr.Image(label="Reference used for geometry")
with gr.Tab("Run details"):
details = gr.JSON(label="Stage history and provenance")
button.click(generate, inputs=[image, prompt, parts, seed, animation_prompt, animation_seconds],
outputs=[model, reference, gallery, download, status, details, video],
concurrency_limit=1, concurrency_id="asset-pipeline", api_name="generate")
def continue_asset(image, source_mesh, prompt, parts, seed, animation_prompt="", animation_seconds=6,
progress=gr.Progress()):
if not image or not source_mesh:
raise gr.Error("Supply both the reference image and source GLB.")
return generate(image, prompt, parts, seed, animation_prompt, animation_seconds,
progress=progress, source_mesh=source_mesh)
continue_button.click(continue_asset, inputs=[image, source_mesh, prompt, parts, seed, animation_prompt, animation_seconds],
outputs=[model, reference, gallery, download, status, details, video],
concurrency_limit=1, concurrency_id="asset-pipeline", api_name="continue_asset")
return demo.queue(max_size=8, default_concurrency_limit=1)
if __name__ == "__main__":
models = load_models()
if role := os.environ.get("STUDIO_WORKER_ROLE"):
from studio.worker import build_worker
demo = build_worker(role, models)
else:
demo = build_app(models)
demo.launch(server_name="0.0.0.0" if os.environ.get("SPACE_ID") else "127.0.0.1",
server_port=int(os.environ.get("PORT", "7860")),
allowed_paths=[str(ROOT / "outputs"), str(ROOT / "assets")],
share=False)