SimamAnything2GS / depth_bridge.py
junaid-simamdigital's picture
Add tested depth and pose fusion bridge
e539df0 verified
Raw History Blame Contribute Delete
5.44 kB
"""Model-agnostic depth/pose bridge for the future Anything2GS backends."""
from __future__ import annotations
import numpy as np
def _intrinsics(intrinsics, width: int, height: int) -> np.ndarray:
if intrinsics is None:
return np.array([[width, 0.0, (width - 1) / 2], [0.0, height, (height - 1) / 2], [0.0, 0.0, 1.0]], dtype=np.float32)
matrix = np.asarray(intrinsics, dtype=np.float32)
if matrix.shape != (3, 3) or not np.isfinite(matrix).all():
raise ValueError("intrinsics must be a finite 3x3 matrix")
return matrix
def _pose(pose) -> np.ndarray:
if pose is None:
return np.eye(4, dtype=np.float32)
matrix = np.asarray(pose, dtype=np.float32)
if matrix.shape == (3, 4):
result = np.eye(4, dtype=np.float32)
result[:3] = matrix
matrix = result
if matrix.shape != (4, 4) or not np.isfinite(matrix).all():
raise ValueError("pose must be a finite 4x4 or 3x4 world-to-camera matrix")
return matrix
def depth_view_to_points(depth, colors, confidence=None, intrinsics=None, pose=None):
"""Project one inverse-relative-depth view into world coordinates.
The convention matches Simam3D: larger normalized depth values are closer
to the camera. Poses are world-to-camera; the returned points are world
points. This function does not estimate depth or camera pose.
"""
depth = np.asarray(depth, dtype=np.float32)
colors = np.asarray(colors, dtype=np.uint8)
if depth.ndim != 2 or colors.shape != (*depth.shape, 3):
raise ValueError("depth must be HxW and colors must be HxWx3")
height, width = depth.shape
confidence = np.ones_like(depth, dtype=np.float32) if confidence is None else np.asarray(confidence, dtype=np.float32)
if confidence.shape != depth.shape:
raise ValueError("confidence must match depth shape")
valid = np.isfinite(depth) & np.isfinite(confidence) & (confidence > 0)
if not valid.any():
return np.empty((0, 3), np.float32), np.empty((0, 3), np.uint8), np.empty((0,), np.float32)
safe_depth = np.nan_to_num(depth, nan=0.0, posinf=1.0, neginf=0.0)
minimum = float(safe_depth[valid].min())
z = 1.0 / np.maximum(safe_depth - minimum + 1e-3, 1e-3)
k = _intrinsics(intrinsics, width, height)
yy, xx = np.nonzero(valid)
camera = np.column_stack([
(xx - k[0, 2]) * z[yy, xx] / max(float(k[0, 0]), 1e-6),
(yy - k[1, 2]) * z[yy, xx] / max(float(k[1, 1]), 1e-6),
z[yy, xx],
]).astype(np.float32)
world_to_camera = _pose(pose)
camera_to_world = np.linalg.inv(world_to_camera)
homogeneous = np.column_stack([camera, np.ones(len(camera), dtype=np.float32)])
points = (homogeneous @ camera_to_world.T)[:, :3].astype(np.float32)
return points, colors[yy, xx], confidence[yy, xx].astype(np.float32)
def fuse_depth_views(views, voxel_size: float = 0.02, max_points: int = 150_000) -> dict[str, np.ndarray]:
"""Fuse projected views by confidence-weighted voxel centroid and color."""
if voxel_size <= 0 or max_points < 1:
raise ValueError("voxel_size and max_points must be positive")
all_points, all_colors, all_weights, all_sources = [], [], [], []
for source, view in enumerate(views):
points, colors, weights = depth_view_to_points(**view)
all_points.append(points)
all_colors.append(colors)
all_weights.append(weights)
all_sources.append(np.full(len(points), source, dtype=np.int32))
if not all_points or not any(len(points) for points in all_points):
return {"points": np.empty((0, 3), np.float32), "colors": np.empty((0, 3), np.uint8), "confidence": np.empty((0,), np.float32), "source_views": np.empty((0,), np.int32), "source_masks": np.empty((0,), np.int32)}
points = np.concatenate(all_points)
colors = np.concatenate(all_colors)
weights = np.concatenate(all_weights)
sources = np.concatenate(all_sources)
keys = np.floor(points / voxel_size).astype(np.int64)
order = np.lexsort((keys[:, 2], keys[:, 1], keys[:, 0]))
keys, points, colors, weights, sources = keys[order], points[order], colors[order], weights[order], sources[order]
unique, starts = np.unique(keys, axis=0, return_index=True)
ends = np.r_[starts[1:], len(keys)]
fused_points, fused_colors, fused_weights, fused_sources, masks = [], [], [], [], []
for start, end in zip(starts, ends):
local = weights[start:end]
total = float(local.sum())
fused_points.append((points[start:end] * local[:, None]).sum(axis=0) / max(total, 1e-8))
fused_colors.append(np.clip(np.rint((colors[start:end] * local[:, None]).sum(axis=0) / max(total, 1e-8)), 0, 255))
fused_weights.append(total / (end - start))
strongest = start + int(np.argmax(local))
fused_sources.append(sources[strongest])
mask = 0
for source in np.unique(sources[start:end]):
if 0 <= source < 30:
mask |= 1 << int(source)
masks.append(mask)
keep = np.argsort(-np.asarray(fused_weights))[:max_points]
return {
"points": np.asarray(fused_points, np.float32)[keep],
"colors": np.asarray(fused_colors, np.uint8)[keep],
"confidence": np.asarray(fused_weights, np.float32)[keep],
"source_views": np.asarray(fused_sources, np.int32)[keep],
"source_masks": np.asarray(masks, np.int32)[keep],
}