dvd-image / app_dvd_image.py
Zhengrui's picture
Remove ZeroGPU smoke test
d61df16 verified
Raw
History Blame
43.9 kB
import os
os.environ.setdefault("SPCONV_ALGO", "native")
os.environ.setdefault("ATTN_BACKEND", "flash_attn")
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
os.environ.setdefault("DVD_MODEL_REPO", "Zhengrui/dvd")
try:
import huggingface_hub
if not hasattr(huggingface_hub, "HfFolder"):
class HfFolder:
@staticmethod
def get_token():
return huggingface_hub.get_token()
@staticmethod
def save_token(token):
return huggingface_hub.login(token=token, add_to_git_credential=False)
huggingface_hub.HfFolder = HfFolder
except Exception:
pass
import gradio as gr
import torch
try:
import spaces
except ImportError:
class _SpacesFallback:
@staticmethod
def GPU(duration=180):
return lambda fn: fn
spaces = _SpacesFallback()
import argparse
import gc
import shutil
import uuid
from pathlib import Path
import gradio_client.utils as gradio_client_utils
import numpy as np
from PIL import Image
from starlette.templating import Jinja2Templates
MAX_SEED = np.iinfo(np.int32).max
RESOLUTION = 64
ROOT_DIR = Path(__file__).resolve().parent
TMP_DIR = ROOT_DIR / "tmp" / "dvd_app"
TMP_DIR.mkdir(parents=True, exist_ok=True)
GEN_DVD_CONFIG = os.environ.get("DVD_GEN_CONFIG", "ckpts/dvd_img.json")
GEN_DVD_CKPT = os.environ.get("DVD_GEN_CKPT", "ckpts/dvd_img.safetensors")
EDIT_DVD_CONFIG = os.environ.get("DVD_EDIT_CONFIG", "ckpts/dvd_img_BSP_ft.json")
EDIT_DVD_CKPT = os.environ.get("DVD_EDIT_CKPT", "ckpts/dvd_img_BSP_ft.safetensors")
TRELLIS_IMAGE_MODEL = os.environ.get("TRELLIS_IMAGE_MODEL", "microsoft/TRELLIS-image-large")
DVD_MODEL_REPO = os.environ.get("DVD_MODEL_REPO")
DVD_MODEL_SUBFOLDER = os.environ.get("DVD_MODEL_SUBFOLDER") or None
DVD_MODEL_REVISION = os.environ.get("DVD_MODEL_REVISION") or None
DVD_MODEL_TOKEN = os.environ.get("DVD_MODEL_TOKEN") or os.environ.get("HF_TOKEN") or None
IMAGE_EXTENSIONS = {".png", ".jpg", ".jpeg", ".webp"}
def parse_camera_position(value: str) -> tuple[float, float, float]:
parts = [float(part.strip()) for part in value.split(",")]
if len(parts) != 3:
raise ValueError("DVD_VOXEL_CAMERA_POSITION must be three comma-separated numbers, e.g. -180,90,3")
return tuple(parts)
VOXEL_CAMERA_POSITION = parse_camera_position(os.environ.get("DVD_VOXEL_CAMERA_POSITION", "-180,90,3"))
dvd_gen_pipeline = None
dvd_edit_pipeline = None
trellis_pipeline = None
DVDImageToVoxelPipeline = None
TrellisImageTo3DPipeline = None
as_voxel_output = None
export_cubified_voxels = None
run_image_stage2_from_dvd_voxels = None
_postprocessing_utils = None
def log_event(message: str):
print(f"[DVD Space] {message}", flush=True)
def _download_file(url: str, dest: Path):
import urllib.request
dest.parent.mkdir(parents=True, exist_ok=True)
if dest.exists() and dest.stat().st_size > 0:
log_event(f"using cached {dest.name}")
return
tmp = dest.with_name(dest.name + ".tmp")
if tmp.exists():
tmp.unlink()
log_event(f"downloading {url} -> {dest}")
urllib.request.urlretrieve(url, tmp)
tmp.replace(dest)
def prefetch_dvd_weights():
from huggingface_hub import hf_hub_download, snapshot_download
if DVD_MODEL_REPO:
filenames = (
"dvd_img.json",
"dvd_img.safetensors",
"dvd_img_BSP_ft.json",
"dvd_img_BSP_ft.safetensors",
)
for filename in filenames:
log_event(f"prefetching {DVD_MODEL_REPO}/{filename}")
hf_hub_download(
repo_id=DVD_MODEL_REPO,
filename=filename,
subfolder=DVD_MODEL_SUBFOLDER,
revision=DVD_MODEL_REVISION,
token=DVD_MODEL_TOKEN,
)
if os.environ.get("DVD_PREFETCH_TRELLIS", "0").lower() not in {"0", "false", "no", "off"}:
log_event(f"prefetching TRELLIS stage2 model {TRELLIS_IMAGE_MODEL}")
snapshot_download(repo_id=TRELLIS_IMAGE_MODEL)
def prefetch_dinov2():
import shutil
import zipfile
hub_dir = Path(torch.hub.get_dir())
hub_dir.mkdir(parents=True, exist_ok=True)
repo_dir = hub_dir / "facebookresearch_dinov2_main"
if repo_dir.exists():
log_event(f"using cached DINOv2 repo {repo_dir}")
else:
zip_path = hub_dir / "main.zip"
_download_file("https://github.com/facebookresearch/dinov2/zipball/main", zip_path)
extract_tmp = hub_dir / "_dvd_dinov2_extract"
if extract_tmp.exists():
shutil.rmtree(extract_tmp)
extract_tmp.mkdir(parents=True, exist_ok=True)
log_event(f"extracting DINOv2 repo to {repo_dir}")
with zipfile.ZipFile(zip_path) as zf:
zf.extractall(extract_tmp)
top_level = zf.namelist()[0].split("/", 1)[0]
shutil.move(str(extract_tmp / top_level), str(repo_dir))
shutil.rmtree(extract_tmp, ignore_errors=True)
ckpt_dir = hub_dir / "checkpoints"
_download_file(
"https://dl.fbaipublicfiles.com/dinov2/dinov2_vitl14/dinov2_vitl14_reg4_pretrain.pth",
ckpt_dir / "dinov2_vitl14_reg4_pretrain.pth",
)
def prefetch_assets():
if os.environ.get("DVD_PREFETCH_ASSETS", "1").lower() in {"0", "false", "no", "off"}:
log_event("asset prefetch disabled")
return
try:
log_event("asset prefetch start")
prefetch_dvd_weights()
prefetch_dinov2()
log_event("asset prefetch done")
except Exception as exc:
log_event(f"asset prefetch failed; continuing without prefetch: {exc}")
def ensure_dvd_imports():
global DVDImageToVoxelPipeline
global TrellisImageTo3DPipeline
global as_voxel_output
global export_cubified_voxels
global run_image_stage2_from_dvd_voxels
if DVDImageToVoxelPipeline is not None:
return
log_event("importing DVD image/TRELLIS modules")
from dvd import (
DVDImageToVoxelPipeline as _DVDImageToVoxelPipeline,
TrellisImageTo3DPipeline as _TrellisImageTo3DPipeline,
as_voxel_output as _as_voxel_output,
export_cubified_voxels as _export_cubified_voxels,
run_image_stage2_from_dvd_voxels as _run_image_stage2_from_dvd_voxels,
)
DVDImageToVoxelPipeline = _DVDImageToVoxelPipeline
TrellisImageTo3DPipeline = _TrellisImageTo3DPipeline
as_voxel_output = _as_voxel_output
export_cubified_voxels = _export_cubified_voxels
run_image_stage2_from_dvd_voxels = _run_image_stage2_from_dvd_voxels
def ensure_pipeline_device(pipeline, device: str, label: str):
target = torch.device(device)
current = pipeline.device
if current != target:
log_event(f"moving {label} pipeline from {current} to {target}")
pipeline.to(target)
log_event(f"{label} pipeline is on {target}")
return pipeline
def ensure_zero_gpu_extensions():
"""Verify TRELLIS CUDA render extensions installed by Space requirements."""
import importlib
missing = []
for module_name in ("nvdiffrast.torch", "diff_gaussian_rasterization"):
try:
importlib.import_module(module_name)
except ImportError as exc:
missing.append(f"{module_name}: {exc}")
if missing:
raise RuntimeError(
"Missing prebuilt CUDA render extensions. Rebuild/redeploy the Space wheels for: "
+ "; ".join(missing)
)
log_event("CUDA render extensions ready")
def get_postprocessing_utils():
global _postprocessing_utils
if _postprocessing_utils is not None:
return _postprocessing_utils
ensure_zero_gpu_extensions()
from trellis.utils import postprocessing_utils
_postprocessing_utils = postprocessing_utils
return _postprocessing_utils
_original_json_schema_to_python_type = gradio_client_utils._json_schema_to_python_type
def _safe_json_schema_to_python_type(schema, defs):
if isinstance(schema, bool):
return "Any"
if isinstance(schema, dict) and isinstance(schema.get("additionalProperties"), bool):
schema = dict(schema)
if schema["additionalProperties"]:
schema["additionalProperties"] = {}
else:
schema.pop("additionalProperties")
return _original_json_schema_to_python_type(schema, defs)
gradio_client_utils._json_schema_to_python_type = _safe_json_schema_to_python_type
_original_template_response = Jinja2Templates.TemplateResponse
def _template_response_compat(self, *args, **kwargs):
if args and isinstance(args[0], str):
name = args[0]
context = args[1] if len(args) > 1 else kwargs.pop("context", None)
if isinstance(context, dict) and "request" in context:
return _original_template_response(self, context["request"], name, context, *args[2:], **kwargs)
return _original_template_response(self, *args, **kwargs)
Jinja2Templates.TemplateResponse = _template_response_compat
def list_asset_files(directory: str, suffixes: set[str]) -> list[Path]:
path = ROOT_DIR / directory
if not path.exists():
return []
return sorted(p for p in path.iterdir() if p.is_file() and p.suffix.lower() in suffixes)
def asset_label(path: Path) -> str:
return path.stem.replace("_", " ")
GENERATION_IMAGE_EXAMPLES = [
(asset_label(path), str(path)) for path in list_asset_files("assets/example_image", IMAGE_EXTENSIONS)
]
EDIT_IMAGE_EXAMPLES = [
(asset_label(path), str(path)) for path in list_asset_files("assets/example_image_edit", IMAGE_EXTENSIONS)
]
EDIT_VOXEL_EXAMPLES = [
(asset_label(path), str(path)) for path in list_asset_files("assets/example_voxel_edit", {".npy", ".npz", ".pt", ".pth"})
]
def voxel_viewer(label: str, exposure: float = 5.0, height: int = 300):
return gr.Model3D(
label=label,
height=height,
camera_position=VOXEL_CAMERA_POSITION,
)
def get_device(device_arg: str) -> str:
if device_arg == "auto":
return "cuda" if torch.cuda.is_available() else "cpu"
return device_arg
def default_server_port():
port = os.environ.get("GRADIO_SERVER_PORT")
return int(port) if port else None
def load_dvd_pipeline_variant(device: str, variant: str):
ensure_dvd_imports()
if DVD_MODEL_REPO:
common_kwargs = {
"device": device,
"subfolder": DVD_MODEL_SUBFOLDER,
"revision": DVD_MODEL_REVISION,
"token": DVD_MODEL_TOKEN,
}
return DVDImageToVoxelPipeline.from_pretrained(DVD_MODEL_REPO, variant=variant, **common_kwargs)
if variant == "base":
return DVDImageToVoxelPipeline.from_files(GEN_DVD_CONFIG, GEN_DVD_CKPT, device=device)
if variant == "bsp":
return DVDImageToVoxelPipeline.from_files(EDIT_DVD_CONFIG, EDIT_DVD_CKPT, device=device)
raise ValueError(f"Unsupported DVD pipeline variant: {variant}")
def load_dvd_pipelines(device: str):
return (
load_dvd_pipeline_variant(device, "base"),
load_dvd_pipeline_variant(device, "bsp"),
)
def release_pipeline(name: str, label: str):
global dvd_gen_pipeline, dvd_edit_pipeline, trellis_pipeline
pipelines = {
"dvd_gen": dvd_gen_pipeline,
"dvd_edit": dvd_edit_pipeline,
"trellis": trellis_pipeline,
}
pipeline = pipelines[name]
if pipeline is None:
return
log_event(f"releasing {label} pipeline")
try:
pipeline.to("cpu")
except Exception as exc:
log_event(f"{label} pipeline CPU release warning: {exc}")
if name == "dvd_gen":
dvd_gen_pipeline = None
elif name == "dvd_edit":
dvd_edit_pipeline = None
elif name == "trellis":
trellis_pipeline = None
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
log_event(f"released {label} pipeline")
def release_dvd_pipelines():
release_pipeline("dvd_gen", "DVD generation")
release_pipeline("dvd_edit", "DVD editing")
def release_trellis_pipeline():
release_pipeline("trellis", "TRELLIS stage2")
def ensure_dvd_gen_pipeline(device: str | None = None):
global dvd_gen_pipeline
device = device or os.environ.get("DVD_SPACE_DEVICE", "cuda")
release_trellis_pipeline()
if dvd_gen_pipeline is None:
log_event(f"loading DVD generation pipeline on {device}")
dvd_gen_pipeline = load_dvd_pipeline_variant(device, "base")
log_event("DVD generation pipeline ready")
else:
ensure_pipeline_device(dvd_gen_pipeline, device, "DVD generation")
return dvd_gen_pipeline
def ensure_dvd_edit_pipeline(device: str | None = None):
global dvd_edit_pipeline
device = device or os.environ.get("DVD_SPACE_DEVICE", "cuda")
release_pipeline("dvd_gen", "DVD generation")
release_trellis_pipeline()
if dvd_edit_pipeline is None:
log_event(f"loading DVD editing pipeline on {device}")
dvd_edit_pipeline = load_dvd_pipeline_variant(device, "bsp")
log_event("DVD editing pipeline ready")
else:
ensure_pipeline_device(dvd_edit_pipeline, device, "DVD editing")
return dvd_edit_pipeline
def ensure_trellis_pipeline(device: str | None = None):
ensure_dvd_imports()
global trellis_pipeline
device = device or os.environ.get("DVD_SPACE_DEVICE", "cuda")
release_dvd_pipelines()
if trellis_pipeline is None:
log_event(f"loading TRELLIS stage2 pipeline on {device}")
trellis_pipeline = TrellisImageTo3DPipeline.from_pretrained(TRELLIS_IMAGE_MODEL)
trellis_pipeline.to(device)
log_event("TRELLIS stage2 pipeline ready")
else:
ensure_pipeline_device(trellis_pipeline, device, "TRELLIS stage2")
return trellis_pipeline
def preload_zero_gpu_models(device: str | None = None, preload: str | None = None):
device = device or os.environ.get("DVD_SPACE_DEVICE", "cuda")
preload = preload or os.environ.get("DVD_STARTUP_PRELOAD", "none")
requested = {
item.strip().lower()
for item in preload.replace(";", ",").split(",")
if item.strip()
}
if requested & {"0", "false", "no", "none", "off"}:
log_event("startup model preload disabled")
return
ensure_dvd_imports()
if requested & {"all", "gen", "generation", "dvd_gen", "image"}:
ensure_dvd_gen_pipeline(device)
if requested & {"all", "edit", "editing", "dvd_edit"}:
ensure_dvd_edit_pipeline(device)
if requested & {"all", "stage2", "trellis"}:
ensure_zero_gpu_extensions()
ensure_trellis_pipeline(device)
def start_session(req: gr.Request):
user_dir = TMP_DIR / str(req.session_hash)
user_dir.mkdir(parents=True, exist_ok=True)
def end_session(req: gr.Request):
user_dir = TMP_DIR / str(req.session_hash)
if user_dir.exists():
shutil.rmtree(user_dir)
def session_path(req: gr.Request, name: str) -> str:
user_dir = TMP_DIR / str(req.session_hash)
user_dir.mkdir(parents=True, exist_ok=True)
return str(user_dir / name)
def worker_path(name: str) -> str:
user_dir = TMP_DIR / f"worker-{uuid.uuid4().hex}"
user_dir.mkdir(parents=True, exist_ok=True)
return str(user_dir / name)
def get_seed(randomize_seed: bool, seed: int) -> int:
return int(np.random.randint(0, MAX_SEED)) if randomize_seed else int(seed)
def dvd_cfg_schedule(mode: str, constant: float, early: float, late: float, split: float):
if mode == "Default schedule":
return None
if mode == "Constant":
return float(constant)
split = float(split)
return lambda t: float(early) if t < split else float(late)
def dvd_sampler_kwargs(
steps: int,
cfg_mode: str,
cfg_constant: float,
cfg_early: float,
cfg_late: float,
cfg_split: float,
) -> dict:
kwargs = {"steps": int(steps)}
cfg_strength = dvd_cfg_schedule(cfg_mode, cfg_constant, cfg_early, cfg_late, cfg_split)
if cfg_strength is not None:
kwargs["cfg_strength"] = cfg_strength
return kwargs
def slat_sampler_params(steps: int, cfg_strength: float) -> dict:
return {
"steps": int(steps),
"cfg_strength": float(cfg_strength),
}
def voxel_output(voxels, resolution: int = RESOLUTION):
ensure_dvd_imports()
if isinstance(voxels, np.lib.npyio.NpzFile):
key = next((k for k in ("coords", "voxels", "samples") if k in voxels.files), voxels.files[0])
voxels = voxels[key]
if isinstance(voxels, np.ndarray):
voxels = torch.as_tensor(voxels)
elif isinstance(voxels, dict):
for key in ("coords", "voxels", "samples"):
if key in voxels:
voxels = voxels[key]
break
else:
raise gr.Error("Voxel dict must contain one of: coords, voxels, samples.")
return voxel_output(voxels, resolution=resolution)
elif isinstance(voxels, (list, tuple)):
voxels = torch.as_tensor(voxels)
output = as_voxel_output(voxels, resolution=resolution)
return as_voxel_output(output.samples.detach().cpu().long(), resolution=resolution)
def voxel_state(voxels, resolution: int = RESOLUTION):
output = voxel_output(voxels, resolution=resolution)
return output.coords_without_batch.detach().cpu().numpy().astype(np.int32)
def load_voxel_file(file, resolution: int = RESOLUTION):
if file is None:
raise gr.Error("Please upload a voxel coordinate file.")
path = None
if isinstance(file, (str, Path)):
path = str(file)
elif isinstance(file, dict):
path = file.get("path") or file.get("name") or file.get("orig_name")
if path is None:
return voxel_output(file, resolution=resolution)
elif hasattr(file, "path"):
path = file.path
elif hasattr(file, "name"):
path = file.name
else:
return voxel_output(file, resolution=resolution)
suffix = Path(path).suffix.lower()
if suffix == ".npy":
data = np.load(path, allow_pickle=False)
elif suffix == ".npz":
data = np.load(path, allow_pickle=False)
elif suffix in {".pt", ".pth"}:
data = torch.load(path, map_location="cpu")
else:
raise gr.Error(f"Unsupported voxel file type: {suffix}. Use .npy, .npz, .pt, or .pth.")
return voxel_output(data, resolution=resolution)
def load_image_from_path(path: str):
if not path:
raise gr.Error("Please select an example image.")
return Image.open(path).convert("RGBA")
def load_generation_example(image_path: str):
return load_image_from_path(image_path)
def load_edit_image_example(image_path: str):
return load_image_from_path(image_path)
@spaces.GPU(duration=60)
def load_edit_voxel_example(voxel_path: str, req: gr.Request):
voxels = load_voxel_file(voxel_path)
mesh_path = voxel_to_mesh(voxels, session_path(req, "example_edit_voxels.glb"))
return voxel_state(voxels), mesh_path
def save_voxel_coords(voxels, path: str) -> str:
ensure_dvd_imports()
output = as_voxel_output(voxels, resolution=RESOLUTION)
np.save(path, output.coords_without_batch.numpy())
return path
def voxel_to_mesh(voxels, path: str) -> str:
ensure_dvd_imports()
return export_cubified_voxels(voxels, path, resolution=RESOLUTION)
def rotate_voxels(voxels, axis: str, req: gr.Request):
if voxels is None:
raise gr.Error("No editing voxels available. Upload voxels or transfer generated voxels first.")
output = load_voxel_file(voxels)
samples = output.samples.clone()
axis_to_dims = {
"x": (2, 3),
"y": (1, 3),
"z": (1, 2),
}
samples = torch.rot90(samples, k=1, dims=axis_to_dims[axis])
rotated = as_voxel_output(samples, resolution=RESOLUTION)
mesh_path = voxel_to_mesh(rotated, session_path(req, f"edit_voxels_rot_{axis}.glb"))
return voxel_state(rotated), mesh_path
@spaces.GPU(duration=60)
def rotate_x(voxels, req: gr.Request):
return rotate_voxels(voxels, "x", req)
@spaces.GPU(duration=60)
def rotate_y(voxels, req: gr.Request):
return rotate_voxels(voxels, "y", req)
@spaces.GPU(duration=60)
def rotate_z(voxels, req: gr.Request):
return rotate_voxels(voxels, "z", req)
def build_edit_mask(
use_1, x0_1, x1_1, y0_1, y1_1, z0_1, z1_1,
use_2, x0_2, x1_2, y0_2, y1_2, z0_2, z1_2,
use_3, x0_3, x1_3, y0_3, y1_3, z0_3, z1_3,
batch_size: int = 1,
):
boxes = [
(use_1, x0_1, x1_1, y0_1, y1_1, z0_1, z1_1),
(use_2, x0_2, x1_2, y0_2, y1_2, z0_2, z1_2),
(use_3, x0_3, x1_3, y0_3, y1_3, z0_3, z1_3),
]
edit_mask = torch.zeros((batch_size, RESOLUTION, RESOLUTION, RESOLUTION), dtype=torch.bool)
any_box = False
for use, x0, x1, y0, y1, z0, z1 in boxes:
if not use:
continue
ranges = [int(x0), int(x1), int(y0), int(y1), int(z0), int(z1)]
x0, x1, y0, y1, z0, z1 = [max(0, min(RESOLUTION, v)) for v in ranges]
if x0 >= x1 or y0 >= y1 or z0 >= z1:
continue
edit_mask[:, x0:x1, y0:y1, z0:z1] = True
any_box = True
if not any_box:
raise gr.Error("Enable at least one valid edit-mask box.")
return edit_mask
def mask_inputs():
return [
box1_use, box1_x0, box1_x1, box1_y0, box1_y1, box1_z0, box1_z1,
box2_use, box2_x0, box2_x1, box2_y0, box2_y1, box2_z0, box2_z1,
box3_use, box3_x0, box3_x1, box3_y0, box3_y1, box3_z0, box3_z1,
]
@spaces.GPU(duration=600)
def generate_voxels(
image: Image.Image,
seed: int,
randomize_seed: bool,
preprocess_image: bool,
dvd_steps: int,
dvd_cfg_mode: str,
dvd_cfg_constant: float,
dvd_cfg_early: float,
dvd_cfg_late: float,
dvd_cfg_split: float,
progress=gr.Progress(track_tqdm=True),
):
progress(0.02, desc="Starting ZeroGPU callback")
log_event(f"generate_voxels start seed={seed} randomize={randomize_seed} steps={dvd_steps}")
if image is None:
log_event("generate_voxels missing image")
raise gr.Error("Please provide a generation image.")
seed = get_seed(randomize_seed, seed)
progress(0.08, desc="Loading DVD generation model")
pipeline = ensure_dvd_gen_pipeline("cuda")
progress(0.18, desc="Sampling voxels")
log_event(f"generate_voxels sampling seed={seed}")
voxels = pipeline.sample_voxels(
image,
seed=seed,
preprocess_image=preprocess_image,
**dvd_sampler_kwargs(dvd_steps, dvd_cfg_mode, dvd_cfg_constant, dvd_cfg_early, dvd_cfg_late, dvd_cfg_split),
)
progress(0.88, desc="Exporting voxel preview")
log_event("generate_voxels sampled; exporting voxel mesh")
mesh_path = voxel_to_mesh(voxels, worker_path("generated_voxels.glb"))
npy_path = save_voxel_coords(voxels, worker_path("generated_voxel64_coords.npy"))
torch.cuda.empty_cache()
log_event(f"generate_voxels done seed={seed} npy={npy_path}")
return voxel_state(voxels), mesh_path, npy_path, seed
@spaces.GPU(duration=600)
def generation_stage2(
image: Image.Image,
voxels,
seed: int,
randomize_seed: bool,
preprocess_image: bool,
slat_steps: int,
slat_cfg_strength: float,
progress=gr.Progress(track_tqdm=True),
):
progress(0.02, desc="Starting ZeroGPU callback")
log_event(f"generation_stage2 start seed={seed} randomize={randomize_seed} steps={slat_steps}")
if image is None:
log_event("generation_stage2 missing image")
raise gr.Error("Please provide the same generation image for TRELLIS stage 2.")
if voxels is None:
raise gr.Error("Generate voxels before running TRELLIS stage 2.")
seed = get_seed(randomize_seed, seed)
log_event(f"generation_stage2 loading voxel state type={type(voxels).__name__}")
voxels = load_voxel_file(voxels)
log_event(f"generation_stage2 voxel state ready shape={tuple(voxels.samples.shape)}")
progress(0.08, desc="Checking CUDA render extensions")
ensure_dvd_imports()
ensure_zero_gpu_extensions()
log_event("generation_stage2 CUDA render extensions ready")
progress(0.16, desc="Loading TRELLIS stage 2")
log_event("generation_stage2 running TRELLIS")
pipeline = ensure_trellis_pipeline()
log_event("generation_stage2 TRELLIS pipeline ready")
outputs = run_image_stage2_from_dvd_voxels(
pipeline,
image,
voxels,
seed=seed,
formats=["gaussian", "mesh"],
preprocess_image=preprocess_image,
slat_sampler_params=slat_sampler_params(slat_steps, slat_cfg_strength),
)
glb = get_postprocessing_utils().to_glb(outputs["gaussian"][0], outputs["mesh"][0])
glb_path = worker_path("generated_stage2.glb")
glb.export(glb_path)
torch.cuda.empty_cache()
log_event(f"generation_stage2 done seed={seed} glb={glb_path}")
return glb_path, glb_path, seed
def transfer_generation_to_editing(voxels, mesh_path):
if voxels is None:
raise gr.Error("No generated voxels to transfer.")
return voxels, mesh_path
@spaces.GPU(duration=60)
def load_edit_voxels(file, voxel_path: str, req: gr.Request):
voxel_source = file if file is not None else voxel_path
if voxel_source in (None, ""):
raise gr.Error("Upload a voxel file or select a preloaded edit voxel.")
voxels = load_voxel_file(voxel_source)
mesh_path = voxel_to_mesh(voxels, session_path(req, "loaded_edit_voxels.glb"))
return voxel_state(voxels), mesh_path
@spaces.GPU(duration=60)
def visualize_edit_mask(
voxels,
use_1, x0_1, x1_1, y0_1, y1_1, z0_1, z1_1,
use_2, x0_2, x1_2, y0_2, y1_2, z0_2, z1_2,
use_3, x0_3, x1_3, y0_3, y1_3, z0_3, z1_3,
req: gr.Request,
):
if voxels is None:
raise gr.Error("No editing voxels available.")
output = load_voxel_file(voxels)
edit_mask = build_edit_mask(
use_1, x0_1, x1_1, y0_1, y1_1, z0_1, z1_1,
use_2, x0_2, x1_2, y0_2, y1_2, z0_2, z1_2,
use_3, x0_3, x1_3, y0_3, y1_3, z0_3, z1_3,
batch_size=output.samples.shape[0],
)
perturbed = output.samples.clone()
perturbed[edit_mask] = torch.randint(0, 2, perturbed[edit_mask].shape, dtype=perturbed.dtype)
mask_preview = as_voxel_output(perturbed, resolution=RESOLUTION)
mesh_path = voxel_to_mesh(mask_preview, session_path(req, "edit_mask_preview.glb"))
return mesh_path
@spaces.GPU(duration=600)
def run_editing(
target_image: Image.Image,
voxels,
seed: int,
randomize_seed: bool,
preprocess_image: bool,
dvd_steps: int,
dvd_cfg_mode: str,
dvd_cfg_constant: float,
dvd_cfg_early: float,
dvd_cfg_late: float,
dvd_cfg_split: float,
use_1, x0_1, x1_1, y0_1, y1_1, z0_1, z1_1,
use_2, x0_2, x1_2, y0_2, y1_2, z0_2, z1_2,
use_3, x0_3, x1_3, y0_3, y1_3, z0_3, z1_3,
progress=gr.Progress(track_tqdm=True),
):
progress(0.02, desc="Starting ZeroGPU callback")
log_event(f"run_editing start seed={seed} randomize={randomize_seed} steps={dvd_steps}")
if target_image is None:
log_event("run_editing missing target image")
raise gr.Error("Please provide a target image for editing.")
if voxels is None:
raise gr.Error("Upload voxels or transfer generated voxels first.")
seed = get_seed(randomize_seed, seed)
output = load_voxel_file(voxels)
edit_mask = build_edit_mask(
use_1, x0_1, x1_1, y0_1, y1_1, z0_1, z1_1,
use_2, x0_2, x1_2, y0_2, y1_2, z0_2, z1_2,
use_3, x0_3, x1_3, y0_3, y1_3, z0_3, z1_3,
batch_size=output.samples.shape[0],
)
keep_mask = ~edit_mask
progress(0.08, desc="Loading DVD editing model")
pipeline = ensure_dvd_edit_pipeline("cuda")
progress(0.18, desc="Sampling edited voxels")
log_event("run_editing sampling")
edited = pipeline.edit_voxels(
target_image,
output,
keep_mask=keep_mask,
seed=seed,
preprocess_image=preprocess_image,
**dvd_sampler_kwargs(dvd_steps, dvd_cfg_mode, dvd_cfg_constant, dvd_cfg_early, dvd_cfg_late, dvd_cfg_split),
)
mesh_path = voxel_to_mesh(edited, worker_path("edited_voxels.glb"))
npy_path = save_voxel_coords(edited, worker_path("edited_voxel64_coords.npy"))
torch.cuda.empty_cache()
log_event(f"run_editing done seed={seed} npy={npy_path}")
return voxel_state(edited), mesh_path, npy_path, seed
@spaces.GPU(duration=600)
def editing_stage2(
target_image: Image.Image,
edited_voxels,
seed: int,
randomize_seed: bool,
preprocess_image: bool,
slat_steps: int,
slat_cfg_strength: float,
progress=gr.Progress(track_tqdm=True),
):
progress(0.02, desc="Starting ZeroGPU callback")
log_event(f"editing_stage2 start seed={seed} randomize={randomize_seed} steps={slat_steps}")
if target_image is None:
log_event("editing_stage2 missing target image")
raise gr.Error("Please provide the target image for TRELLIS stage 2.")
if edited_voxels is None:
raise gr.Error("Run editing before TRELLIS stage 2.")
seed = get_seed(randomize_seed, seed)
log_event(f"editing_stage2 loading voxel state type={type(edited_voxels).__name__}")
edited_voxels = load_voxel_file(edited_voxels)
log_event(f"editing_stage2 voxel state ready shape={tuple(edited_voxels.samples.shape)}")
progress(0.08, desc="Checking CUDA render extensions")
ensure_dvd_imports()
ensure_zero_gpu_extensions()
log_event("editing_stage2 CUDA render extensions ready")
progress(0.16, desc="Loading TRELLIS stage 2")
pipeline = ensure_trellis_pipeline()
log_event("editing_stage2 TRELLIS pipeline ready")
outputs = run_image_stage2_from_dvd_voxels(
pipeline,
target_image,
edited_voxels,
seed=seed,
formats=["gaussian", "mesh"],
preprocess_image=preprocess_image,
slat_sampler_params=slat_sampler_params(slat_steps, slat_cfg_strength),
)
glb = get_postprocessing_utils().to_glb(outputs["gaussian"][0], outputs["mesh"][0])
glb_path = worker_path("edited_stage2.glb")
glb.export(glb_path)
torch.cuda.empty_cache()
return glb_path, glb_path, seed
APP_CSS = """
#editing-three-col {
align-items: flex-start;
}
#editing-three-col > div {
min-width: 220px !important;
}
"""
with gr.Blocks(
delete_cache=(600, 600),
title="DVD + TRELLIS Voxel Generation and Editing",
css=APP_CSS,
fill_width=True,
) as demo:
gr.Markdown(
"""
## DVD Voxel Generation and Editing
DVD generates or edits a 64^3 voxel structure first. TRELLIS stage 2 is run only when you click the stage-2 button.
"""
)
generated_voxels_state = gr.State()
edit_voxels_state = gr.State()
edited_voxels_state = gr.State()
with gr.Tab("Generation"):
with gr.Row():
with gr.Column():
gen_example = gr.Dropdown(
choices=GENERATION_IMAGE_EXAMPLES,
label="Preloaded Generation Images",
value=None,
interactive=True,
)
gen_image = gr.Image(label="Condition Image", format="png", image_mode="RGBA", type="pil", height=300)
with gr.Accordion("Generation Settings", open=False):
gen_seed = gr.Slider(0, MAX_SEED, value=0, step=1, label="Seed")
gen_randomize = gr.Checkbox(value=True, label="Randomize seed")
gen_preprocess = gr.Checkbox(value=False, label="DVD preprocess image")
gen_dvd_steps = gr.Slider(1, 512, value=256, step=1, label="DVD voxel steps")
gen_dvd_cfg_mode = gr.Radio(
["Default schedule", "Constant", "Two-stage"],
value="Default schedule",
label="DVD voxel CFG mode",
)
gen_dvd_cfg_constant = gr.Slider(
0.0,
5.0,
value=0.7,
step=0.05,
label="DVD voxel constant CFG",
)
with gr.Row():
gen_dvd_cfg_early = gr.Slider(0.0, 5.0, value=0.4, step=0.05, label="DVD CFG early t<0.5")
gen_dvd_cfg_late = gr.Slider(0.0, 5.0, value=0.7, step=0.05, label="DVD CFG late")
gen_dvd_cfg_split = gr.Slider(0.0, 1.0, value=0.5, step=0.05, label="DVD CFG switch time")
gen_stage2_preprocess = gr.Checkbox(value=True, label="TRELLIS preprocess image for stage 2")
gen_slat_steps = gr.Slider(1, 50, value=25, step=1, label="TRELLIS stage-2 steps")
gen_slat_cfg = gr.Slider(0.0, 10.0, value=5.0, step=0.1, label="TRELLIS stage-2 CFG")
gen_btn = gr.Button("1. Generate DVD Voxels")
gen_stage2_btn = gr.Button("2. Run TRELLIS Stage 2", interactive=True)
transfer_btn = gr.Button("Move Generated Voxels To Editing")
with gr.Column():
gen_voxel_view = voxel_viewer("Generated / Cubified Voxels", exposure=5.0, height=320)
gen_npy_download = gr.DownloadButton(label="Download Voxel Coords (.npy)", interactive=False)
gen_stage2_view = gr.Model3D(label="TRELLIS Stage 2 GLB", height=320)
gen_glb_download = gr.DownloadButton(label="Download Stage 2 GLB", interactive=False)
with gr.Tab("Editing"):
with gr.Row(equal_height=False, elem_id="editing-three-col"):
with gr.Column(scale=1, min_width=220):
gr.Markdown("### 1. Source Voxels")
edit_image_example = gr.Dropdown(
choices=EDIT_IMAGE_EXAMPLES,
label="Preloaded Edit Target Images",
value=None,
interactive=True,
)
edit_target_image = gr.Image(label="Target Image", format="png", image_mode="RGBA", type="pil", height=260)
edit_voxel_example = gr.Dropdown(
choices=EDIT_VOXEL_EXAMPLES,
label="Preloaded Edit Voxels",
value=None,
interactive=True,
)
edit_file = gr.File(label="Upload Voxel Coords (.npy, .npz, .pt, .pth)")
load_edit_btn = gr.Button("Load Selected / Uploaded Voxels")
with gr.Row():
rot_x_btn = gr.Button("Rotate X 90")
rot_y_btn = gr.Button("Rotate Y 90")
rot_z_btn = gr.Button("Rotate Z 90")
edit_voxel_view = voxel_viewer("Current Editing Voxels", exposure=10.0, height=300)
with gr.Column(scale=1, min_width=220):
gr.Markdown("### Edit Mask Boxes\nEach enabled box is an edit region. The union of enabled boxes is regenerated.")
with gr.Accordion("Box 1", open=True):
box1_use = gr.Checkbox(value=True, label="Use box 1")
with gr.Row():
box1_x0 = gr.Slider(0, RESOLUTION, value=0, step=1, label="x0")
box1_x1 = gr.Slider(0, RESOLUTION, value=RESOLUTION, step=1, label="x1")
with gr.Row():
box1_y0 = gr.Slider(0, RESOLUTION, value=0, step=1, label="y0")
box1_y1 = gr.Slider(0, RESOLUTION, value=RESOLUTION, step=1, label="y1")
with gr.Row():
box1_z0 = gr.Slider(0, RESOLUTION, value=32, step=1, label="z0")
box1_z1 = gr.Slider(0, RESOLUTION, value=RESOLUTION, step=1, label="z1")
with gr.Accordion("Box 2", open=False):
box2_use = gr.Checkbox(value=False, label="Use box 2")
with gr.Row():
box2_x0 = gr.Slider(0, RESOLUTION, value=0, step=1, label="x0")
box2_x1 = gr.Slider(0, RESOLUTION, value=RESOLUTION, step=1, label="x1")
with gr.Row():
box2_y0 = gr.Slider(0, RESOLUTION, value=0, step=1, label="y0")
box2_y1 = gr.Slider(0, RESOLUTION, value=RESOLUTION, step=1, label="y1")
with gr.Row():
box2_z0 = gr.Slider(0, RESOLUTION, value=0, step=1, label="z0")
box2_z1 = gr.Slider(0, RESOLUTION, value=16, step=1, label="z1")
with gr.Accordion("Box 3", open=False):
box3_use = gr.Checkbox(value=False, label="Use box 3")
with gr.Row():
box3_x0 = gr.Slider(0, RESOLUTION, value=0, step=1, label="x0")
box3_x1 = gr.Slider(0, RESOLUTION, value=RESOLUTION, step=1, label="x1")
with gr.Row():
box3_y0 = gr.Slider(0, RESOLUTION, value=0, step=1, label="y0")
box3_y1 = gr.Slider(0, RESOLUTION, value=RESOLUTION, step=1, label="y1")
with gr.Row():
box3_z0 = gr.Slider(0, RESOLUTION, value=16, step=1, label="z0")
box3_z1 = gr.Slider(0, RESOLUTION, value=32, step=1, label="z1")
preview_mask_btn = gr.Button("Preview Edit Region By Perturbing")
edit_mask_view = voxel_viewer("Edit Region Preview", exposure=5.0, height=300)
with gr.Column(scale=1, min_width=220):
gr.Markdown("### 3. Edited Result")
with gr.Accordion("Editing Settings", open=False):
edit_seed = gr.Slider(0, MAX_SEED, value=0, step=1, label="Seed")
edit_randomize = gr.Checkbox(value=True, label="Randomize seed")
edit_preprocess = gr.Checkbox(value=True, label="DVD preprocess target image")
edit_dvd_steps = gr.Slider(1, 512, value=128, step=1, label="DVD edit steps")
edit_dvd_cfg_mode = gr.Radio(
["Default schedule", "Constant", "Two-stage"],
value="Default schedule",
label="DVD edit CFG mode",
)
edit_dvd_cfg_constant = gr.Slider(
0.0,
5.0,
value=0.45,
step=0.05,
label="DVD edit constant CFG",
)
with gr.Row():
edit_dvd_cfg_early = gr.Slider(0.0, 5.0, value=0.45, step=0.05, label="DVD CFG early t<0.5")
edit_dvd_cfg_late = gr.Slider(0.0, 5.0, value=0.45, step=0.05, label="DVD CFG late")
edit_dvd_cfg_split = gr.Slider(0.0, 1.0, value=0.5, step=0.05, label="DVD CFG switch time")
edit_stage2_preprocess = gr.Checkbox(value=True, label="TRELLIS preprocess target image for stage 2")
edit_slat_steps = gr.Slider(1, 50, value=25, step=1, label="TRELLIS stage-2 steps")
edit_slat_cfg = gr.Slider(0.0, 10.0, value=5.0, step=0.1, label="TRELLIS stage-2 CFG")
edit_btn = gr.Button("Run DVD Editing")
edit_stage2_btn = gr.Button("Run TRELLIS Stage 2")
edited_voxel_view = voxel_viewer("Edited / Cubified Voxels", exposure=5.0, height=300)
edited_npy_download = gr.DownloadButton(label="Download Edited Voxel Coords (.npy)", interactive=False)
edit_stage2_view = gr.Model3D(label="Edited TRELLIS Stage 2 GLB", height=300)
edited_glb_download = gr.DownloadButton(label="Download Edited Stage 2 GLB", interactive=False)
demo.load(start_session)
demo.unload(end_session)
gen_example.change(load_generation_example, inputs=[gen_example], outputs=[gen_image])
edit_image_example.change(load_edit_image_example, inputs=[edit_image_example], outputs=[edit_target_image])
edit_voxel_example.change(
load_edit_voxel_example,
inputs=[edit_voxel_example],
outputs=[edit_voxels_state, edit_voxel_view],
)
gen_btn.click(
generate_voxels,
inputs=[
gen_image,
gen_seed,
gen_randomize,
gen_preprocess,
gen_dvd_steps,
gen_dvd_cfg_mode,
gen_dvd_cfg_constant,
gen_dvd_cfg_early,
gen_dvd_cfg_late,
gen_dvd_cfg_split,
],
outputs=[generated_voxels_state, gen_voxel_view, gen_npy_download, gen_seed],
).then(lambda: gr.DownloadButton(interactive=True), outputs=[gen_npy_download])
gen_stage2_btn.click(
generation_stage2,
inputs=[
gen_image,
generated_voxels_state,
gen_seed,
gen_randomize,
gen_stage2_preprocess,
gen_slat_steps,
gen_slat_cfg,
],
outputs=[gen_stage2_view, gen_glb_download, gen_seed],
).then(lambda: gr.DownloadButton(interactive=True), outputs=[gen_glb_download])
transfer_btn.click(
transfer_generation_to_editing,
inputs=[generated_voxels_state, gen_voxel_view],
outputs=[edit_voxels_state, edit_voxel_view],
)
load_edit_btn.click(
load_edit_voxels,
inputs=[edit_file, edit_voxel_example],
outputs=[edit_voxels_state, edit_voxel_view],
)
rot_x_btn.click(rotate_x, inputs=[edit_voxels_state], outputs=[edit_voxels_state, edit_voxel_view])
rot_y_btn.click(rotate_y, inputs=[edit_voxels_state], outputs=[edit_voxels_state, edit_voxel_view])
rot_z_btn.click(rotate_z, inputs=[edit_voxels_state], outputs=[edit_voxels_state, edit_voxel_view])
preview_mask_btn.click(
visualize_edit_mask,
inputs=[edit_voxels_state] + mask_inputs(),
outputs=[edit_mask_view],
)
edit_btn.click(
run_editing,
inputs=[
edit_target_image,
edit_voxels_state,
edit_seed,
edit_randomize,
edit_preprocess,
edit_dvd_steps,
edit_dvd_cfg_mode,
edit_dvd_cfg_constant,
edit_dvd_cfg_early,
edit_dvd_cfg_late,
edit_dvd_cfg_split,
] + mask_inputs(),
outputs=[edited_voxels_state, edited_voxel_view, edited_npy_download, edit_seed],
).then(lambda: gr.DownloadButton(interactive=True), outputs=[edited_npy_download])
edit_stage2_btn.click(
editing_stage2,
inputs=[
edit_target_image,
edited_voxels_state,
edit_seed,
edit_randomize,
edit_stage2_preprocess,
edit_slat_steps,
edit_slat_cfg,
],
outputs=[edit_stage2_view, edited_glb_download, edit_seed],
).then(lambda: gr.DownloadButton(interactive=True), outputs=[edited_glb_download])
def parse_args():
parser = argparse.ArgumentParser()
parser.add_argument("--device", default="auto", help="cuda, cpu, or auto")
parser.add_argument("--share", action="store_true")
parser.add_argument("--server-name", default=os.environ.get("GRADIO_SERVER_NAME", "0.0.0.0"))
parser.add_argument("--server-port", type=int, default=default_server_port())
return parser.parse_args()
if __name__ == "__main__":
args = parse_args()
# Keep startup free of DVD/TRELLIS imports; only prefetch raw assets.
prefetch_assets()
preload_zero_gpu_models(os.environ.get("DVD_SPACE_DEVICE", "cuda"))
demo.queue().launch(
share=args.share,
server_name=args.server_name,
server_port=args.server_port,
show_api=False,
show_error=True,
ssr_mode=False,
)