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:
co-authored by
Claude Fable 5
parent
918c882abd
commit
2f9aa76478
+23
-3
@@ -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))
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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,
|
||||
]
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user