mantrakp's picture
Isolate reference inference in a dedicated ZeroGPU worker
5c331a4 verified
Raw History Blame Contribute Delete
5.13 kB
"""ZeroGPU workers inherit startup models; only CPU inputs/results cross their process queues."""
import json
import math
import os
from pathlib import Path
import sys
import spaces
ROOT = Path(__file__).resolve().parents[1]
RUNTIME = ROOT / ".runtime"
LOCK = json.loads((ROOT / "models.lock.json").read_text())
_MODELS = {}
@spaces.GPU(duration=60)
def reference_gpu(image, prompt, seed):
import torch
kwargs = {"prompt": prompt + ". Single complete asset, full silhouette visible, plain background, no text.",
"height": 1024, "width": 1024, "num_inference_steps": 4, "guidance_scale": 1.0,
"generator": torch.Generator("cuda").manual_seed(seed)}
if image is not None:
kwargs["image"] = image.convert("RGB")
return _MODELS["reference"](**kwargs).images[0]
@spaces.GPU(duration=180)
def geometry_gpu(image, seed, output):
import o_voxel
import torch
pipeline = _MODELS["geometry"]
with torch.inference_mode():
meshes, latents = pipeline.run(
image, seed=seed, pipeline_type="1024_cascade", return_latent=True,
sparse_structure_sampler_params={"steps": 12},
shape_slat_sampler_params={"steps": 12}, tex_slat_sampler_params={"steps": 12})
mesh = meshes[0]
mesh.simplify(16777216)
glb = o_voxel.postprocess.to_glb(
vertices=mesh.vertices, faces=mesh.faces, attr_volume=mesh.attrs, coords=mesh.coords,
attr_layout=pipeline.pbr_attr_layout, grid_size=latents[2],
aabb=[[-0.5, -0.5, -0.5], [0.5, 0.5, 0.5]], decimation_target=50000,
texture_size=1024, remesh=True, remesh_band=1, remesh_project=0, use_tqdm=False)
glb.export(output)
return output
@spaces.GPU(duration=90)
def segment_gpu(mesh, count, seed):
from .partfield import segment
return segment(_MODELS["segment"], mesh, count, seed)
def upscale_duration(parts):
tiles = sum(math.ceil(part.visual.material.baseColorTexture.width / 256) *
math.ceil(part.visual.material.baseColorTexture.height / 256) for _, part in parts)
return min(180, max(30, math.ceil(15 + tiles * 1.5)))
@spaces.GPU(duration=upscale_duration)
def upscale_gpu(parts):
from .upscale import upscale_parts
events = []
stats = upscale_parts(parts, _MODELS["sr"], events.append)
textures = {key: part.visual.material.baseColorTexture for key, part in parts}
return textures, stats, events
class NativeModels:
def __init__(self, role=None):
from scripts.prepare_runtime import worker_role
self.role = worker_role(role)
marker = "prepared.json" if self.role == "all" else f"prepared-{self.role}.json"
if not (RUNTIME / marker).exists():
raise RuntimeError(f"Prepare the {self.role} worker runtime before loading GPU models.")
if self.role in ("all", "reference"):
import torch
from diffusers import Flux2KleinPipeline
_MODELS["reference"] = Flux2KleinPipeline.from_pretrained(
"black-forest-labs/FLUX.2-klein-4B", revision=LOCK["black-forest-labs/FLUX.2-klein-4B"],
torch_dtype=torch.bfloat16).to("cuda")
if self.role in ("all", "geometry"):
sys.path.insert(0, str(RUNTIME / "trellis"))
os.environ.setdefault("ATTN_BACKEND", "flash_attn")
from trellis2.pipelines import Trellis2ImageTo3DPipeline
pipeline = Trellis2ImageTo3DPipeline.from_pretrained(str(RUNTIME / "trellis-model"))
pipeline.low_vram = True
pipeline.cuda()
_MODELS["geometry"] = pipeline
if self.role in ("all", "segment"):
sys.path.insert(0, str(RUNTIME / "PartField"))
from huggingface_hub import hf_hub_download
from .partfield import load_model
checkpoint = hf_hub_download("mikaelaangel/partfield-ckpt", "model_objaverse.ckpt",
revision=LOCK["mikaelaangel/partfield-ckpt"])
_MODELS["segment"] = load_model(RUNTIME / "PartField", checkpoint)
if self.role in ("all", "texture"):
from spandrel import ModelLoader
_MODELS["sr"] = ModelLoader().load_from_file(str(RUNTIME / "RealESRGAN_x4plus.pth")).eval().cuda()
if self.role in ("all", "motion"):
from .motion import load_motion_model
load_motion_model(RUNTIME, LOCK)
# Static functions avoid serializing model instances as a `self` argument to ZeroGPU workers.
reference = staticmethod(reference_gpu)
geometry = staticmethod(geometry_gpu)
segment = staticmethod(segment_gpu)
from .motion import generate_motion
motion = staticmethod(generate_motion)
def upscale(self, parts, log):
textures, stats, events = upscale_gpu(parts)
# Worker mutations do not propagate to the parent process. Apply returned CPU images explicitly.
for key, part in parts:
part.visual.material.baseColorTexture = textures[key]
for event in events:
log(event)
return stats