Compare commits

...
Author SHA1 Message Date
google-labs-jules[bot] d49f3f03ce Fix readability 2025-05-22 20:18:58 +00:00
2 changed files with 403 additions and 177 deletions
+374 -166
View File
@@ -95,76 +95,46 @@ class FisheyeDepthEstimator:
depth_full, = de_node.estimate_depth(image, model_name, depth_scale)
mask_full = (depth_full > 0).float()
# 2) Pinhole orientations (5 views)
rotations = [
(0, 0, 0), # front
(0, 45, 0), # right
(0, -45, 0), # left
(45, 0, 0), # up
(-45, 0, 0), # down
]
# 2) Generate Pinhole Views
fisheye_depths, fisheye_masks = self._generate_pinhole_views(
image,
de_node, z2r_node, ri_node, rd_node,
fisheye_fov, pinhole_fov,
pin_w, pin_h, fish_w, fish_h,
model_name, depth_scale, median_blur_kernel
)
fisheye_depths = []
fisheye_masks = []
# euler → matrix
def euler_to_matrix(pitch, yaw, roll):
p, y, r = map(math.radians, (pitch, yaw, roll))
Rx = torch.tensor([[1,0,0],[0,math.cos(p),-math.sin(p)],[0,math.sin(p),math.cos(p)]], dtype=torch.float32)
Ry = torch.tensor([[math.cos(y),0,math.sin(y)],[0,1,0],[-math.sin(y),0,math.cos(y)]], dtype=torch.float32)
Rz = torch.tensor([[math.cos(r),-math.sin(r),0],[math.sin(r),math.cos(r),0],[0,0,1]], dtype=torch.float32)
R = Rz @ Ry @ Rx
M = torch.eye(4, dtype=torch.float32)
M[:3, :3] = R
return M
# 3) Process each orientation
for pitch, yaw, roll in rotations:
M = euler_to_matrix(pitch, yaw, roll)
M_np = M.numpy()
M_inv = torch.inverse(M).numpy()
# fisheye → pinhole
img_pin, mask_pin = ri_node.reproject_image(
image,
input_horiszontal_fov = fisheye_fov,
output_horiszontal_fov= pinhole_fov,
input_projection = "FISHEYE",
output_projection = "PINHOLE",
output_width = pin_w,
output_height = pin_h,
transform_matrix = M_np,
feathering = 0,
)
# estimate pinhole depth
depth_pin, = de_node.estimate_depth(img_pin, model_name, depth_scale, median_blur_kernel=median_blur_kernel)
depth_pin, = z2r_node.depth_to_ray_depth(
depth_pin,
pinhole_fov,
)
# pinhole → fisheye
fish_depth, fish_mask = rd_node.reproject_depth(
depth_pin,
input_horizontal_fov = pinhole_fov,
output_horizontal_fov= fisheye_fov,
input_projection = "PINHOLE",
output_projection = "FISHEYE",
output_width = fish_w,
output_height = fish_h,
transform_matrix = M_inv,
)
# squeeze mask to [B,H,W]
fish_mask = fish_mask.squeeze(1)
fisheye_depths.append(fish_depth) # [B,H,W]
fisheye_masks.append(fish_mask)
fisheye_depths.append(depth_full) # [B,H,W]
fisheye_masks.append(mask_full.squeeze(-1)) # [B,H,W 1]
# merged mask
merged_mask = torch.sum(torch.stack(fisheye_masks), dim=0) > 0.5
# print(fisheye_depths[0].shape, fisheye_depths[-1].shape, merged_mask.shape)
# 4) Merge in sequence
d_acc, m_acc = self._merge_depths(
fisheye_depths, fisheye_masks,
ren_node, comb_node,
mode, softmerge_radius
)
# 5) Circular mask
ys = torch.arange(fish_h, device=d_acc.device).view(1, fish_h, 1)
xs = torch.arange(fish_w, device=d_acc.device).view(1, 1, fish_w)
cy = (fish_h - 1) / 2.0
cx = (fish_w - 1) / 2.0
dist2 = (ys - cy)**2 + (xs - cx)**2
radius2 = (min(fish_w, fish_h) / 2.0)**2
circ_mask = (dist2 <= radius2).float()
return d_acc, circ_mask
def _merge_depths(
self,
fisheye_depths: list,
fisheye_masks: list,
ren_node: DepthRenormalizer,
comb_node: CombineDepthsNode,
mode: str,
softmerge_radius: int,
) -> Tuple[torch.Tensor, torch.Tensor]:
d_acc = fisheye_depths[0]
m_acc = fisheye_masks[0]
for d_new, m_new in zip(fisheye_depths[1:-1], fisheye_masks[1:-1]):
@@ -187,20 +157,120 @@ class FisheyeDepthEstimator:
m_acc,
d_norm,
m_new_last,
mode = "SRC",
mode = "SRC", # Use SRC for the full fisheye to preserve its details
invert_mask = False,
softmerge_radius = softmerge_radius
)
return d_acc, m_acc
# 5) Circular mask
ys = torch.arange(fish_h, device=d_acc.device).view(1, fish_h, 1)
xs = torch.arange(fish_w, device=d_acc.device).view(1, 1, fish_w)
cy = (fish_h - 1) / 2.0
cx = (fish_w - 1) / 2.0
dist2 = (ys - cy)**2 + (xs - cx)**2
radius2 = (min(fish_w, fish_h) / 2.0)**2
circ_mask = (dist2 <= radius2).float()
return d_acc, circ_mask
def _generate_pinhole_views(
self,
image: torch.Tensor,
de_node: DepthEstimatorNode,
z2r_node: ZDepthToRayDepthNode,
ri_node: ReprojectImage,
rd_node: ReprojectDepth,
fisheye_fov: float,
pinhole_fov: float,
pin_w: int,
pin_h: int,
fish_w: int,
fish_h: int,
model_name: str,
depth_scale: float,
median_blur_kernel: int,
) -> Tuple[list, list]:
rotations = [
(0, 0, 0), # front
(0, 45, 0), # right
(0, -45, 0), # left
(45, 0, 0), # up
(-45, 0, 0), # down
]
fisheye_depths = []
fisheye_masks = []
for pitch, yaw, roll in rotations:
M = self._euler_to_matrix(pitch, yaw, roll)
M_np = M.numpy()
M_inv = torch.inverse(M).numpy()
fish_depth, fish_mask = self._process_view(
image, M_np, M_inv,
de_node, z2r_node, ri_node, rd_node,
fisheye_fov, pinhole_fov,
pin_w, pin_h, fish_w, fish_h,
model_name, depth_scale, median_blur_kernel
)
fisheye_depths.append(fish_depth)
fisheye_masks.append(fish_mask)
return fisheye_depths, fisheye_masks
def _process_view(
self,
image: torch.Tensor,
M_np: np.ndarray,
M_inv: np.ndarray,
de_node: DepthEstimatorNode,
z2r_node: ZDepthToRayDepthNode,
ri_node: ReprojectImage,
rd_node: ReprojectDepth,
fisheye_fov: float,
pinhole_fov: float,
pin_w: int,
pin_h: int,
fish_w: int,
fish_h: int,
model_name: str,
depth_scale: float,
median_blur_kernel: int,
) -> Tuple[torch.Tensor, torch.Tensor]:
# fisheye → pinhole
img_pin, mask_pin = ri_node.reproject_image(
image,
input_horiszontal_fov = fisheye_fov,
output_horiszontal_fov= pinhole_fov,
input_projection = "FISHEYE",
output_projection = "PINHOLE",
output_width = pin_w,
output_height = pin_h,
transform_matrix = M_np,
feathering = 0,
)
# estimate pinhole depth
depth_pin, = de_node.estimate_depth(img_pin, model_name, depth_scale, median_blur_kernel=median_blur_kernel)
depth_pin, = z2r_node.depth_to_ray_depth(
depth_pin,
pinhole_fov,
)
# pinhole → fisheye
fish_depth, fish_mask = rd_node.reproject_depth(
depth_pin,
input_horizontal_fov = pinhole_fov,
output_horizontal_fov= fisheye_fov,
input_projection = "PINHOLE",
output_projection = "FISHEYE",
output_width = fish_w,
output_height = fish_h,
transform_matrix = M_inv,
)
# squeeze mask to [B,H,W]
fish_mask = fish_mask.squeeze(1)
return fish_depth, fish_mask
def _euler_to_matrix(self, pitch, yaw, roll):
p, y, r = map(math.radians, (pitch, yaw, roll))
Rx = torch.tensor([[1,0,0],[0,math.cos(p),-math.sin(p)],[0,math.sin(p),math.cos(p)]], dtype=torch.float32)
Ry = torch.tensor([[math.cos(y),0,math.sin(y)],[0,1,0],[-math.sin(y),0,math.cos(y)]], dtype=torch.float32)
Rz = torch.tensor([[math.cos(r),-math.sin(r),0],[math.sin(r),math.cos(r),0],[0,0,1]], dtype=torch.float32)
R = Rz @ Ry @ Rx
M = torch.eye(4, dtype=torch.float32)
M[:3, :3] = R
return M
class PointcloudTrajectoryEnricher:
"""
@@ -288,106 +358,244 @@ class PointcloudTrajectoryEnricher:
debug_img = torch.zeros((1, height, width, 3), device=device)
debug_depth = torch.zeros((1, height, width, 1), device=device)
enriched_pc = pointcloud
# Initialize debug_img and debug_depth which will be updated in the loop
# and will hold the values from the last processed view.
debug_img = torch.zeros((1, height, width, 3), device=device)
debug_depth = torch.zeros((1, height, width, 1), device=device)
# loop over trajectory (limit or full)
for M in tqdm(trajectory[:15], desc="Enriching trajectory"):
M_np = M.cpu().numpy()
M_inv = np.linalg.inv(M_np)
enriched_pc, view_debug_img, view_debug_depth = self._process_single_view(
M, enriched_pc, device,
proj_node, outpaint_node, depth_node, renorm_node,
depth2pc_node, transform_node, clean_node, zdepth_node,
camera_type, horizontal_fov, width, height,
patch_projection, patch_horiz_fov, patch_res,
patch_phi, patch_theta, prompt,
num_inference_steps, guidance_scale, mask_blur,
voxel_size, min_points_per_voxel, model_name
)
debug_img = view_debug_img
debug_depth = view_debug_depth
return enriched_pc, debug_img, debug_depth
# transform and select front points
rotated, = transform_node.transform_pointcloud(enriched_pc, M_np)
pc_front = rotated[rotated[:, 2] > 0]
def _process_single_view(
self,
M_matrix: torch.Tensor,
current_enriched_pc: torch.Tensor,
device: torch.device,
proj_node: ProjectPointCloud,
outpaint_node: OutpaintAnyProjection,
depth_node: DepthEstimatorNode,
renorm_node: DepthRenormalizer,
depth2pc_node: DepthToPointCloud,
transform_node: TransformPointCloud,
clean_node: PointCloudCleaner,
zdepth_node: ZDepthToRayDepthNode,
camera_type: str,
horizontal_fov: float,
width: int,
height: int,
patch_projection: str,
patch_horiz_fov: float,
patch_res: int,
patch_phi: float,
patch_theta: float,
prompt: str,
num_inference_steps: int,
guidance_scale: float,
mask_blur: int,
voxel_size: float,
min_points_per_voxel: int,
model_name: str,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
M_np = M_matrix.cpu().numpy()
M_inv = np.linalg.inv(M_np)
# clean front points
pc_front, = clean_node.clean_pointcloud(
pc_front,
voxel_size=voxel_size,
min_points_per_voxel=min_points_per_voxel,
width=4096,
height=4096,
)
img, mask, depth_map, pc_front = self._prepare_view_data(
current_enriched_pc, M_np, transform_node, clean_node, proj_node,
voxel_size, min_points_per_voxel, camera_type, horizontal_fov,
width, height
)
# project to image + depth
img, mask, depth_map = proj_node.project_pointcloud(
pc_front,
camera_type,
horizontal_fov,
width,
height,
point_size=3,
return_inverse_depth=False,
)
debug_img = img
# fill nan in depthmap with (-1)
# outpaint missing regions
hole_mask = (mask < 0.5).float()
out_img, out_mask = outpaint_node.outpaint_any(
img,
input_projection = camera_type,
input_horiz_fov = horizontal_fov,
output_projection = camera_type,
output_horiz_fov = horizontal_fov,
output_width = width,
output_height = height,
patch_projection = patch_projection,
patch_horiz_fov = patch_horiz_fov,
patch_res = patch_res,
patch_phi = patch_phi,
patch_theta = patch_theta,
prompt = prompt,
num_inference_steps = num_inference_steps,
cached = False,
guidance_scale = guidance_scale,
mask_blur = mask_blur,
mask = hole_mask,
debug = False,
)
debug_img = out_img
# estimate and renormalize depth
nan_mask = torch.isnan(depth_map)
# …and replace them with –1.0 (in-place)
depth_map[nan_mask] = 0
# clip from -1 to 1000
depth_map = torch.clamp(depth_map, 0, 1000.0)
new_depth, = depth_node.estimate_depth(out_img, model_name, depth_scale=1.0)
new_depth, = zdepth_node.depth_to_ray_depth(
new_depth,
horizontal_fov,
)
# renormalize depth
norm_depth, = renorm_node.renormalize_depth(
new_depth,
depth_map,
depth_mask=(mask>=0.5)*1,
guidance_mask=(mask<0.5)*1,
use_inverse=False,
)
# median blur on depth
k = 5
d = norm_depth.permute(0,3,1,2) # [B,1,H,W]
pad = k//2
pd = F.pad(d, (pad, pad, pad, pad), mode='reflect')
patches = pd.unfold(2, k, 1).unfold(3, k, 1)
patches = patches.contiguous().view(d.shape[0], d.shape[1], d.shape[2], d.shape[3], k*k)
d, _ = patches.median(dim=-1)
norm_depth = d.permute(0,2,3,1) # [B,H,W,1]
debug_depth = norm_depth*hole_mask.unsqueeze(0).unsqueeze(-1)+depth_map*(1-hole_mask.unsqueeze(0).unsqueeze(-1))
# fill nan in depthmap with (-1)
# outpaint missing regions
hole_mask = (mask < 0.5).float()
out_img = self._outpaint_missing_regions(
img, hole_mask, # Pass hole_mask instead of the full mask
outpaint_node, camera_type, horizontal_fov, width, height,
patch_projection, patch_horiz_fov, patch_res,
patch_phi, patch_theta, prompt,
num_inference_steps, guidance_scale, mask_blur
)
# estimate and renormalize depth
norm_depth, debug_depth_view = self._estimate_and_refine_depth(
out_img, depth_map, mask, hole_mask,
depth_node, zdepth_node, renorm_node,
model_name, horizontal_fov
)
# back to pointcloud
pc_new, = depth2pc_node.depth_to_pointcloud(
out_img,
camera_type,
horizontal_fov,
depth_scale=1.0,
invert_depth=False,
depthmap=norm_depth,
mask=hole_mask,
)
# back to pointcloud
pc_world = self._convert_depth_to_world_pointcloud(
out_img, norm_depth, hole_mask, M_inv,
depth2pc_node, transform_node,
camera_type, horizontal_fov
)
# enriched_pc is not rotated
current_enriched_pc = torch.cat([current_enriched_pc, pc_world.to(device)], dim=0)
return current_enriched_pc, out_img, debug_depth_view # Return out_img and the depth for this view
pc_world, = transform_node.transform_pointcloud(pc_new, M_inv)
# enriched_pc is not rotated
enriched_pc = torch.cat([enriched_pc, pc_world.to(device)], dim=0)
return enriched_pc, debug_img, norm_depth
def _convert_depth_to_world_pointcloud(
self,
out_img: torch.Tensor,
norm_depth: torch.Tensor,
hole_mask: torch.Tensor,
M_inv: np.ndarray,
depth2pc_node: DepthToPointCloud,
transform_node: TransformPointCloud,
camera_type: str,
horizontal_fov: float,
) -> torch.Tensor:
pc_new, = depth2pc_node.depth_to_pointcloud(
out_img,
camera_type,
horizontal_fov,
depth_scale=1.0,
invert_depth=False,
depthmap=norm_depth,
mask=hole_mask,
)
pc_world, = transform_node.transform_pointcloud(pc_new, M_inv)
return pc_world
def _estimate_and_refine_depth(
self,
out_img: torch.Tensor,
depth_map: torch.Tensor,
original_mask: torch.Tensor, # Mask from projection
hole_mask: torch.Tensor,
depth_node: DepthEstimatorNode,
zdepth_node: ZDepthToRayDepthNode,
renorm_node: DepthRenormalizer,
model_name: str,
horizontal_fov: float,
) -> Tuple[torch.Tensor, torch.Tensor]:
nan_mask = torch.isnan(depth_map)
depth_map[nan_mask] = 0 # In-place modification
depth_map = torch.clamp(depth_map, 0, 1000.0)
new_depth, = depth_node.estimate_depth(out_img, model_name, depth_scale=1.0)
new_depth, = zdepth_node.depth_to_ray_depth(
new_depth,
horizontal_fov,
)
# renormalize depth
# Use original_mask for depth_mask as it represents valid projected areas
norm_depth, = renorm_node.renormalize_depth(
new_depth,
depth_map,
depth_mask=(original_mask >= 0.5) * 1,
guidance_mask=(hole_mask >= 0.5) * 1, # hole_mask is appropriate here
use_inverse=False,
)
# median blur on depth
k = 5
d = norm_depth.permute(0,3,1,2) # [B,1,H,W]
pad = k//2
pd = F.pad(d, (pad, pad, pad, pad), mode='reflect')
patches = pd.unfold(2, k, 1).unfold(3, k, 1)
patches = patches.contiguous().view(d.shape[0], d.shape[1], d.shape[2], d.shape[3], k*k)
d_median, _ = patches.median(dim=-1) # Renamed to avoid conflict
norm_depth_blurred = d_median.permute(0,2,3,1) # [B,H,W,1]
# Create debug_depth_view using the blurred normalized depth for holes
# and the original depth_map for non-holes.
debug_depth_view = norm_depth_blurred * hole_mask.unsqueeze(0).unsqueeze(-1) + \
depth_map * (1 - hole_mask.unsqueeze(0).unsqueeze(-1))
return norm_depth_blurred, debug_depth_view
def _outpaint_missing_regions(
self,
img: torch.Tensor,
hole_mask: torch.Tensor, # Expects the specific hole_mask
outpaint_node: OutpaintAnyProjection,
camera_type: str,
horizontal_fov: float,
width: int,
height: int,
patch_projection: str,
patch_horiz_fov: float,
patch_res: int,
patch_phi: float,
patch_theta: float,
prompt: str,
num_inference_steps: int,
guidance_scale: float,
mask_blur: int,
) -> torch.Tensor: # Returns only out_img, out_mask is not used later
out_img, _ = outpaint_node.outpaint_any( # Assign out_mask to _
img,
input_projection = camera_type,
input_horiz_fov = horizontal_fov,
output_projection = camera_type,
output_horiz_fov = horizontal_fov,
output_width = width,
output_height = height,
patch_projection = patch_projection,
patch_horiz_fov = patch_horiz_fov,
patch_res = patch_res,
patch_phi = patch_phi,
patch_theta = patch_theta,
prompt = prompt,
num_inference_steps = num_inference_steps,
cached = False,
guidance_scale = guidance_scale,
mask_blur = mask_blur,
mask = hole_mask, # Use the passed hole_mask
debug = False,
)
return out_img
def _prepare_view_data(
self,
current_enriched_pc: torch.Tensor,
M_np: np.ndarray,
transform_node: TransformPointCloud,
clean_node: PointCloudCleaner,
proj_node: ProjectPointCloud,
voxel_size: float,
min_points_per_voxel: int,
camera_type: str,
horizontal_fov: float,
width: int,
height: int,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
# transform and select front points
rotated, = transform_node.transform_pointcloud(current_enriched_pc, M_np)
pc_front = rotated[rotated[:, 2] > 0]
# clean front points
pc_front, = clean_node.clean_pointcloud(
pc_front,
voxel_size=voxel_size,
min_points_per_voxel=min_points_per_voxel,
width=4096, # Consider passing these as params if they vary
height=4096, # Consider passing these as params if they vary
)
# project to image + depth
img, mask, depth_map = proj_node.project_pointcloud(
pc_front,
camera_type,
horizontal_fov,
width,
height,
point_size=3,
return_inverse_depth=False,
)
return img, mask, depth_map, pc_front
NODE_CLASS_MAPPINGS = {"FisheyeDepthEstimator": FisheyeDepthEstimator,
"PointcloudTrajectoryEnricher": PointcloudTrajectoryEnricher}
+29 -11
View File
@@ -129,7 +129,7 @@ class OutpaintAnyProjection:
patch_mask = torch.ones_like(patch_mask) * 1 if debug else patch_mask
# 5) Reproject inpainted patch back
back_img, back_mask = reproj.reproject_image(
back_img, inpainted_patch_reproj_mask_raw = reproj.reproject_image( # Store raw mask output
inpainted_patch,
patch_horiz_fov, output_horiz_fov,
patch_projection, output_projection,
@@ -138,19 +138,37 @@ class OutpaintAnyProjection:
transform_matrix=rot_m,
feathering=0,
)
back_mask = normalize_mask(back_mask).bool() # True where patch contributes
# inpainted_patch_coverage_mask: 1.0 where the reprojected inpainted patch has content, 0.0 otherwise.
inpainted_patch_coverage_mask = normalize_mask(back_mask) # Ensure it's float [0,1]
# original coverage: True = had data, False = hole
orig_covered = ~base_mask.bool()
base_img=base_img * orig_covered.unsqueeze(-1)
# fill only holes where back_mask is False
filled = back_img * (~back_mask.unsqueeze(-1))*base_mask.unsqueeze(-1)
final_img = base_img+filled
# --- Define Masks based on Conventions ---
# base_mask: 1.0 where original reprojected image has content, 0.0 for holes.
# initial_hole_mask: 1.0 where original reprojected image has holes (inverse of base_mask).
initial_hole_mask = 1.0 - base_mask
# inpainted_patch_coverage_mask: 1.0 where reprojected inpainted patch has content.
# anything that’s still a hole after back‐projection needs inpaint
needs_inpaint = (~((orig_covered) | (~back_mask))).to(torch.float32)
# --- Compositing Logic ---
# Goal: Inpainted patch takes precedence in overlapping areas. Original content is used elsewhere.
return final_img, needs_inpaint
# Contribution from the original image:
# Valid original pixels, excluding areas covered by the inpainted patch.
original_content_contribution = base_img * base_mask.unsqueeze(-1) * \
(1.0 - inpainted_patch_coverage_mask.unsqueeze(-1))
# Contribution from the inpainted patch (reprojected as back_img):
# Valid inpainted pixels, where the patch provides coverage.
inpainted_patch_contribution = back_img * inpainted_patch_coverage_mask.unsqueeze(-1)
# Combine:
final_img = original_content_contribution + inpainted_patch_contribution
# --- needs_inpaint_mask Derivation ---
# Identifies areas that were initially holes AND remain un-filled by the reprojected inpainted patch.
# These are areas that still require inpainting if a further pass was to be made.
not_covered_by_inpainted_patch = 1.0 - inpainted_patch_coverage_mask
needs_inpaint_mask = initial_hole_mask * not_covered_by_inpainted_patch # Element-wise multiplication (AND logic)
return final_img, needs_inpaint_mask
# register
NODE_CLASS_MAPPINGS = {