Zhengrui commited on
Commit
bec7c36
·
verified ·
1 Parent(s): 6c3b5ae

Try trellis-style GPU setup before DVD import

Browse files
Files changed (2) hide show
  1. README.md +1 -1
  2. app_dvd_image_trellis_style.py +202 -0
README.md CHANGED
@@ -5,7 +5,7 @@ colorFrom: blue
5
  colorTo: green
6
  sdk: gradio
7
  sdk_version: 5.34.2
8
- app_file: app_space_image.py
9
  pinned: false
10
  license: mit
11
  models:
 
5
  colorTo: green
6
  sdk: gradio
7
  sdk_version: 5.34.2
8
+ app_file: app_dvd_image_trellis_style.py
9
  pinned: false
10
  license: mit
11
  models:
app_dvd_image_trellis_style.py ADDED
@@ -0,0 +1,202 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import random
3
+ import uuid
4
+ from pathlib import Path
5
+
6
+ os.environ.setdefault("SPCONV_ALGO", "native")
7
+ os.environ.setdefault("ATTN_BACKEND", "flash_attn")
8
+ os.environ.setdefault("SPARSE_ATTN_BACKEND", "flash_attn")
9
+ os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
10
+ os.environ.setdefault("DVD_MODEL_REPO", "Zhengrui/dvd")
11
+
12
+ import spaces
13
+ import gradio as gr
14
+ import numpy as np
15
+ import torch
16
+
17
+
18
+ MAX_SEED = 2**31 - 1
19
+ ROOT_DIR = Path(__file__).resolve().parent
20
+ TMP_DIR = ROOT_DIR / "tmp" / "dvd_image_trellis_style"
21
+ TMP_DIR.mkdir(parents=True, exist_ok=True)
22
+ IMAGE_EXTENSIONS = {".png", ".jpg", ".jpeg", ".webp"}
23
+ EXAMPLE_DIR = ROOT_DIR / "assets" / "example_image"
24
+ EXAMPLES = [
25
+ str(path)
26
+ for path in sorted(EXAMPLE_DIR.iterdir())
27
+ if path.is_file() and path.suffix.lower() in IMAGE_EXTENSIONS
28
+ ] if EXAMPLE_DIR.exists() else []
29
+
30
+
31
+ def log_event(message: str):
32
+ print(f"[DVD TrellisStyle] {message}", flush=True)
33
+
34
+
35
+ @spaces.GPU(duration=60)
36
+ def first_gpu_setup():
37
+ log_event("first_gpu_setup start")
38
+ if not torch.cuda.is_available():
39
+ raise RuntimeError("CUDA unavailable during first_gpu_setup")
40
+ value = torch.ones((1,), device="cuda").sum().item()
41
+ name = torch.cuda.get_device_name(0)
42
+ log_event(f"first_gpu_setup ok device={name} value={value}")
43
+
44
+
45
+ # Match trellis-community: acquire a ZeroGPU worker once before importing TRELLIS.
46
+ first_gpu_setup()
47
+
48
+ from dvd import DVDImageToVoxelPipeline, export_cubified_voxels
49
+
50
+
51
+ def worker_path(name: str) -> str:
52
+ path = TMP_DIR / f"worker-{uuid.uuid4().hex}"
53
+ path.mkdir(parents=True, exist_ok=True)
54
+ return str(path / name)
55
+
56
+
57
+ def cfg_schedule(mode: str, constant: float, early: float, late: float, split: float):
58
+ if mode == "Constant":
59
+ return float(constant)
60
+ if mode == "Two-stage":
61
+ split = float(split)
62
+ early = float(early)
63
+ late = float(late)
64
+ return lambda t: early if t < split else late
65
+ return None
66
+
67
+
68
+ repo = os.environ.get("DVD_MODEL_REPO", "Zhengrui/dvd")
69
+ subfolder = os.environ.get("DVD_MODEL_SUBFOLDER") or None
70
+ revision = os.environ.get("DVD_MODEL_REVISION") or None
71
+ token = os.environ.get("DVD_MODEL_TOKEN") or os.environ.get("HF_TOKEN") or None
72
+
73
+ log_event(f"loading DVD image pipeline from {repo}")
74
+ dvd_pipe = DVDImageToVoxelPipeline.from_pretrained(
75
+ repo,
76
+ variant="base",
77
+ subfolder=subfolder,
78
+ revision=revision,
79
+ token=token,
80
+ )
81
+ log_event("moving DVD image pipeline to cuda")
82
+ dvd_pipe.to("cuda")
83
+ log_event("DVD image pipeline ready on cuda")
84
+
85
+
86
+ @spaces.GPU(duration=30)
87
+ def zero_gpu_smoke_test():
88
+ log_event("zero_gpu_smoke_test start")
89
+ if not torch.cuda.is_available():
90
+ log_event("zero_gpu_smoke_test no cuda")
91
+ return "CUDA unavailable inside ZeroGPU worker"
92
+ value = torch.ones((1,), device="cuda").sum().item()
93
+ name = torch.cuda.get_device_name(0)
94
+ log_event(f"zero_gpu_smoke_test done device={name} value={value}")
95
+ return f"OK: {name}, value={value}"
96
+
97
+
98
+ @spaces.GPU(duration=240)
99
+ def generate_voxels(
100
+ image,
101
+ seed: int,
102
+ randomize_seed: bool,
103
+ preprocess_image: bool,
104
+ dvd_steps: int,
105
+ dvd_cfg_mode: str,
106
+ dvd_cfg_constant: float,
107
+ dvd_cfg_early: float,
108
+ dvd_cfg_late: float,
109
+ dvd_cfg_split: float,
110
+ progress=gr.Progress(track_tqdm=True),
111
+ ):
112
+ progress(0.01, desc="Starting ZeroGPU callback")
113
+ log_event(f"generate_voxels start seed={seed} randomize={randomize_seed} steps={dvd_steps}")
114
+ if image is None:
115
+ raise gr.Error("Please provide an image.")
116
+
117
+ seed = random.randint(0, MAX_SEED) if randomize_seed else int(seed)
118
+ sampler_kwargs = {"steps": int(dvd_steps)}
119
+ schedule = cfg_schedule(
120
+ dvd_cfg_mode,
121
+ dvd_cfg_constant,
122
+ dvd_cfg_early,
123
+ dvd_cfg_late,
124
+ dvd_cfg_split,
125
+ )
126
+ if schedule is not None:
127
+ sampler_kwargs["cfg_strength"] = schedule
128
+
129
+ progress(0.08, desc="Sampling DVD voxels")
130
+ voxels = dvd_pipe.sample_voxels(
131
+ image,
132
+ seed=seed,
133
+ preprocess_image=preprocess_image,
134
+ **sampler_kwargs,
135
+ )
136
+
137
+ progress(0.88, desc="Exporting voxel preview")
138
+ mesh_path = worker_path("generated_voxels.glb")
139
+ npy_path = worker_path("generated_voxel64_coords.npy")
140
+ export_cubified_voxels(voxels, mesh_path)
141
+ np.save(npy_path, voxels.coords_without_batch.detach().cpu().numpy().astype(np.int32))
142
+ torch.cuda.empty_cache()
143
+ log_event(f"generate_voxels done seed={seed} mesh={mesh_path} npy={npy_path}")
144
+ return mesh_path, npy_path, int(seed), f"Done. seed={seed}"
145
+
146
+
147
+ with gr.Blocks(title="DVD Image", fill_width=True) as demo:
148
+ gr.Markdown("## DVD Image Voxel Generation")
149
+ with gr.Row():
150
+ smoke_btn = gr.Button("ZeroGPU Smoke Test")
151
+ smoke_out = gr.Textbox(label="ZeroGPU Status", interactive=False)
152
+ smoke_btn.click(zero_gpu_smoke_test, outputs=smoke_out)
153
+
154
+ with gr.Row(equal_height=False):
155
+ with gr.Column():
156
+ image = gr.Image(label="Input Image", format="png", image_mode="RGBA", type="pil", height=320)
157
+ if EXAMPLES:
158
+ gr.Examples(examples=EXAMPLES[:12], inputs=image, examples_per_page=6)
159
+ with gr.Accordion("DVD Settings", open=False):
160
+ seed = gr.Slider(0, MAX_SEED, value=0, step=1, label="Seed")
161
+ randomize_seed = gr.Checkbox(value=True, label="Randomize seed")
162
+ preprocess_image = gr.Checkbox(value=True, label="DVD preprocess image")
163
+ dvd_steps = gr.Slider(1, 512, value=256, step=1, label="DVD steps")
164
+ dvd_cfg_mode = gr.Radio(
165
+ ["Default schedule", "Constant", "Two-stage"],
166
+ value="Default schedule",
167
+ label="DVD CFG mode",
168
+ )
169
+ dvd_cfg_constant = gr.Slider(0.0, 5.0, value=0.7, step=0.05, label="Constant CFG")
170
+ dvd_cfg_early = gr.Slider(0.0, 5.0, value=0.4, step=0.05, label="Early CFG")
171
+ dvd_cfg_late = gr.Slider(0.0, 5.0, value=0.7, step=0.05, label="Late CFG")
172
+ dvd_cfg_split = gr.Slider(0.0, 1.0, value=0.5, step=0.05, label="CFG switch time")
173
+ gen_btn = gr.Button("Generate DVD Voxels", variant="primary")
174
+ with gr.Column():
175
+ voxel_view = gr.Model3D(
176
+ label="Generated / Cubified Voxels",
177
+ height=360,
178
+ camera_position=(-180, 90, 3),
179
+ )
180
+ npy_download = gr.DownloadButton(label="Download Voxel Coords (.npy)", interactive=False)
181
+ status = gr.Textbox(label="Status", interactive=False)
182
+
183
+ gen_btn.click(
184
+ generate_voxels,
185
+ inputs=[
186
+ image,
187
+ seed,
188
+ randomize_seed,
189
+ preprocess_image,
190
+ dvd_steps,
191
+ dvd_cfg_mode,
192
+ dvd_cfg_constant,
193
+ dvd_cfg_early,
194
+ dvd_cfg_late,
195
+ dvd_cfg_split,
196
+ ],
197
+ outputs=[voxel_view, npy_download, seed, status],
198
+ ).then(lambda: gr.DownloadButton(interactive=True), outputs=[npy_download])
199
+
200
+
201
+ if __name__ == "__main__":
202
+ demo.launch(show_api=False, show_error=True, ssr_mode=False)