Wat3R / visual_util.py
lsxi77777's picture
Update Wat3R Space demo
efb793f verified
Raw
History Blame Contribute Delete
7.86 kB
from __future__ import annotations
import numpy as np
import trimesh
from matplotlib import colormaps
from scipy.spatial.transform import Rotation
def predictions_to_glb(
predictions: dict,
conf_thres: float = 50.0,
filter_by_frames: str = "All",
remove_far_points: bool = True,
static_only: bool = False,
show_cam: bool = True,
prediction_mode: str = "Depthmap and Camera Branch",
) -> trimesh.Scene:
"""Convert Wat3R predictions to a GLB scene."""
if not isinstance(predictions, dict):
raise TypeError("predictions must be a dictionary")
points, confidence = _select_points(predictions, prediction_mode)
images = _images_to_nhwc(predictions["images"])
depth = _depth_to_shw(predictions.get("depth"))
static_mask = _mask_to_shw(predictions.get("static_mask"))
extrinsics = predictions["extrinsic"]
frame_index = _parse_frame_index(filter_by_frames)
if frame_index is not None:
points = points[frame_index][None]
confidence = confidence[frame_index][None]
images = images[frame_index][None]
if depth is not None:
depth = depth[frame_index][None]
if static_mask is not None:
static_mask = static_mask[frame_index][None]
extrinsics = extrinsics[frame_index][None]
vertices = points.reshape(-1, 3)
colors = (images.reshape(-1, 3).clip(0, 1) * 255).astype(np.uint8)
confidence = confidence.reshape(-1)
if conf_thres == 0:
conf_threshold = 0.0
else:
conf_threshold = np.percentile(confidence, conf_thres)
valid_mask = (confidence >= conf_threshold) & (confidence > 1e-5) & np.isfinite(vertices).all(axis=1)
if remove_far_points and depth is not None:
valid_mask &= _two_mean_near_mask(depth).reshape(-1)
if static_only:
if static_mask is None:
valid_mask &= False
else:
valid_mask &= static_mask.reshape(-1)
vertices = vertices[valid_mask]
colors = colors[valid_mask]
scene_scale = _scene_scale(vertices)
scene = trimesh.Scene()
scene.add_geometry(trimesh.PointCloud(vertices=vertices, colors=colors))
extrinsics_4x4 = _to_4x4(extrinsics)
if show_cam:
cmap = colormaps.get_cmap("gist_rainbow")
for index, world_to_camera in enumerate(extrinsics_4x4):
camera_to_world = np.linalg.inv(world_to_camera)
rgba = cmap(index / max(1, len(extrinsics_4x4)))
color = tuple(int(255 * value) for value in rgba[:3])
_add_camera(scene, camera_to_world, color, scene_scale)
if len(extrinsics_4x4) > 0:
_align_scene_to_first_camera(scene, extrinsics_4x4[0])
return scene
def _select_points(predictions: dict, prediction_mode: str):
if "Pointmap" in prediction_mode and "world_points" in predictions:
return (
predictions["world_points"],
predictions.get("world_points_conf", np.ones_like(predictions["world_points"][..., 0])),
)
return (
predictions["world_points_from_depth"],
predictions.get("depth_conf", np.ones_like(predictions["world_points_from_depth"][..., 0])),
)
def _images_to_nhwc(images: np.ndarray) -> np.ndarray:
if images.ndim != 4:
raise ValueError(f"Expected images with shape [S, C, H, W] or [S, H, W, C], got {images.shape}")
if images.shape[1] == 3:
images = np.transpose(images, (0, 2, 3, 1))
return images
def _depth_to_shw(depth: np.ndarray | None) -> np.ndarray | None:
if depth is None:
return None
depth = np.asarray(depth)
if depth.ndim == 5 and depth.shape[0] == 1:
depth = depth[0]
if depth.ndim == 4 and depth.shape[-1] == 1:
depth = depth[..., 0]
if depth.ndim != 3:
return None
return depth
def _mask_to_shw(mask: np.ndarray | None) -> np.ndarray | None:
if mask is None:
return None
mask = np.asarray(mask)
if mask.ndim == 5 and mask.shape[0] == 1:
mask = mask[0]
if mask.ndim == 4 and mask.shape[-1] == 1:
mask = mask[..., 0]
if mask.ndim != 3:
return None
return mask.astype(bool)
def _two_mean_near_mask(
depth: np.ndarray,
num_iters: int = 20,
clip_percentiles: tuple[float, float] = (1.0, 99.0),
) -> np.ndarray:
"""Split each depth map into near/far clusters and keep the near cluster."""
depth = np.asarray(depth, dtype=np.float32)
masks = []
for frame_depth in depth:
finite = np.isfinite(frame_depth)
if not finite.any():
masks.append(np.zeros(frame_depth.shape, dtype=bool))
continue
values = frame_depth[finite]
low, high = np.percentile(values, clip_percentiles)
values_clipped = np.clip(values, low, high)
center_1 = float(values_clipped.min())
center_2 = float(values_clipped.max())
if abs(center_1 - center_2) < 1e-8:
masks.append(finite)
continue
for _ in range(num_iters):
assign_1 = np.abs(values_clipped - center_1) <= np.abs(values_clipped - center_2)
if assign_1.any():
center_1 = float(values_clipped[assign_1].mean())
if (~assign_1).any():
center_2 = float(values_clipped[~assign_1].mean())
near_center = min(center_1, center_2)
far_center = max(center_1, center_2)
threshold = 0.5 * (near_center + far_center) * 0.9
mask = finite & (frame_depth < threshold)
masks.append(mask if mask.any() else finite)
return np.stack(masks, axis=0)
def _parse_frame_index(frame_filter: str):
if frame_filter in (None, "All", "all"):
return None
try:
return int(str(frame_filter).split(":")[0])
except (TypeError, ValueError, IndexError):
return None
def _to_4x4(extrinsics: np.ndarray) -> np.ndarray:
extrinsics = np.asarray(extrinsics)
if extrinsics.ndim != 3 or extrinsics.shape[1:] != (3, 4):
raise ValueError(f"Expected extrinsics with shape [S, 3, 4], got {extrinsics.shape}")
matrices = np.tile(np.eye(4, dtype=np.float32), (extrinsics.shape[0], 1, 1))
matrices[:, :3, :4] = extrinsics
return matrices
def _scene_scale(vertices: np.ndarray) -> float:
if vertices.size == 0:
return 1.0
lower = np.percentile(vertices, 5, axis=0)
upper = np.percentile(vertices, 95, axis=0)
scale = float(np.linalg.norm(upper - lower))
return scale if scale > 1e-6 else 1.0
def _add_camera(scene: trimesh.Scene, camera_to_world: np.ndarray, color: tuple[int, int, int], scene_scale: float) -> None:
cam_width = scene_scale * 0.05
cam_height = scene_scale * 0.1
rotate_45 = np.eye(4)
rotate_45[:3, :3] = Rotation.from_euler("z", 45, degrees=True).as_matrix()
rotate_45[2, 3] = -cam_height
transform = camera_to_world @ _opencv_to_opengl() @ rotate_45
cone = trimesh.creation.cone(cam_width, cam_height, sections=4)
vertices = _transform_points(transform, cone.vertices)
faces = cone.faces
mesh = trimesh.Trimesh(vertices=vertices, faces=faces)
mesh.visual.face_colors[:, :3] = color
scene.add_geometry(mesh)
def _align_scene_to_first_camera(scene: trimesh.Scene, first_world_to_camera: np.ndarray) -> None:
align_rotation = np.eye(4)
align_rotation[:3, :3] = Rotation.from_euler("y", 180, degrees=True).as_matrix()
scene.apply_transform(np.linalg.inv(first_world_to_camera) @ _opencv_to_opengl() @ align_rotation)
def _opencv_to_opengl() -> np.ndarray:
matrix = np.eye(4)
matrix[1, 1] = -1
matrix[2, 2] = -1
return matrix
def _transform_points(transform: np.ndarray, points: np.ndarray) -> np.ndarray:
points_h = np.concatenate([points, np.ones((points.shape[0], 1))], axis=1)
return (points_h @ transform.T)[:, :3]