happy0612 commited on
Commit
b42bc05
Β·
verified Β·
1 Parent(s): 2eaf553

Add GeoNeXt-SVD backend source

Browse files
Files changed (3) hide show
  1. backend-svd/Dockerfile +29 -0
  2. backend-svd/README.md +15 -0
  3. backend-svd/app.py +209 -0
backend-svd/Dockerfile ADDED
@@ -0,0 +1,29 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ FROM nvidia/cuda:12.1.1-cudnn8-runtime-ubuntu22.04
2
+
3
+ ENV DEBIAN_FRONTEND=noninteractive \
4
+ PYTHONUNBUFFERED=1 \
5
+ PIP_NO_CACHE_DIR=1 \
6
+ HF_HOME=/data/.huggingface
7
+
8
+ RUN apt-get update && apt-get install -y --no-install-recommends \
9
+ git python3 python3-pip python3-dev libgl1 libglib2.0-0 \
10
+ && rm -rf /var/lib/apt/lists/*
11
+
12
+ WORKDIR /app
13
+
14
+ RUN git clone --depth 1 https://github.com/Creative-Intelligence-Studio/GeoNeXt.git /app/GeoNeXt
15
+
16
+ RUN python3 -m pip install --upgrade pip \
17
+ && python3 -m pip install torch==2.3.1 torchvision==0.18.1 \
18
+ --index-url https://download.pytorch.org/whl/cu121 \
19
+ && python3 -m pip install -r /app/GeoNeXt/requirements-svd.txt \
20
+ && python3 -m pip install \
21
+ gradio==4.44.1 pydantic==2.10.6 fastapi==0.115.6 \
22
+ && python3 -m pip install -e /app/GeoNeXt
23
+
24
+ COPY app.py /app/app.py
25
+
26
+ EXPOSE 7860
27
+
28
+ CMD ["python3", "/app/app.py"]
29
+
backend-svd/README.md ADDED
@@ -0,0 +1,15 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ title: GeoNeXt-SVD
3
+ emoji: 🌐
4
+ colorFrom: indigo
5
+ colorTo: purple
6
+ sdk: docker
7
+ app_port: 7860
8
+ pinned: false
9
+ license: apache-2.0
10
+ ---
11
+
12
+ # GeoNeXt-SVD
13
+
14
+ Official demo for **Video Generative Models as Geometry Learner**.
15
+
backend-svd/app.py ADDED
@@ -0,0 +1,209 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import importlib.util
2
+ import os
3
+ import sys
4
+ import threading
5
+ from contextlib import nullcontext
6
+ from pathlib import Path
7
+
8
+ os.environ["NO_PROXY"] = "localhost,127.0.0.1,0.0.0.0"
9
+ os.environ["no_proxy"] = os.environ["NO_PROXY"]
10
+
11
+ import gradio as gr
12
+ import numpy as np
13
+ import torch
14
+ from diffusers import AutoencoderKL, UNetSpatioTemporalConditionModel
15
+ from huggingface_hub import snapshot_download
16
+ from PIL import Image
17
+
18
+
19
+ ROOT = Path("/app/GeoNeXt")
20
+ if not ROOT.exists():
21
+ ROOT = Path(__file__).resolve().parents[2]
22
+ sys.path.insert(0, str(ROOT))
23
+ sys.path.insert(0, str(ROOT / "GeoNeXt-SVD"))
24
+
25
+
26
+ def load_module(name, path):
27
+ spec = importlib.util.spec_from_file_location(name, path)
28
+ module = importlib.util.module_from_spec(spec)
29
+ spec.loader.exec_module(module)
30
+ return module
31
+
32
+
33
+ svd = load_module("geonext_svd_inference", ROOT / "GeoNeXt-SVD" / "inference.py")
34
+ pipeline_module = load_module("geonext_svd_pipeline", ROOT / "GeoNeXt-SVD" / "pipeline.py")
35
+ GeoNeXtPipeline = pipeline_module.GeoNeXtPipeline
36
+
37
+ device = torch.device("cuda")
38
+ dtype = torch.float16
39
+ checkpoint_root = snapshot_download(
40
+ repo_id="happy0612/GeoNeXt",
41
+ allow_patterns="GeoNeXt-SVD/**",
42
+ )
43
+ checkpoint = str(Path(checkpoint_root) / "GeoNeXt-SVD")
44
+
45
+ vae = AutoencoderKL.from_pretrained(
46
+ "stabilityai/sd-vae-ft-mse",
47
+ torch_dtype=dtype,
48
+ subfolder=None,
49
+ )
50
+ unet = UNetSpatioTemporalConditionModel.from_pretrained(
51
+ checkpoint,
52
+ subfolder="unet",
53
+ torch_dtype=dtype,
54
+ low_cpu_mem_usage=False,
55
+ )
56
+ pipe = GeoNeXtPipeline.from_pretrained(
57
+ svd.SVD_BASE_MODEL,
58
+ unet=unet,
59
+ vae=vae,
60
+ variant="fp16",
61
+ torch_dtype=dtype,
62
+ low_cpu_mem_usage=False,
63
+ ).to(device)
64
+ pipe.set_progress_bar_config(disable=True)
65
+
66
+ inference_lock = threading.Lock()
67
+
68
+
69
+ @torch.inference_mode()
70
+ def predict(image, steps):
71
+ if image is None:
72
+ raise gr.Error("Please upload an image first.")
73
+
74
+ original = Image.fromarray(np.asarray(image, dtype=np.uint8), mode="RGB")
75
+ resized = svd._resize(original, 768, "long")
76
+ width, height = resized.size
77
+ width = max(64, round(width / 64) * 64)
78
+ height = max(64, round(height / 64) * 64)
79
+ resized = resized.resize((width, height), Image.Resampling.BICUBIC)
80
+ generator = torch.Generator(device=device).manual_seed(0)
81
+
82
+ context = torch.autocast("cuda", dtype=dtype)
83
+ with inference_lock, context:
84
+ prediction = pipe(
85
+ resized,
86
+ num_frames=3,
87
+ width=width,
88
+ height=height,
89
+ min_guidance_scale=1.0,
90
+ max_guidance_scale=1.2,
91
+ noise_aug_strength=0.0,
92
+ decode_chunk_size=8,
93
+ generator=generator,
94
+ motion_bucket_id=127,
95
+ fps=7,
96
+ num_inference_steps=int(steps),
97
+ )
98
+
99
+ depth = prediction.geo_res[0].mean(dim=1).squeeze().float().cpu().numpy()
100
+ normal = prediction.geo_res[1].squeeze().permute(1, 2, 0).float().cpu().numpy()
101
+ depth = np.asarray(
102
+ Image.fromarray(depth, mode="F").resize(original.size, Image.Resampling.BILINEAR)
103
+ )
104
+ normal = np.stack(
105
+ [
106
+ np.asarray(
107
+ Image.fromarray(normal[..., channel], mode="F").resize(
108
+ original.size, Image.Resampling.BILINEAR
109
+ )
110
+ )
111
+ for channel in range(3)
112
+ ],
113
+ axis=-1,
114
+ )
115
+
116
+ from utils.visualization import depth_to_vis, normal_to_vis
117
+
118
+ depth_vis = depth_to_vis(np.clip(depth, 0.0, 1.0), reverse_color=True)
119
+ normal_vis = normal_to_vis(np.clip(normal, -1.0, 1.0))
120
+ return np.asarray(depth_vis), np.asarray(normal_vis)
121
+
122
+
123
+ header = """
124
+ <div id="geonext-header">
125
+ <h1>GeoNeXt</h1>
126
+ <h3>Video Generative Models as Geometry Learner</h3>
127
+ <p>Predict monocular depth and surface normals with video generative priors.</p>
128
+ <div class="geonext-links">
129
+ <a href="https://huggingface.co/spaces/happy0612/GeoNeXt">🌐 Unified Demo</a>
130
+ <a href="https://arxiv.org/abs/2608.28549" target="_blank">πŸ“„ Paper</a>
131
+ <a href="https://happy-hsy.github.io/projects/GeoNeXt/" target="_blank">🌐 Project Page</a>
132
+ <a href="https://github.com/Creative-Intelligence-Studio/GeoNeXt" target="_blank">πŸ’» GitHub</a>
133
+ <a href="https://huggingface.co/happy0612/GeoNeXt" target="_blank">πŸ€— Model Checkpoints</a>
134
+ </div>
135
+ </div>
136
+ """
137
+
138
+ css = """
139
+ .gradio-container { max-width: 1180px !important; margin: 0 auto !important; }
140
+ #geonext-header { text-align: center; padding: 1.5rem 0 1.1rem; }
141
+ #geonext-header h1 { font-size: 2.7rem; line-height: 1; margin: 0 0 0.55rem; }
142
+ #geonext-header h3 { font-size: 1.25rem; font-weight: 600; margin: 0 0 0.45rem; }
143
+ #geonext-header p { color: var(--body-text-color-subdued); margin: 0 0 1rem; }
144
+ .geonext-links { display: flex; justify-content: center; flex-wrap: wrap; gap: 0.55rem; }
145
+ .geonext-links a { border: 1px solid var(--border-color-primary); border-radius: 999px; color: var(--body-text-color); padding: 0.4rem 0.8rem; text-decoration: none !important; }
146
+ .geonext-links a:hover { border-color: var(--color-accent); color: var(--color-accent); }
147
+ #run-button { min-height: 46px; font-weight: 700; }
148
+ #input-panel, #output-panel { min-width: 0; }
149
+ """
150
+
151
+ examples = [
152
+ str(path)
153
+ for path in sorted((ROOT / "assets" / "input").glob("*"))
154
+ if path.suffix.lower() in {".jpg", ".jpeg", ".png"}
155
+ ]
156
+
157
+ with gr.Blocks(
158
+ theme=gr.themes.Soft(primary_hue="indigo", neutral_hue="slate"),
159
+ css=css,
160
+ ) as demo:
161
+ gr.HTML(header)
162
+ with gr.Row():
163
+ with gr.Column(scale=5, elem_id="input-panel"):
164
+ input_image = gr.Image(
165
+ type="numpy", image_mode="RGB", label="Input Image", height=520
166
+ )
167
+ with gr.Column(scale=7, elem_id="output-panel"):
168
+ with gr.Tabs():
169
+ with gr.Tab("Depth"):
170
+ depth_output = gr.Image(
171
+ type="numpy", label="Predicted Depth", format="png",
172
+ height=520, interactive=False,
173
+ )
174
+ with gr.Tab("Surface Normal"):
175
+ normal_output = gr.Image(
176
+ type="numpy", label="Predicted Surface Normal", format="png",
177
+ height=520, interactive=False,
178
+ )
179
+
180
+ inference_steps = gr.Slider(
181
+ minimum=1,
182
+ maximum=5,
183
+ value=5,
184
+ step=1,
185
+ label="Inference Steps",
186
+ info="More steps may improve quality but take longer.",
187
+ )
188
+ run_button = gr.Button(
189
+ "Run GeoNeXt-SVD", variant="primary", elem_id="run-button"
190
+ )
191
+
192
+ run_button.click(
193
+ fn=predict,
194
+ inputs=[input_image, inference_steps],
195
+ outputs=[depth_output, normal_output],
196
+ concurrency_limit=1,
197
+ api_name="predict",
198
+ )
199
+ if examples:
200
+ gr.Examples(examples=examples, inputs=input_image, label="Try an example")
201
+
202
+
203
+ if __name__ == "__main__":
204
+ is_space = bool(os.environ.get("SPACE_ID"))
205
+ demo.queue(default_concurrency_limit=1).launch(
206
+ server_name="0.0.0.0" if is_space else "127.0.0.1",
207
+ server_port=int(os.environ.get("PORT", "7860")),
208
+ share=not is_space,
209
+ )