Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d49f3f03ce |
+374
-166
@@ -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}
|
||||
|
||||
@@ -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 = {
|
||||
|
||||
Reference in New Issue
Block a user