dvd-image / app_dvd_text.py
Zhengrui's picture
Deploy isolated DVD image Space
31f6f71 verified
Raw
History Blame
37.5 kB
import argparse
import os
import shutil
import uuid
from pathlib import Path
os.environ.setdefault("SPCONV_ALGO", "native")
os.environ.setdefault("ATTN_BACKEND", "flash_attn")
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
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
try:
import spaces
except ImportError:
class _SpacesFallback:
@staticmethod
def GPU(duration=180):
return lambda fn: fn
spaces = _SpacesFallback()
import gradio as gr
import gradio_client.utils as gradio_client_utils
import numpy as np
import torch
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_text_app"
TMP_DIR.mkdir(parents=True, exist_ok=True)
GEN_DVD_CONFIG = os.environ.get("DVD_TEXT_GEN_CONFIG", "ckpts/dvd_text.json")
GEN_DVD_CKPT = os.environ.get("DVD_TEXT_GEN_CKPT", "ckpts/dvd_text.safetensors")
EDIT_DVD_CONFIG = os.environ.get("DVD_TEXT_EDIT_CONFIG", "ckpts/dvd_text_BSP_ft.json")
EDIT_DVD_CKPT = os.environ.get("DVD_TEXT_EDIT_CKPT", "ckpts/dvd_text_BSP_ft.safetensors")
TRELLIS_TEXT_MODEL = os.environ.get("TRELLIS_TEXT_MODEL", "microsoft/TRELLIS-text-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
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"))
TEXT_PROMPTS = [
"A small helper robot with a round body, sky-blue paint, screen face, and tiny tool arms.",
"A cottage with a large wizard hat roof.",
"A mushroom house with a tiny round door and warm windows.",
"A tiny steampunk locomotive with brass pipes and a chimney.",
]
dvd_gen_pipeline = None
dvd_edit_pipeline = None
trellis_pipeline = None
DVDTextToVoxelPipeline = None
TrellisTextTo3DPipeline = None
as_voxel_output = None
export_cubified_voxels = None
run_text_stage2_from_dvd_voxels = None
_postprocessing_utils = None
def log_event(message: str):
print(f"[DVD Space] {message}", flush=True)
def ensure_dvd_imports():
global DVDTextToVoxelPipeline
global TrellisTextTo3DPipeline
global as_voxel_output
global export_cubified_voxels
global run_text_stage2_from_dvd_voxels
if DVDTextToVoxelPipeline is not None:
return
log_event("importing DVD text/TRELLIS modules")
from dvd import (
DVDTextToVoxelPipeline as _DVDTextToVoxelPipeline,
TrellisTextTo3DPipeline as _TrellisTextTo3DPipeline,
as_voxel_output as _as_voxel_output,
export_cubified_voxels as _export_cubified_voxels,
run_text_stage2_from_dvd_voxels as _run_text_stage2_from_dvd_voxels,
)
DVDTextToVoxelPipeline = _DVDTextToVoxelPipeline
TrellisTextTo3DPipeline = _TrellisTextTo3DPipeline
as_voxel_output = _as_voxel_output
export_cubified_voxels = _export_cubified_voxels
run_text_stage2_from_dvd_voxels = _run_text_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("_", " ")
EDIT_VOXEL_EXAMPLES = [
(asset_label(path), str(path)) for path in list_asset_files("assets/example_voxel_edit", {".npy", ".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 DVDTextToVoxelPipeline.from_pretrained(DVD_MODEL_REPO, variant=variant, **common_kwargs)
if variant == "base":
return DVDTextToVoxelPipeline.from_files(GEN_DVD_CONFIG, GEN_DVD_CKPT, device=device)
if variant == "bsp":
return DVDTextToVoxelPipeline.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 ensure_dvd_gen_pipeline(device: str | None = None):
global dvd_gen_pipeline
device = device or os.environ.get("DVD_SPACE_DEVICE", "cuda")
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")
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")
if trellis_pipeline is None:
log_event(f"loading TRELLIS stage2 pipeline on {device}")
trellis_pipeline = TrellisTextTo3DPipeline.from_pretrained(TRELLIS_TEXT_MODEL)
trellis_pipeline.to(device)
log_event("TRELLIS stage2 pipeline ready")
else:
ensure_pipeline_device(trellis_pipeline, device, "TRELLIS stage2")
return trellis_pipeline
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)
@spaces.GPU(duration=30)
def zero_gpu_smoke_test():
log_event("zero_gpu_smoke_test start")
if not torch.cuda.is_available():
log_event("zero_gpu_smoke_test no cuda")
return "CUDA unavailable inside ZeroGPU worker"
value = torch.ones((1,), device="cuda").sum().item()
name = torch.cuda.get_device_name(0)
log_event(f"zero_gpu_smoke_test done device={name} value={value}")
return f"OK: {name}, value={value}"
def get_seed(randomize_seed: bool, seed: int) -> int:
return int(np.random.randint(0, MAX_SEED)) if randomize_seed else int(seed)
def load_prompt_example(prompt: str):
if not prompt:
raise gr.Error("Please select a preloaded prompt.")
return prompt
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_state(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_state(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 load_voxel_file(file, resolution: int = RESOLUTION):
if file is None:
raise gr.Error("Please upload a voxel coordinate file or select a preloaded edit voxel.")
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_state(file, resolution=resolution)
elif hasattr(file, "path"):
path = file.path
elif hasattr(file, "name"):
path = file.name
else:
return voxel_state(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_state(data, resolution=resolution)
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=120)
def generate_voxels(
prompt: str,
seed: int,
randomize_seed: 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 not prompt:
log_event("generate_voxels missing prompt")
raise gr.Error("Please provide a generation prompt.")
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(
prompt,
seed=seed,
**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_text_voxels.glb"))
npy_path = save_voxel_coords(voxels, worker_path("generated_text_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(
prompt: str,
voxels,
seed: int,
randomize_seed: 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 not prompt:
log_event("generation_stage2 missing prompt")
raise gr.Error("Please provide the same generation prompt for TRELLIS stage 2.")
if voxels is None:
raise gr.Error("Generate voxels before running TRELLIS stage 2.")
seed = get_seed(randomize_seed, seed)
voxels = load_voxel_file(voxels)
progress(0.08, desc="Checking CUDA render extensions")
ensure_dvd_imports()
ensure_zero_gpu_extensions()
progress(0.16, desc="Loading TRELLIS stage 2")
log_event("generation_stage2 running TRELLIS")
outputs = run_text_stage2_from_dvd_voxels(
ensure_trellis_pipeline(),
prompt,
voxels,
seed=seed,
formats=["gaussian", "mesh"],
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_text_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, prompt):
if voxels is None:
raise gr.Error("No generated voxels to transfer.")
return voxels, mesh_path, prompt
@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_text_edit_voxels.glb"))
return voxel_state(voxels), mesh_path
@spaces.GPU(duration=60)
def load_edit_voxel_example(voxel_path: str, req: gr.Request):
return load_edit_voxels(None, voxel_path, req)
@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, "text_edit_mask_preview.glb"))
return mesh_path
@spaces.GPU(duration=120)
def run_editing(
prompt: str,
voxels,
seed: int,
randomize_seed: 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 not prompt:
log_event("run_editing missing prompt")
raise gr.Error("Please provide a target prompt 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(
prompt,
output,
keep_mask=keep_mask,
seed=seed,
**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_text_voxels.glb"))
npy_path = save_voxel_coords(edited, worker_path("edited_text_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(
prompt: str,
edited_voxels,
seed: int,
randomize_seed: 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 not prompt:
log_event("editing_stage2 missing prompt")
raise gr.Error("Please provide the target prompt for TRELLIS stage 2.")
if edited_voxels is None:
raise gr.Error("Run editing before TRELLIS stage 2.")
seed = get_seed(randomize_seed, seed)
edited_voxels = load_voxel_file(edited_voxels)
progress(0.08, desc="Checking CUDA render extensions")
ensure_dvd_imports()
ensure_zero_gpu_extensions()
progress(0.16, desc="Loading TRELLIS stage 2")
outputs = run_text_stage2_from_dvd_voxels(
ensure_trellis_pipeline(),
prompt,
edited_voxels,
seed=seed,
formats=["gaussian", "mesh"],
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_text_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 Text Voxel Generation and Editing",
css=APP_CSS,
fill_width=True,
) as demo:
gr.Markdown(
"""
## DVD Text Voxel Generation and Editing
Text prompts condition DVD voxel generation/editing first. TRELLIS text stage 2 runs only when you click the stage-2 button.
"""
)
with gr.Row():
zero_gpu_smoke_btn = gr.Button("ZeroGPU Smoke Test")
zero_gpu_smoke_out = gr.Textbox(label="ZeroGPU Status", interactive=False)
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_prompt_example = gr.Dropdown(
choices=TEXT_PROMPTS,
label="Preloaded Generation Prompts",
value=None,
interactive=True,
)
gen_prompt = gr.Textbox(
label="Condition Prompt",
value=TEXT_PROMPTS[0],
lines=4,
)
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_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=1.0, 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=1.2, step=0.05, label="DVD CFG late")
gen_dvd_cfg_split = gr.Slider(0.0, 1.0, value=0.8, step=0.05, label="DVD CFG switch time")
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 Text Stage 2")
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 Text 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_prompt_example = gr.Dropdown(
choices=TEXT_PROMPTS,
label="Preloaded Edit Prompts",
value=None,
interactive=True,
)
edit_prompt = gr.Textbox(
label="Target Edit Prompt",
value="A cottage with a large wizard hat roof.",
lines=4,
)
edit_voxel_example = gr.Dropdown(
choices=EDIT_VOXEL_EXAMPLES,
label="Preloaded Edit Voxels",
value=EDIT_VOXEL_EXAMPLES[0][1] if EDIT_VOXEL_EXAMPLES else 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=5.0, height=300)
with gr.Column(scale=1, min_width=220):
gr.Markdown("### 2. 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=28, 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=0, 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=16, 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=RESOLUTION, 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=16, step=1, label="y0")
box3_y1 = gr.Slider(0, RESOLUTION, value=32, step=1, label="y1")
with gr.Row():
box3_z0 = gr.Slider(0, RESOLUTION, value=0, step=1, label="z0")
box3_z1 = gr.Slider(0, RESOLUTION, value=RESOLUTION, 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_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_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 Text Editing")
edit_stage2_btn = gr.Button("Run TRELLIS Text 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 Text 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)
zero_gpu_smoke_btn.click(zero_gpu_smoke_test, outputs=[zero_gpu_smoke_out])
gen_prompt_example.change(load_prompt_example, inputs=[gen_prompt_example], outputs=[gen_prompt])
edit_prompt_example.change(load_prompt_example, inputs=[edit_prompt_example], outputs=[edit_prompt])
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_prompt,
gen_seed,
gen_randomize,
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_prompt,
generated_voxels_state,
gen_seed,
gen_randomize,
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, gen_prompt],
outputs=[edit_voxels_state, edit_voxel_view, edit_prompt],
)
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_prompt,
edit_voxels_state,
edit_seed,
edit_randomize,
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_prompt,
edited_voxels_state,
edit_seed,
edit_randomize,
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()
# Space entrypoint: keep startup CPU-only and let @spaces.GPU callbacks
# acquire the ZeroGPU worker directly from this app file.
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,
)