refactor pointcloud projector for better hole-filling mechanism

This commit is contained in:
Alexander Kharin
2025-07-07 20:15:10 +02:00
parent d67cd61c3c
commit 696a4ad763
+68 -92
View File
@@ -335,7 +335,8 @@ class ProjectPointCloud:
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
device = pointcloud.device
coords = pointcloud[:, :3]
colors = pointcloud[:, 3:].float()
# use only RGB channels
colors = pointcloud[:, 3:6].float()
# 1) Filter points in front of the camera
mask_front = coords[:, 2] > 0
@@ -352,115 +353,90 @@ class ProjectPointCloud:
u, v, depth = XYZ_to_equirect(X, Y, Z, output_horizontal_fov)
# 3) Rasterize to pixel indices
px = (u * (output_width - 1) / 2) + (output_width - 1) / 2
py = (v * (output_height - 1) / 2) + (output_height - 1) / 2
ix = px.round().clamp(0, output_width - 1).long()
iy = py.round().clamp(0, output_height - 1).long()
pix = iy * output_width + ix
M = output_width * output_height
W, H = output_width, output_height
px = (u * (W - 1) / 2) + (W - 1) / 2
py = (v * (H - 1) / 2) + (H - 1) / 2
ix = px.round().clamp(0, W - 1).long()
iy = py.round().clamp(0, H - 1).long()
pix = iy * W + ix
M = W * H
# —— NEW: drop any invalid / NaN→int_min projections ——
# 4) Keep only valid projections
valid = (pix >= 0) & (pix < M)
depth = depth[valid]
depth = depth[valid]
colors = colors[valid]
pix = pix[valid]
order = torch.arange(depth.size(0), device=device)
# rebuild your "order" to match
pix = pix[valid]
# 4) Allocate or reuse buffers
if not hasattr(self, '_z_front') or self._z_front.numel() != M:
self._z_front = torch.empty((M,), device=device)
self._z_back = torch.empty((M,), device=device)
self._idx = torch.full((M,), -1, dtype=torch.long, device=device)
self._flat = torch.zeros((M, 4), device=device)
z_front = self._z_front
z_back = self._z_back
idxbuf = self._idx
flat = self._flat
# 5) Front z-buffer pass (nearest)
z_front.fill_(float('inf'))
# 5) Nearest z-buffer (front)
z_front = torch.full((M,), float('inf'), device=device)
z_front.scatter_reduce_(0, pix, depth, reduce='amin', include_self=True)
# select only the nearest sample per pixel
sel_front = depth == z_front[pix]
order = torch.arange(depth.size(0), device=device)
order_m = torch.where(sel_front, order, depth.size(0))
idxbuf.fill_(depth.size(0))
idxbuf.scatter_reduce_(0, pix, order_m, reduce='amin', include_self=True)
win_front = order == idxbuf[pix]
pix_sel = pix[sel_front]
col_sel = colors[sel_front]
flat.fill_(0)
flat[pix[win_front]] = colors[win_front]
img4 = flat.view(output_height, output_width, 4)
rgb = img4[..., :3].clamp(0, 255)
alpha = (img4[..., 3] > 0).float()
mask_init= (img4[..., 3] > 0) # initial mask
rgb *= alpha.unsqueeze(-1)
depth_img = z_front.view(output_height, output_width)
rgb_HR = rgb
# pack into image buffer
flat_rgb = torch.zeros((M, 3), device=device)
flat_rgb[pix_sel] = col_sel.clamp(0, 255)
rgb = flat_rgb.view(H, W, 3)
# 6) Back z-buffer pass (farthest) for hole-filling
# create initial depth & mask
depth_img = z_front.view(H, W)
mask_init = (depth_img < float('inf'))
mask = mask_init.clone()
depth_inftozero= depth_img.clone()
# replace inf with 0
depth_inftozero[depth_inftozero == float('inf')] = 0
# 6) Hole-filling + small-hole closure
mask = mask_init.clone()
if point_size > 1:
z_back.fill_(-float('inf'))
z_back.scatter_reduce_(0, pix, depth, reduce='amax', include_self=True)
sel_back = depth == z_back[pix]
order_m = torch.where(sel_back, order, -1)
idxbuf.fill_(-1)
idxbuf.scatter_reduce_(0, pix, order_m, reduce='amax', include_self=True)
win_back = idxbuf[pix] >= 0
flat.fill_(0)
flat[pix[win_back]] = colors[win_back]
back4 = flat.view(output_height, output_width, 4)
rgb_back = back4[..., :3].clamp(0,255)
alpha_back = (back4[..., 3] > 0).float()
r = point_size // 2
k = 2*r + 1
pad = r
# tensors for conv
depth_t = depth_inftozero.unsqueeze(0).unsqueeze(0) # [1,1,H,W]
mask_t = mask.float().unsqueeze(0).unsqueeze(0) # [1,1,H,W]
rgb_t = rgb.permute(2,0,1).unsqueeze(0) # [1,3,H,W]
# define kernel
kernel = torch.ones((1,1,k,k), device=device)
# fill holes where front missed
hole = (alpha == 0) & (alpha_back > 0)
rgb[hole] = rgb_back[hole]
alpha[hole] = 1.0
depth_img[hole] = z_back.view(output_height, output_width)[hole]
# average-based fill
sum_d = F.conv2d(depth_t * mask_t, kernel, padding=pad)
cnt = F.conv2d(mask_t, kernel, padding=pad).clamp(min=1.0)
kernel_rgb = kernel.repeat(3,1,1,1)
sum_c = F.conv2d(rgb_t * mask_t, kernel_rgb, padding=pad, groups=3)
avg_d = sum_d / cnt
avg_c = sum_c / cnt
cnt_map = cnt.squeeze(0).squeeze(0)
# 7) Median-filter _only_ in hole regions
if hole.any():
# prepare for kornia median_blur: [B,C,H,W]
rgb_t = rgb.permute(2,0,1).unsqueeze(0) # [1,3,H,W]
# apply median filter
rgb_med = median_blur(rgb_t, (point_size, point_size))
# back to HWC
rgb_med = rgb_med.squeeze(0).permute(1,2,0)
# merge only at hole locations
rgb[hole] = rgb_med[hole]
# alpha already set to 1.0 for holes
# 8 apply median blur to mask if point_size > 1 and to initial image
mask_t = alpha.unsqueeze(0).unsqueeze(0) # [1,1,H,W]
pad = point_size // 2
ksize = (point_size, point_size)
# morphological closing on initial mask to find small holes
mask0_t = mask_init.float().unsqueeze(0).unsqueeze(0)
grow = F.max_pool2d(mask0_t, (k,k), stride=1, padding=pad)
shrink = 1 - F.max_pool2d(1-grow, (k,k), stride=1, padding=pad)
closed = shrink.squeeze(0).squeeze(0) > 0
small_holes = closed & ~mask_init
# b) grow (dilate) mask by max‑pool
mask_grow = F.max_pool2d(mask_t, kernel_size=ksize, stride=1, padding=pad)
# fill only small, fillable holes
hole = small_holes & (cnt_map > 0)
filled_c = avg_c.squeeze(0).permute(1,2,0)
filled_d = avg_d.squeeze(0).squeeze(0)
rgb[hole] = filled_c[hole]
depth_img[hole] = filled_d[hole]
mask = closed
# c) shrink (erode) by inverting, max‑pool, then inverting back
mask_shrink = 1.0 - F.max_pool2d(1.0 - mask_grow, kernel_size=ksize, stride=1, padding=pad)
# d) back to [H,W] and use as our new alpha
alpha = mask_shrink.squeeze(0).squeeze(0)
print(1)
# e) median‑filter the *whole* RGB image
# prep for kornia: [B,C,H,W]
rgb_t_full = rgb.permute(2,0,1).unsqueeze(0) # [1,3,H,W]
rgb_med_full = median_blur(rgb_t_full, ksize) # [1,3,H,W]
rgb_med_full = rgb_med_full.squeeze(0).permute(1,2,0) # [H,W,3]
# g) refill *only* the original holes with the median result
rgb[~mask_init] = rgb_med_full[~mask_init]
# 9) Pack and return with original script shapes
img = rgb.unsqueeze(0) # [1,H,W,3]
mask_out = alpha # [H,W]
depth4 = depth_img.unsqueeze(0).unsqueeze(-1) # [1,H,W,1]
# 7) Pack outputs
img = rgb.unsqueeze(0) # [1,H,W,3]
mask_out = mask # [H,W]
depth4 = depth_img.unsqueeze(0).unsqueeze(-1) # [1,H,W,1]
if return_inverse_depth:
depth4 = 1.0 / depth4.clamp(min=1e-6)
depth4 = depth4 * mask_out.unsqueeze(0).unsqueeze(-1)
return img, mask_out, depth4
class PointCloudUnion:
"""
Combine two point clouds into one.