Exclude fisheye out-of-circle regions from 4D lifting and rendering

For FISHEYE inputs, the [-1,1] uv square contains the corners beyond the
image circle (r > 1, view angles beyond fov/2). Those regions carry no
scene content, yet:
- _uv_to_dirs clamped r > 1 onto the rim, so MotionMaskFromDepth could
  flag garbage-depth corners as dynamic and TracksToTrajectories lifted
  corner tracks to junk 3D control trajectories;
- _projection_valid only checked the square, so points at angles beyond
  fov/2 that project diagonally (e.g. r=1.33 at u=v~0.94) were treated as
  in-image by the motion-mask warp check and SplitSplatsByMask;
- render_gaussians had the same square-only cull, painting behind-camera
  splats into the corners of FISHEYE renders (and mirror-projecting
  behind-camera points in PINHOLE renders - Z>0 cull added to match
  GS4D's _projection_valid).

Add _uv_in_fov helper, apply it in MotionMaskFromDepth (corner pixels can
never be flagged dynamic) and TracksToTrajectories (corner samples are
invalid, fully-out tracks dropped), extend _projection_valid and the
render cull with the r <= 1 circle test.

New smoke test 11 covers all three paths: flickering-corner depth stays
static while real in-circle motion is flagged, a corner track is dropped,
and a beyond-fov splat falls outside an all-ones mask (11/11 pass).

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Alexander Kharin
2026-07-17 15:36:22 +03:00
co-authored by Claude Fable 5
parent 918c882abd
commit 2f9aa76478
3 changed files with 88 additions and 3 deletions
+23 -3
View File
@@ -180,6 +180,18 @@ def _uv_grid(height: int, width: int, device: torch.device) -> Tuple[torch.Tenso
return u, v
def _uv_in_fov(u: torch.Tensor, v: torch.Tensor, projection: str) -> torch.Tensor:
"""True where normalized uv lies inside the projection's actual image region.
For FISHEYE the [-1,1] square contains the corners beyond the image circle
(r > 1, i.e. view angles beyond fov/2); pixels there carry no scene content
(black corners / garbage depth) and must not be lifted, warped or tracked.
"""
if projection == "FISHEYE":
return (u * u + v * v) <= 1.0 + 1e-6
return torch.ones_like(u, dtype=torch.bool)
def _uv_to_dirs(u: torch.Tensor, v: torch.Tensor, projection: str, horizontal_fov: float) -> torch.Tensor:
"""Unit ray directions [...,3] in camera frame for normalized uv in [-1,1].
@@ -222,7 +234,9 @@ def _project_xyz(
def _projection_valid(
u: torch.Tensor, v: torch.Tensor, Z: torch.Tensor, projection: str
) -> torch.Tensor:
"""In-image validity for projected points; pinhole additionally requires Z>0."""
"""In-image validity for projected points; pinhole additionally requires Z>0,
fisheye requires the point inside the image circle (r <= 1), not just the
[-1,1] square — angles beyond fov/2 can otherwise land in the corners."""
valid = (
torch.isfinite(u)
& torch.isfinite(v)
@@ -233,6 +247,8 @@ def _projection_valid(
)
if projection == "PINHOLE":
valid = valid & (Z > 1e-6)
elif projection == "FISHEYE":
valid = valid & ((u * u + v * v) <= 1.0 + 1e-6)
return valid
@@ -480,6 +496,7 @@ class MotionMaskFromDepth:
u, v = _uv_grid(H, W, target_device)
dirs = _uv_to_dirs(u, v, input_projection, input_horizontal_fov) # [H,W,3]
in_fov = _uv_in_fov(u, v, input_projection) # excludes fisheye corners
gap = max(1, int(frame_gap))
dynamic = torch.zeros((T, H, W), device=target_device)
@@ -490,7 +507,7 @@ class MotionMaskFromDepth:
tr_t = poses[t, :3, 3]
# world = (cam - t) @ R (inverse of cam = world @ R.T + t)
world = (cam_pts - tr_t) @ R_t
has_depth = depth_t > 1e-6
has_depth = (depth_t > 1e-6) & in_fov
flagged = torch.zeros((H, W), dtype=torch.bool, device=target_device)
for t2 in (t + gap, t - gap):
@@ -693,7 +710,10 @@ class TracksToTrajectories:
# world = (cam - t) @ R per frame.
world = torch.bmm(cam - tr.unsqueeze(1), R)
in_bounds = (u >= -1.0) & (u <= 1.0) & (v >= -1.0) & (v <= 1.0)
in_bounds = (
(u >= -1.0) & (u <= 1.0) & (v >= -1.0) & (v <= 1.0)
& _uv_in_fov(u, v, input_projection)
)
valid = vis & in_bounds & (d > 1e-6) & torch.isfinite(world).all(dim=-1)
world = torch.where(valid.unsqueeze(-1), world, torch.zeros_like(world))
+6
View File
@@ -1303,6 +1303,12 @@ def render_gaussians(
u, v, depth = _xyz_to_equirect(X, Y, Z, camera_horizontal_fov)
valid = (u >= -1.0) & (u <= 1.0) & (v >= -1.0) & (v <= 1.0)
if camera_projection == "PINHOLE":
# behind-camera points otherwise mirror-project into the frame
valid = valid & (Z > 1e-6)
elif camera_projection == "FISHEYE":
# keep the image circle only: angles beyond fov/2 land in the corners
valid = valid & ((u * u + v * v) <= 1.0 + 1e-6)
if not valid.any():
return _empty_render(output_width, output_height, dev)
+59
View File
@@ -18,6 +18,7 @@ with small synthetic data:
8. align_depth_scale (world_nodes, contract C4) + DepthEdgeFilter
9. FuseSplats weighted voxel fusion
10. SphereSplatSeed pano -> splat sphere -> render round-trip
11. Fisheye image-circle exclusion (motion mask / tracks / splat split)
"""
import math
@@ -445,6 +446,63 @@ def test_10_sphere_splat_seed():
# --------------------------------------------------------------------------- #
# Runner
# --------------------------------------------------------------------------- #
def test_11_fisheye_circle_exclusion():
"""Fisheye pixels/points outside the image circle (r > 1) must be ignored:
corners never become 'dynamic', corner tracks never lift to 3D, and points
at view angles beyond fov/2 are not matched against the mask."""
T, H, W = 6, 32, 32
fov = 180.0
# --- MotionMaskFromDepth: wildly flickering corners must stay static ----
depth = torch.full((T, H, W), 5.0)
depth[:, 14:18, 14:18] = torch.linspace(5.0, 2.0, T).view(T, 1, 1) # real motion
u = torch.linspace(-1, 1, W).view(1, W).expand(H, W)
v = torch.linspace(-1, 1, H).view(H, 1).expand(H, W)
corners = (u * u + v * v) > 1.0 + 1e-6
for t in range(T): # garbage depth flicker outside the circle
depth[t][corners] = 1.0 if t % 2 == 0 else 10.0
identity = torch.eye(4).unsqueeze(0).expand(T, 4, 4).contiguous()
(dyn,) = GS4D_nodes.MotionMaskFromDepth().motion_mask(
depth_seq=depth, trajectory=identity, input_projection="FISHEYE",
input_horizontal_fov=fov, threshold=0.10, frame_gap=2, dilate=0,
device="cpu",
)
assert dyn[:, corners].max().item() == 0.0, "fisheye corners were flagged dynamic"
assert dyn[:, 14:18, 14:18].max().item() == 1.0, "real in-circle motion missed"
# --- TracksToTrajectories: corner track must be dropped -----------------
tracks = torch.zeros(T, 2, 2)
tracks[:, 0, 0] = W / 2.0 # center track
tracks[:, 0, 1] = H / 2.0
tracks[:, 1, 0] = 0.0 # corner track (u=v=-1, r=1.414)
tracks[:, 1, 1] = 0.0
visibility = torch.ones(T, 2)
traj3d, track_ok = GS4D_nodes.TracksToTrajectories().tracks_to_trajectories(
tracks=tracks, visibility=visibility, depth_seq=depth,
input_projection="FISHEYE", input_horizontal_fov=fov,
min_visible_frac=0.5, device="cpu",
)
assert bool(track_ok[0]), "in-circle track unexpectedly invalid"
assert not bool(track_ok[1]), "out-of-circle corner track was lifted to 3D"
# --- SplitSplatsByMask: angle beyond fov/2 lands diagonally inside the
# [-1,1] square (r=1.33, u=v~0.94) but must not be matched to the mask ---
front = torch.tensor([[0.0, 0.0, 1.0]])
theta = math.radians(120.0)
behind_diag = torch.tensor([[
math.sin(theta) * math.cos(math.radians(45.0)),
math.sin(theta) * math.sin(math.radians(45.0)),
math.cos(theta),
]])
splats = make_splats(torch.cat([front, behind_diag], dim=0))
inside, outside = GS4D_nodes.SplitSplatsByMask().split_splats(
splats=splats, mask=torch.ones(H, W), projection="FISHEYE",
horizontal_fov=fov, threshold=0.5, device="cpu",
)
assert inside.xyz.shape[0] == 1, "expected only the in-fov splat inside"
assert outside.xyz.shape[0] == 1, "beyond-fov splat must fall outside"
TESTS = [
test_01_interpolate_se3,
test_02_render_gaussians_shapes_and_empty,
@@ -456,6 +514,7 @@ TESTS = [
test_08_align_depth_scale_and_depth_edge_filter,
test_09_fuse_splats,
test_10_sphere_splat_seed,
test_11_fisheye_circle_exclusion,
]