Wat3R / app.py
lsxi77777's picture
Update paper arXiv link
1b285b0 verified
Raw
History Blame Contribute Delete
20.6 kB
from __future__ import annotations
import gc
import glob
import os
import re
import shutil
import time
from datetime import datetime
import cv2
import gradio as gr
import numpy as np
import torch
from visual_util import predictions_to_glb
from wat3r.models.wat3r import Wat3R
from wat3r.utils.geometry import unproject_depth_map_to_point_map
from wat3r.utils.load_fn import load_and_preprocess_images
from wat3r.utils.pose_enc import pose_encoding_to_extri_intri
from wat3r.utils.static_mask import build_static_masks, depth_foreground_mask
try:
import spaces
except ImportError:
spaces = None
def gpu(duration: int = 120):
if spaces is None:
def decorator(function):
return function
return decorator
return spaces.GPU(duration=duration)
MODEL_ID = os.environ.get("WAT3R_MODEL_ID", "lsxi77777/Wat3R")
PAPER_URL = os.environ.get("WAT3R_PAPER_URL", "https://arxiv.org/abs/2607.08772")
CODE_URL = os.environ.get("WAT3R_CODE_URL", "https://github.com/LSXI7/Wat3R")
VIDEO_FRAME_INTERVAL_SECONDS = float(os.environ.get("WAT3R_VIDEO_FRAME_INTERVAL", "1.0"))
print(f"Initializing Wat3R model from {MODEL_ID}...")
model = Wat3R.from_pretrained(MODEL_ID).eval()
def _empty_cache():
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
def _device():
return torch.device("cuda" if torch.cuda.is_available() else "cpu")
def _amp_dtype(device: torch.device):
if device.type != "cuda":
return torch.float32
return torch.bfloat16 if torch.cuda.get_device_capability(device)[0] >= 8 else torch.float16
def _to_numpy_without_batch(predictions: dict) -> dict:
output = {}
for key, value in predictions.items():
if isinstance(value, torch.Tensor):
value = value.detach().cpu().numpy()
if value.ndim > 0 and value.shape[0] == 1:
value = value[0]
output[key] = value
return output
@gpu(duration=120)
def run_model(target_dir: str, model: Wat3R) -> dict:
print(f"Processing images from {target_dir}")
device = _device()
model = model.to(device).eval()
image_names = sorted(glob.glob(os.path.join(target_dir, "images", "*")), key=_natural_key)
if len(image_names) == 0:
raise ValueError("No images found. Please upload a video or images first.")
images = load_and_preprocess_images(image_names, mode="max", target_size=518).to(device)
print(f"Preprocessed images shape: {images.shape}")
with torch.no_grad():
with torch.cuda.amp.autocast(dtype=_amp_dtype(device), enabled=device.type == "cuda"):
predictions = model(images)
extrinsic, intrinsic = pose_encoding_to_extri_intri(predictions["pose_enc"], images.shape[-2:])
predictions["extrinsic"] = extrinsic
predictions["intrinsic"] = intrinsic
predictions.pop("pose_enc_list", None)
predictions_np = _to_numpy_without_batch(predictions)
predictions_np["world_points_from_depth"] = unproject_depth_map_to_point_map(
predictions_np["depth"],
predictions_np["extrinsic"],
predictions_np["intrinsic"],
)
_empty_cache()
return predictions_np
def handle_uploads(input_video, input_images):
_empty_cache()
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S_%f")
target_dir = f"input_images_{timestamp}"
target_dir_images = os.path.join(target_dir, "images")
if os.path.exists(target_dir):
shutil.rmtree(target_dir)
os.makedirs(target_dir_images, exist_ok=True)
image_paths = []
if input_images is not None:
if not isinstance(input_images, list):
input_images = [input_images]
for file_data in input_images:
file_path = _uploaded_file_path(file_data)
dst_path = os.path.join(target_dir_images, os.path.basename(file_path))
shutil.copy(file_path, dst_path)
image_paths.append(dst_path)
if input_video is not None:
video_path = _uploaded_file_path(input_video)
video = cv2.VideoCapture(video_path)
fps = video.get(cv2.CAP_PROP_FPS)
frame_interval = max(1, int(fps * VIDEO_FRAME_INTERVAL_SECONDS)) if fps > 0 else 1
count = 0
frame_index = 0
while True:
ok, frame = video.read()
if not ok:
break
if count % frame_interval == 0:
image_path = os.path.join(target_dir_images, f"{frame_index:06}.png")
cv2.imwrite(image_path, frame)
image_paths.append(image_path)
frame_index += 1
count += 1
video.release()
image_paths = sorted(image_paths, key=_natural_key)
return target_dir, image_paths
def update_gallery_on_upload(input_video, input_images):
if not input_video and not input_images:
return None, None, None, "Please upload a video or images.", _visibility_tolerance_update(0)
target_dir, image_paths = handle_uploads(input_video, input_images)
return (
None,
target_dir,
image_paths,
"Upload complete. Click Reconstruct to begin.",
_visibility_tolerance_update(len(image_paths)),
)
def load_example_scene(scene_name):
if not scene_name:
return None, None, None, "Please select an example scene.", _visibility_tolerance_update(0)
image_paths = _example_scene_paths(scene_name)
if not image_paths:
return None, None, None, f"No images found for example scene: {scene_name}", _visibility_tolerance_update(0)
target_dir, copied_paths = handle_uploads(None, image_paths)
return (
None,
target_dir,
copied_paths,
f"Loaded example scene {scene_name} ({len(copied_paths)} images). Click Reconstruct to begin.",
_visibility_tolerance_update(len(copied_paths)),
)
def load_example_scene_from_gallery(event: gr.SelectData):
index = event.index
if isinstance(index, (list, tuple)):
index = index[0]
scene_names = _example_scene_choices()
try:
scene_name = scene_names[int(index)]
except (TypeError, ValueError, IndexError):
return None, None, None, "Please select an example scene.", _visibility_tolerance_update(0)
return load_example_scene(scene_name)
@gpu(duration=120)
def gradio_demo(
target_dir,
conf_thres=50.0,
remove_far_points=True,
static_only=False,
visibility_tolerance=2,
show_cam=True,
prediction_mode="Depthmap and Camera Branch",
):
if not target_dir or target_dir == "None" or not os.path.isdir(target_dir):
return None, "No valid input directory. Please upload first."
start_time = time.time()
_empty_cache()
target_dir_images = os.path.join(target_dir, "images")
image_files = sorted(os.listdir(target_dir_images), key=_natural_key) if os.path.isdir(target_dir_images) else []
frame_filter = "All"
predictions = run_model(target_dir, model)
np.savez(os.path.join(target_dir, "predictions.npz"), **predictions)
visibility_tolerance = _clamp_visibility_tolerance(
visibility_tolerance, _prediction_frame_count(predictions)
)
_attach_static_mask(
predictions,
target_dir=target_dir,
static_only=static_only,
visibility_tolerance=visibility_tolerance,
remove_far_points=remove_far_points,
)
glbfile = _glb_path(
target_dir,
conf_thres,
frame_filter,
remove_far_points,
static_only,
visibility_tolerance,
show_cam,
prediction_mode,
)
scene = predictions_to_glb(
predictions,
conf_thres=conf_thres,
filter_by_frames=frame_filter,
remove_far_points=remove_far_points,
static_only=static_only,
show_cam=show_cam,
prediction_mode=prediction_mode,
)
scene.export(file_obj=glbfile)
del predictions
_empty_cache()
elapsed = time.time() - start_time
return glbfile, f"Reconstruction complete ({len(image_files)} frames, {elapsed:.2f}s)."
@gpu(duration=120)
def update_visualization(
target_dir,
conf_thres,
remove_far_points,
static_only,
visibility_tolerance,
show_cam,
prediction_mode,
):
if not target_dir or target_dir == "None" or not os.path.isdir(target_dir):
return None, "No reconstruction available. Please click Reconstruct first."
predictions_path = os.path.join(target_dir, "predictions.npz")
if not os.path.exists(predictions_path):
return None, "No saved predictions found. Please click Reconstruct first."
predictions = _load_predictions(predictions_path)
visibility_tolerance = _clamp_visibility_tolerance(
visibility_tolerance, _prediction_frame_count(predictions)
)
_attach_static_mask(
predictions,
target_dir=target_dir,
static_only=static_only,
visibility_tolerance=visibility_tolerance,
remove_far_points=remove_far_points,
)
frame_filter = "All"
glbfile = _glb_path(
target_dir,
conf_thres,
frame_filter,
remove_far_points,
static_only,
visibility_tolerance,
show_cam,
prediction_mode,
)
if not os.path.exists(glbfile):
scene = predictions_to_glb(
predictions,
conf_thres=conf_thres,
filter_by_frames=frame_filter,
remove_far_points=remove_far_points,
static_only=static_only,
show_cam=show_cam,
prediction_mode=prediction_mode,
)
scene.export(file_obj=glbfile)
return glbfile, "Visualization updated."
def _glb_path(
target_dir,
conf_thres,
frame_filter,
remove_far_points,
static_only,
visibility_tolerance,
show_cam,
prediction_mode,
):
safe_frame = str(frame_filter).replace(".", "_").replace(":", "").replace(" ", "_")
safe_mode = prediction_mode.replace(" ", "_")
safe_visibility_tolerance = int(visibility_tolerance) if static_only else "off"
return os.path.join(
target_dir,
(
f"glbscene_conf{conf_thres}_frame{safe_frame}_far{remove_far_points}"
f"_static{static_only}_vtol{safe_visibility_tolerance}_cam{show_cam}_pred{safe_mode}.glb"
),
)
def _load_predictions(predictions_path):
with np.load(predictions_path, allow_pickle=True) as loaded:
return {key: loaded[key] for key in loaded.keys()}
def _uploaded_file_path(file_data):
if isinstance(file_data, dict):
return file_data.get("path") or file_data.get("name")
if hasattr(file_data, "path"):
return file_data.path
if hasattr(file_data, "name"):
return file_data.name
return file_data
def _attach_static_mask(
predictions,
*,
target_dir,
static_only,
visibility_tolerance,
remove_far_points,
):
if not static_only:
predictions.pop("static_mask", None)
return
depth = np.asarray(predictions["depth"])
if depth.ndim == 4 and depth.shape[-1] == 1:
depth = depth[..., 0]
if depth.ndim != 3:
raise ValueError(f"Expected depth with shape [S, H, W] or [S, H, W, 1], got {predictions['depth'].shape}")
visibility_tolerance = _clamp_visibility_tolerance(visibility_tolerance, depth.shape[0])
static_mask_path = _static_mask_path(target_dir, visibility_tolerance, remove_far_points)
if os.path.exists(static_mask_path):
predictions["static_mask"] = np.load(static_mask_path)
return
if depth.shape[0] < 2:
predictions["static_mask"] = np.isfinite(depth) & (depth > 0)
np.save(static_mask_path, predictions["static_mask"])
return
device = _device()
depth_tensor = torch.from_numpy(np.ascontiguousarray(depth)).float().unsqueeze(0).to(device)
extrinsic_tensor = torch.from_numpy(np.ascontiguousarray(predictions["extrinsic"])).float().unsqueeze(0).to(device)
intrinsic_tensor = torch.from_numpy(np.ascontiguousarray(predictions["intrinsic"])).float().unsqueeze(0).to(device)
candidate_masks = None
if remove_far_points:
candidate_masks = depth_foreground_mask(depth_tensor)
valid_depth = torch.isfinite(depth_tensor) & (depth_tensor > 0)
empty_candidate = ~candidate_masks.flatten(-2).any(dim=-1)
if empty_candidate.any():
candidate_masks[empty_candidate] = valid_depth[empty_candidate]
static_mask = build_static_masks(
depth=depth_tensor,
extrinsics=extrinsic_tensor,
intrinsics=intrinsic_tensor,
candidate_masks=candidate_masks,
visibility_tolerance=int(visibility_tolerance),
)
predictions["static_mask"] = static_mask.squeeze(0).detach().cpu().numpy()
np.save(static_mask_path, predictions["static_mask"])
_empty_cache()
def _static_mask_path(target_dir, visibility_tolerance, remove_far_points):
return os.path.join(
target_dir,
f"static_mask_vtol{int(visibility_tolerance)}_far{bool(remove_far_points)}.npy",
)
def _prediction_frame_count(predictions):
depth = np.asarray(predictions["depth"])
if depth.ndim == 4 and depth.shape[-1] == 1:
depth = depth[..., 0]
if depth.ndim != 3:
return 0
return int(depth.shape[0])
def _clamp_visibility_tolerance(value, num_frames):
max_value = max(0, int(num_frames))
try:
value = int(value)
except (TypeError, ValueError):
value = 0
return min(max(value, 0), max_value)
def _visibility_tolerance_update(num_images, value=2):
max_value = max(0, int(num_images))
value = _clamp_visibility_tolerance(value, max_value)
return gr.Slider(maximum=max_value, value=value, interactive=max_value > 0)
def _example_scene_map():
scene_map = {}
for root_dir in ("examples", "exampes"):
for scene_dir in sorted(glob.glob(os.path.join(root_dir, "scene*")), key=_natural_key):
if not os.path.isdir(scene_dir):
continue
image_paths = _scene_image_paths(scene_dir)
if not image_paths:
continue
scene_name = os.path.basename(scene_dir)
scene_map.setdefault(scene_name, image_paths)
return scene_map
def _example_scene_choices():
return list(_example_scene_map().keys())
def _example_scene_gallery_value():
scene_map = _example_scene_map()
return [(paths[0], scene_name) for scene_name, paths in scene_map.items() if paths]
def _example_scene_paths(scene_name):
return _example_scene_map().get(scene_name, [])
def _scene_image_paths(scene_dir):
image_extensions = ("*.jpg", "*.jpeg", "*.png", "*.bmp", "*.webp")
image_paths = []
for pattern in image_extensions:
image_paths.extend(glob.glob(os.path.join(scene_dir, pattern)))
return sorted(image_paths, key=_natural_key)
def _natural_key(path):
base = os.path.basename(path)
parts = re.split(r"(\d+)", base)
return [int(part) if part.isdigit() else part.lower() for part in parts]
def _header_html():
return f"""
<div style="text-align: center; max-width: 960px; margin: 0 auto 18px auto;">
<h1 style="margin-bottom: 0.25em;">🌊 Wat3R</h1>
<h3 style="font-weight: 500; margin-top: 0;">
Underwater 3D Geometry Learning without Annotations
</h3>
<p style="margin: 0.75em auto 1em auto;">
Upload a video or a set of images. Wat3R predicts camera poses, depth maps,
and 3D point maps, then exports an interactive GLB reconstruction.
</p>
<p>
<a href="{PAPER_URL}" target="_blank"
style="display: inline-block; padding: 8px 14px; margin: 4px;
border-radius: 8px; background: #f0f0f0; text-decoration: none; color: #111;">
📄 Paper
</a>
<a href="{CODE_URL}" target="_blank"
style="display: inline-block; padding: 8px 14px; margin: 4px;
border-radius: 8px; background: #f0f0f0; text-decoration: none; color: #111;">
💻 Code
</a>
</p>
</div>
"""
def clear_model3d():
return None
def update_log():
return "Loading and reconstructing..."
theme = gr.themes.Ocean()
with gr.Blocks(theme=theme) as demo:
gr.HTML(_header_html())
target_dir_output = gr.Textbox(label="Target Dir", visible=False, value="None")
example_scene_gallery_value = _example_scene_gallery_value()
with gr.Row():
with gr.Column(scale=2):
input_video = gr.Video(label="Upload Video", interactive=True)
input_images = gr.File(file_count="multiple", label="Upload Images", interactive=True)
image_gallery = gr.Gallery(
label="Preview",
columns=4,
height="300px",
show_download_button=True,
object_fit="contain",
preview=True,
)
with gr.Column(scale=4):
log_output = gr.Markdown("Please upload a video or images, then click Reconstruct.")
reconstruction_output = gr.Model3D(height=520, zoom_speed=0.5, pan_speed=0.5)
with gr.Row():
submit_btn = gr.Button("Reconstruct", variant="primary")
clear_btn = gr.ClearButton(
[input_video, input_images, reconstruction_output, log_output, target_dir_output, image_gallery]
)
prediction_mode = gr.Radio(
["Depthmap and Camera Branch", "Pointmap Branch"],
label="Prediction Mode",
value="Depthmap and Camera Branch",
)
conf_thres = gr.Slider(minimum=0, maximum=100, value=50, step=0.1, label="Confidence Threshold (%)")
show_cam = gr.Checkbox(label="Show Camera", value=True)
remove_far_points = gr.Checkbox(label="Remove Far Points (Two-Mean)", value=True)
static_only = gr.Checkbox(label="Show Static Points Only", value=False)
visibility_tolerance = gr.Slider(
minimum=0,
maximum=0,
value=0,
step=1,
label="Visibility Tolerance",
interactive=False,
)
if example_scene_gallery_value:
gr.Markdown("### Example Scenes")
example_gallery = gr.Gallery(
value=example_scene_gallery_value,
label="Click a scene to load it",
columns=min(4, len(example_scene_gallery_value)),
height=220,
show_download_button=False,
object_fit="contain",
preview=False,
)
example_gallery.select(
fn=load_example_scene_from_gallery,
inputs=[],
outputs=[
reconstruction_output,
target_dir_output,
image_gallery,
log_output,
visibility_tolerance,
],
)
submit_btn.click(fn=clear_model3d, inputs=[], outputs=[reconstruction_output]).then(
fn=update_log, inputs=[], outputs=[log_output]
).then(
fn=gradio_demo,
inputs=[
target_dir_output,
conf_thres,
remove_far_points,
static_only,
visibility_tolerance,
show_cam,
prediction_mode,
],
outputs=[reconstruction_output, log_output],
)
for component in [conf_thres, remove_far_points, static_only, visibility_tolerance, show_cam, prediction_mode]:
component.change(
update_visualization,
[
target_dir_output,
conf_thres,
remove_far_points,
static_only,
visibility_tolerance,
show_cam,
prediction_mode,
],
[reconstruction_output, log_output],
)
input_video.change(
fn=update_gallery_on_upload,
inputs=[input_video, input_images],
outputs=[reconstruction_output, target_dir_output, image_gallery, log_output, visibility_tolerance],
)
input_images.change(
fn=update_gallery_on_upload,
inputs=[input_video, input_images],
outputs=[reconstruction_output, target_dir_output, image_gallery, log_output, visibility_tolerance],
)
if __name__ == "__main__":
demo.queue(max_size=20).launch(show_error=True)