dvd-image / app_dvd_image_trellis_style.py
Zhengrui's picture
Try trellis-style GPU setup before DVD import
bec7c36 verified
Raw
History Blame Contribute Delete
7.12 kB
import os
import random
import uuid
from pathlib import Path
os.environ.setdefault("SPCONV_ALGO", "native")
os.environ.setdefault("ATTN_BACKEND", "flash_attn")
os.environ.setdefault("SPARSE_ATTN_BACKEND", "flash_attn")
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
os.environ.setdefault("DVD_MODEL_REPO", "Zhengrui/dvd")
import spaces
import gradio as gr
import numpy as np
import torch
MAX_SEED = 2**31 - 1
ROOT_DIR = Path(__file__).resolve().parent
TMP_DIR = ROOT_DIR / "tmp" / "dvd_image_trellis_style"
TMP_DIR.mkdir(parents=True, exist_ok=True)
IMAGE_EXTENSIONS = {".png", ".jpg", ".jpeg", ".webp"}
EXAMPLE_DIR = ROOT_DIR / "assets" / "example_image"
EXAMPLES = [
str(path)
for path in sorted(EXAMPLE_DIR.iterdir())
if path.is_file() and path.suffix.lower() in IMAGE_EXTENSIONS
] if EXAMPLE_DIR.exists() else []
def log_event(message: str):
print(f"[DVD TrellisStyle] {message}", flush=True)
@spaces.GPU(duration=60)
def first_gpu_setup():
log_event("first_gpu_setup start")
if not torch.cuda.is_available():
raise RuntimeError("CUDA unavailable during first_gpu_setup")
value = torch.ones((1,), device="cuda").sum().item()
name = torch.cuda.get_device_name(0)
log_event(f"first_gpu_setup ok device={name} value={value}")
# Match trellis-community: acquire a ZeroGPU worker once before importing TRELLIS.
first_gpu_setup()
from dvd import DVDImageToVoxelPipeline, export_cubified_voxels
def worker_path(name: str) -> str:
path = TMP_DIR / f"worker-{uuid.uuid4().hex}"
path.mkdir(parents=True, exist_ok=True)
return str(path / name)
def cfg_schedule(mode: str, constant: float, early: float, late: float, split: float):
if mode == "Constant":
return float(constant)
if mode == "Two-stage":
split = float(split)
early = float(early)
late = float(late)
return lambda t: early if t < split else late
return None
repo = os.environ.get("DVD_MODEL_REPO", "Zhengrui/dvd")
subfolder = os.environ.get("DVD_MODEL_SUBFOLDER") or None
revision = os.environ.get("DVD_MODEL_REVISION") or None
token = os.environ.get("DVD_MODEL_TOKEN") or os.environ.get("HF_TOKEN") or None
log_event(f"loading DVD image pipeline from {repo}")
dvd_pipe = DVDImageToVoxelPipeline.from_pretrained(
repo,
variant="base",
subfolder=subfolder,
revision=revision,
token=token,
)
log_event("moving DVD image pipeline to cuda")
dvd_pipe.to("cuda")
log_event("DVD image pipeline ready on cuda")
@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}"
@spaces.GPU(duration=240)
def generate_voxels(
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.01, desc="Starting ZeroGPU callback")
log_event(f"generate_voxels start seed={seed} randomize={randomize_seed} steps={dvd_steps}")
if image is None:
raise gr.Error("Please provide an image.")
seed = random.randint(0, MAX_SEED) if randomize_seed else int(seed)
sampler_kwargs = {"steps": int(dvd_steps)}
schedule = cfg_schedule(
dvd_cfg_mode,
dvd_cfg_constant,
dvd_cfg_early,
dvd_cfg_late,
dvd_cfg_split,
)
if schedule is not None:
sampler_kwargs["cfg_strength"] = schedule
progress(0.08, desc="Sampling DVD voxels")
voxels = dvd_pipe.sample_voxels(
image,
seed=seed,
preprocess_image=preprocess_image,
**sampler_kwargs,
)
progress(0.88, desc="Exporting voxel preview")
mesh_path = worker_path("generated_voxels.glb")
npy_path = worker_path("generated_voxel64_coords.npy")
export_cubified_voxels(voxels, mesh_path)
np.save(npy_path, voxels.coords_without_batch.detach().cpu().numpy().astype(np.int32))
torch.cuda.empty_cache()
log_event(f"generate_voxels done seed={seed} mesh={mesh_path} npy={npy_path}")
return mesh_path, npy_path, int(seed), f"Done. seed={seed}"
with gr.Blocks(title="DVD Image", fill_width=True) as demo:
gr.Markdown("## DVD Image Voxel Generation")
with gr.Row():
smoke_btn = gr.Button("ZeroGPU Smoke Test")
smoke_out = gr.Textbox(label="ZeroGPU Status", interactive=False)
smoke_btn.click(zero_gpu_smoke_test, outputs=smoke_out)
with gr.Row(equal_height=False):
with gr.Column():
image = gr.Image(label="Input Image", format="png", image_mode="RGBA", type="pil", height=320)
if EXAMPLES:
gr.Examples(examples=EXAMPLES[:12], inputs=image, examples_per_page=6)
with gr.Accordion("DVD Settings", open=False):
seed = gr.Slider(0, MAX_SEED, value=0, step=1, label="Seed")
randomize_seed = gr.Checkbox(value=True, label="Randomize seed")
preprocess_image = gr.Checkbox(value=True, label="DVD preprocess image")
dvd_steps = gr.Slider(1, 512, value=256, step=1, label="DVD steps")
dvd_cfg_mode = gr.Radio(
["Default schedule", "Constant", "Two-stage"],
value="Default schedule",
label="DVD CFG mode",
)
dvd_cfg_constant = gr.Slider(0.0, 5.0, value=0.7, step=0.05, label="Constant CFG")
dvd_cfg_early = gr.Slider(0.0, 5.0, value=0.4, step=0.05, label="Early CFG")
dvd_cfg_late = gr.Slider(0.0, 5.0, value=0.7, step=0.05, label="Late CFG")
dvd_cfg_split = gr.Slider(0.0, 1.0, value=0.5, step=0.05, label="CFG switch time")
gen_btn = gr.Button("Generate DVD Voxels", variant="primary")
with gr.Column():
voxel_view = gr.Model3D(
label="Generated / Cubified Voxels",
height=360,
camera_position=(-180, 90, 3),
)
npy_download = gr.DownloadButton(label="Download Voxel Coords (.npy)", interactive=False)
status = gr.Textbox(label="Status", interactive=False)
gen_btn.click(
generate_voxels,
inputs=[
image,
seed,
randomize_seed,
preprocess_image,
dvd_steps,
dvd_cfg_mode,
dvd_cfg_constant,
dvd_cfg_early,
dvd_cfg_late,
dvd_cfg_split,
],
outputs=[voxel_view, npy_download, seed, status],
).then(lambda: gr.DownloadButton(interactive=True), outputs=[npy_download])
if __name__ == "__main__":
demo.launch(show_api=False, show_error=True, ssr_mode=False)