mirror of
https://github.com/harry7557558/spirula-studio.git
synced 2026-10-02 02:44:54 +08:00
fused per pixel loss function
This commit is contained in:
@@ -28,7 +28,7 @@ def extract_function_declarations(code):
|
||||
# Match the function name
|
||||
\s+\b\w+\b\s*
|
||||
# Match the function parameters
|
||||
\([^)]*\)
|
||||
\(.*?\)\s
|
||||
""", re.MULTILINE | re.VERBOSE | re.DOTALL)
|
||||
|
||||
matches = function_decl_pattern.findall(code)
|
||||
@@ -160,6 +160,7 @@ path = "spirulae_splat/splat/cuda/csrc/"
|
||||
generate_header(path+"SphericalHarmonics.cu", path+"SphericalHarmonics.cuh")
|
||||
generate_header(path+"BackgroundSphericalHarmonics.cu", path+"BackgroundSphericalHarmonics.cuh")
|
||||
generate_header(path+"PerSplatLoss.cu", path+"PerSplatLoss.cuh")
|
||||
generate_header(path+"PerPixelLoss.cu", path+"PerPixelLoss.cuh")
|
||||
generate_header(path+"PixelWise.cu", path+"PixelWise.cuh")
|
||||
generate_header(path+"Projection.cu", path+"Projection.cuh")
|
||||
generate_header(path+"ProjectionEval3D.cu", path+"ProjectionEval3D.cuh")
|
||||
|
||||
@@ -62,8 +62,8 @@ class SupervisionLosses(torch.nn.Module):
|
||||
|
||||
self.depth_weight = config.depth_supervision_weight
|
||||
self.normal_weight = config.normal_supervision_weight
|
||||
self.alpha_weight = config.alpha_supervision_weight
|
||||
self.alpha_weight_under = config.alpha_supervision_weight_under
|
||||
self.alpha_weight = config.alpha_loss_weight
|
||||
self.alpha_weight_under = config.alpha_loss_weight_under
|
||||
|
||||
@staticmethod
|
||||
def get_alpha_loss(x, y):
|
||||
|
||||
@@ -17,15 +17,7 @@ from fused_ssim import fused_ssim
|
||||
|
||||
from fused_bilagrid import BilateralGrid, slice, total_variation_loss
|
||||
|
||||
|
||||
def pearson_correlation_loss(y_pred, y_true):
|
||||
y_pred_flat = y_pred.flatten()
|
||||
y_true_flat = y_true.flatten()
|
||||
|
||||
stacked_data = torch.stack((y_pred_flat, y_true_flat))
|
||||
correlation_matrix = torch.corrcoef(stacked_data)
|
||||
|
||||
return 1 - correlation_matrix[0, 1]
|
||||
from typing import Optional
|
||||
|
||||
|
||||
class _MaskGradient(torch.autograd.Function):
|
||||
@@ -87,6 +79,66 @@ class _ComputePerSplatLosses(torch.autograd.Function):
|
||||
return (*v_inputs, *([None]*len(hyperparams)))
|
||||
|
||||
|
||||
class _ComputePerPixelLosses(torch.autograd.Function):
|
||||
|
||||
@staticmethod
|
||||
def forward(
|
||||
ctx,
|
||||
render_rgb: Optional[torch.Tensor],
|
||||
ref_rgb: Optional[torch.Tensor],
|
||||
render_depth: Optional[torch.Tensor],
|
||||
ref_depth: Optional[torch.Tensor],
|
||||
render_normal: Optional[torch.Tensor],
|
||||
depth_normal: Optional[torch.Tensor],
|
||||
ref_normal: Optional[torch.Tensor],
|
||||
render_alpha: Optional[torch.Tensor],
|
||||
rgb_dist: Optional[torch.Tensor],
|
||||
depth_dist: Optional[torch.Tensor],
|
||||
normal_dist: Optional[torch.Tensor],
|
||||
ref_alpha: Optional[torch.Tensor],
|
||||
mask: Optional[torch.Tensor],
|
||||
depth_mask: Optional[torch.Tensor],
|
||||
normal_mask: Optional[torch.Tensor],
|
||||
weights
|
||||
):
|
||||
|
||||
tensors = (
|
||||
render_rgb,
|
||||
ref_rgb,
|
||||
render_depth,
|
||||
ref_depth,
|
||||
render_normal,
|
||||
depth_normal,
|
||||
ref_normal,
|
||||
render_alpha,
|
||||
rgb_dist,
|
||||
depth_dist,
|
||||
normal_dist,
|
||||
ref_alpha,
|
||||
mask,
|
||||
depth_mask,
|
||||
normal_mask
|
||||
)
|
||||
|
||||
losses, raw_losses = _C.compute_per_pixel_losses_forward(
|
||||
*tensors, weights
|
||||
)
|
||||
|
||||
ctx.weights = weights
|
||||
ctx.save_for_backward(*tensors, raw_losses)
|
||||
|
||||
return losses
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, v_losses):
|
||||
grads = _C.compute_per_pixel_losses_backward(
|
||||
*ctx.saved_tensors,
|
||||
ctx.weights,
|
||||
v_losses
|
||||
)
|
||||
return *grads, *([None]*(len(ctx.needs_input_grad)-len(grads)))
|
||||
|
||||
|
||||
|
||||
class SplatTrainingLosses(torch.nn.Module):
|
||||
|
||||
@@ -180,255 +232,164 @@ class SplatTrainingLosses(torch.nn.Module):
|
||||
return self.config.alpha_reg_weight * \
|
||||
min(self.step / max(self.config.alpha_reg_warmup, 1), 1)
|
||||
|
||||
# @torch.compile(**_TORCH_COMPILE_ARGS)
|
||||
def _image_loss_0(self, gt_img, pred_img, pred_img_e, mask=None):
|
||||
pred_img_e = torch.clip(pred_img_e, 0.0, 1.0)
|
||||
pred_img = torch.clip(pred_img, 0.0, 1.0)
|
||||
if mask is None:
|
||||
Ll1_e = torch.abs(gt_img - pred_img_e).mean()
|
||||
Ll1 = torch.abs(gt_img - pred_img).mean()
|
||||
Ll2_e = ((gt_img-pred_img_e)**2).mean()
|
||||
else: # TODO: pass mask here and see how it goes
|
||||
num_channels = gt_img.shape[-1]
|
||||
inv_denom = 1.0 / torch.clamp(mask.sum() * num_channels, min=1.0)
|
||||
Ll1_e = (mask * torch.abs(gt_img - pred_img_e)).sum() * inv_denom
|
||||
Ll1 = (mask * torch.abs(gt_img - pred_img)).sum() * inv_denom
|
||||
Ll2_e = (mask * (gt_img-pred_img_e)**2).sum() * inv_denom
|
||||
|
||||
gt_img_bchw = gt_img.permute(0, 3, 1, 2).contiguous()
|
||||
pred_img_bchw = pred_img_e.permute(0, 3, 1, 2).contiguous()
|
||||
return gt_img_bchw, pred_img_bchw, Ll1_e, Ll1, Ll2_e
|
||||
|
||||
# @torch.compile(dynamic=False)
|
||||
def _image_loss_1(self, Ll1_e, Ll1, ssim, ssim_lambda, exposure_reg_image):
|
||||
return torch.lerp(torch.lerp(Ll1_e, 1-ssim, ssim_lambda), Ll1, exposure_reg_image)
|
||||
|
||||
def image_loss(self, gt_img, pred_img, pred_img_e, exposure_reg_image):
|
||||
gt_img_bchw, pred_img_bchw, Ll1_e, Ll1, Ll2_e = self._image_loss_0(gt_img, pred_img, pred_img_e)
|
||||
# ssim = fused_ssim(pred_img_bchw, gt_img_bchw, padding="valid")
|
||||
ssim = fused_ssim(pred_img_bchw, gt_img_bchw, padding="same")
|
||||
|
||||
ssim_lambda = self.config.ssim_lambda * min(self.step/max(self.config.ssim_warmup,1), 1)
|
||||
return self._image_loss_1(Ll1_e, Ll1, ssim, ssim_lambda, exposure_reg_image), Ll2_e, ssim
|
||||
|
||||
# @torch.compile(**_TORCH_COMPILE_ARGS)
|
||||
def alpha_reg(self, alpha):
|
||||
weight_alpha_reg = self.get_alpha_reg_weight()
|
||||
if self.config.randomize_background:
|
||||
reg_alpha = 1.0 - alpha**2 # push to 1
|
||||
else:
|
||||
# reg_alpha = torch.log(4.0*torch.clip(alpha*(1.0-alpha), min=1e-2))
|
||||
reg_alpha = 4.0*alpha*(1.0-alpha)
|
||||
return weight_alpha_reg * reg_alpha.mean()
|
||||
|
||||
def forward(self, step: int, batch, outputs):
|
||||
self.step = step
|
||||
|
||||
# mask out of bound (e.g. fisheye circle)
|
||||
camera_mask = None
|
||||
# TODO
|
||||
# ssplat_camera = outputs["ssplat_camera"] # type: _Camera
|
||||
# if ssplat_camera.is_distorted():
|
||||
# undist_map = ssplat_camera.get_undist_map()
|
||||
# camera_mask = torch.isfinite(undist_map.sum(-1, True))
|
||||
# if not camera_mask.all():
|
||||
# for key in ['rgb', 'depth', 'alpha', 'background']:
|
||||
# if key in outputs:
|
||||
# outputs[key] = _MaskGradient.apply(outputs[key], camera_mask)
|
||||
device = outputs['rgb'].device
|
||||
camera = outputs["camera"]
|
||||
|
||||
gt_depth_mask, gt_normal_mask = None, None
|
||||
if 'depth' in batch and len(batch['depth'].shape) == 3:
|
||||
batch['depth'] = batch['depth'].unsqueeze(-1)
|
||||
pred_rgb = outputs["rgb"]
|
||||
pred_depth = outputs["depth"] if 'depth' in outputs else None
|
||||
pred_normal = outputs["normal"] if 'normal' in outputs else None
|
||||
pred_depth_normal = outputs["depth_normal"] if 'depth_normal' in outputs else None
|
||||
pred_alpha = outputs["alpha"] if 'alpha' in outputs else None
|
||||
|
||||
gt_rgb, gt_depth, gt_normal, gt_alpha = None, None, None, None # for loss
|
||||
gt_rgb_mask, gt_depth_mask, gt_normal_mask, gt_alpha_mask = None, None, None, None # for masking
|
||||
|
||||
# load alpha
|
||||
if "mask" in batch:
|
||||
batch_mask = self._downscale_if_required(batch['mask'].to(device).float()).bool()
|
||||
gt_rgb_mask = batch_mask
|
||||
if self.config.apply_loss_for_mask:
|
||||
gt_alpha = batch_mask
|
||||
|
||||
# load depth
|
||||
if 'depth' in batch:
|
||||
batch['depth'] = batch['depth'].to(device)
|
||||
gt_depth_mask = (batch['depth'] != 0.0)
|
||||
if 'normal' in batch:
|
||||
batch['normal'] = batch['normal'].to(device)
|
||||
gt_normal_mask = (batch['normal'].sum(-1) > -2.366)
|
||||
batch['normal'] = F.normalize(batch['normal'], dim=-1)
|
||||
gt_depth = self._downscale_if_required(batch['depth'].to(device))
|
||||
if len(gt_depth.shape) == 3:
|
||||
gt_depth = gt_depth.unsqueeze(-1)
|
||||
gt_depth_mask = (gt_depth != 0.0)
|
||||
|
||||
# apply bilateral grid
|
||||
gt_depth = batch.get("depth", None)
|
||||
gt_normal = batch.get("normal", None)
|
||||
if self.config.use_bilateral_grid_for_geometry:
|
||||
camera = outputs["camera"]
|
||||
if camera.metadata is not None and "cam_idx" in camera.metadata:
|
||||
# mask sky
|
||||
none_sky_mask = gt_depth < torch.amax(
|
||||
gt_depth, dim=(1,2,3), keepdims=True).detach().item()
|
||||
gt_depth_mask = gt_depth_mask & none_sky_mask
|
||||
if gt_alpha is not None:
|
||||
gt_alpha = gt_alpha & none_sky_mask
|
||||
else:
|
||||
gt_alpha = none_sky_mask
|
||||
|
||||
# apply bilagrid
|
||||
if self.config.use_bilateral_grid_for_geometry and \
|
||||
(camera.metadata is not None and "cam_idx" in camera.metadata):
|
||||
# TODO: fused kernel
|
||||
# TODO: might not be the best way to use RGB bilagrid
|
||||
if gt_normal is not None:
|
||||
B, H, W, C = gt_normal.shape
|
||||
gt_normal = self.apply_bilateral_grid(
|
||||
self.bil_grids_normal,
|
||||
0.5+0.5*gt_normal, camera.metadata["cam_idx"], H, W
|
||||
) * 2.0 - 1.0
|
||||
gt_normal = F.normalize(gt_normal, dim=-1)
|
||||
if gt_depth is not None:
|
||||
B, H, W, C = gt_depth.shape
|
||||
gt_depth = gt_depth / torch.mean(gt_depth, dim=(1, 2, 3)) # TODO: median might be better
|
||||
gt_depth = gt_depth / (gt_depth + 1.0)
|
||||
gt_depth = self.apply_bilateral_grid(
|
||||
self.bil_grids_depth,
|
||||
gt_depth.repeat(1, 1, 1, 3), camera.metadata["cam_idx"], H, W
|
||||
)[..., :1]
|
||||
gt_depth = gt_depth / (1.0 - gt_depth).clip(max=0.999)
|
||||
B, H, W, C = gt_depth.shape
|
||||
gt_depth = gt_depth * (
|
||||
gt_depth_mask.float().sum(dim=(1, 2, 3)) /
|
||||
(gt_depth * gt_depth_mask.float()).sum(dim=(1, 2, 3)) # TODO: fix zero division
|
||||
)
|
||||
gt_depth = gt_depth / (gt_depth + 1.0)
|
||||
gt_depth = self.apply_bilateral_grid(
|
||||
self.bil_grids_depth,
|
||||
gt_depth.repeat(1, 1, 1, 3), camera.metadata["cam_idx"], H, W
|
||||
)[..., :1]
|
||||
gt_depth = gt_depth / (1.0 - gt_depth).clip(max=0.999)
|
||||
|
||||
# load normal
|
||||
if 'normal' in batch:
|
||||
gt_normal = self._downscale_if_required(batch['normal'].to(device))
|
||||
gt_normal_mask = (gt_normal.sum(-1, True) > -2.366) # background is (-1, -1, -1)
|
||||
|
||||
# apply bilagrid
|
||||
if self.config.use_bilateral_grid_for_geometry and \
|
||||
(camera.metadata is not None and "cam_idx" in camera.metadata):
|
||||
# TODO: fused kernel
|
||||
# TODO: might not be the best way to use RGB bilagrid
|
||||
B, H, W, C = gt_normal.shape
|
||||
gt_normal = self.apply_bilateral_grid(
|
||||
self.bil_grids_normal,
|
||||
0.5+0.5*gt_normal, camera.metadata["cam_idx"], H, W
|
||||
) * 2.0 - 1.0
|
||||
|
||||
# load RGB
|
||||
if self.config.fit == "rgb":
|
||||
gt_img_rgba = self.get_gt_img(batch["image"].to(device))
|
||||
elif self.config.fit == "depth":
|
||||
gt_img_rgba = self.get_gt_img(gt_depth)
|
||||
gt_img_rgba = gt_img_rgba / gt_img_rgba.mean()
|
||||
gt_img_rgba = (gt_img_rgba / (1.0 + gt_img_rgba)).repeat(1, 1, 1, 3)
|
||||
elif self.config.fit in ["normal", "depth_normal"]:
|
||||
gt_img_rgba = 0.5+0.5*self.get_gt_img(gt_normal)
|
||||
gt_img = self.composite_with_background(gt_img_rgba, outputs["background"])
|
||||
pred_img = outputs["rgb"]
|
||||
gt_img_rgba = gt_depth
|
||||
elif self.config.fit in ["normal"]:
|
||||
gt_img_rgba = 0.5+0.5*F.normalize(gt_normal, dim=-1)
|
||||
gt_rgb = self.composite_with_background(gt_img_rgba, outputs["background"])
|
||||
|
||||
# alpha channel for bounded objects - apply a cost on rendered alpha
|
||||
alpha_loss = 0.0
|
||||
# update alpha if image is RGBA
|
||||
if gt_img_rgba.shape[-1] == 4 and self.config.alpha_loss_weight > 0.0:
|
||||
alpha = gt_img_rgba[..., -1].unsqueeze(-1)
|
||||
alpha_loss = alpha_loss + SupervisionLosses.get_alpha_loss(outputs['alpha'], alpha)
|
||||
gt_rgb_mask = gt_rgb_mask & alpha if gt_rgb_mask is None else alpha
|
||||
if self.config.apply_loss_for_mask:
|
||||
gt_alpha = gt_alpha & alpha if gt_rgb_mask is None else alpha
|
||||
|
||||
# separate mask for dynamic objects, text, etc.
|
||||
# simply don't consider it when evaluating loss, unless theres's no alpha channel, where a cost is applied
|
||||
# do this to make SSIM happier
|
||||
mask = None
|
||||
if "mask" in batch:
|
||||
# batch["mask"] : [H, W, 1]
|
||||
mask = self._downscale_if_required(batch["mask"])
|
||||
mask = mask.float().to(gt_img.device)
|
||||
assert mask.shape[:-1] == gt_img.shape[:-1] == pred_img.shape[:-1]
|
||||
# can be little bit sketchy for the SSIM loss
|
||||
gt_img = torch.lerp(outputs["background"], gt_img, mask)
|
||||
pred_img = torch.lerp(outputs["background"], pred_img, mask)
|
||||
|
||||
# If alpha channel is not specified, apply loss
|
||||
if isinstance(alpha_loss, float) and alpha_loss == 0.0 and self.config.alpha_loss_weight > 0.0:
|
||||
alpha_loss = alpha_loss + SupervisionLosses.get_alpha_loss(outputs['alpha'], mask)
|
||||
if gt_rgb_mask is not None:
|
||||
gt_rgb = torch.where(gt_rgb_mask, gt_rgb, outputs["background"])
|
||||
pred_rgb = torch.where(gt_rgb_mask, pred_rgb, outputs["background"])
|
||||
|
||||
alpha_loss = self.config.alpha_loss_weight * alpha_loss
|
||||
|
||||
# depth supervision
|
||||
depth_supervision_loss, normal_supervision_loss, alpha_supervision_loss = 0.0, 0.0, 0.0
|
||||
normal_reg = 0.0
|
||||
(weight_depth_dist_reg, weight_normal_dist_reg, weight_rgb_dist_reg), weight_normal_reg = \
|
||||
self.get_2dgs_reg_weights()
|
||||
if "depth" in batch and 'depth' in outputs \
|
||||
and self.step > self.config.supervision_warmup \
|
||||
and (self.config.depth_supervision_weight > 0.0 or \
|
||||
self.config.alpha_supervision_weight > 0.0 or \
|
||||
self.config.alpha_supervision_weight_under > 0.0
|
||||
):
|
||||
if gt_depth.ndim == 3:
|
||||
gt_depth = gt_depth.unsqueeze(-1)
|
||||
batch_depth = self._downscale_if_required(gt_depth.to(device))
|
||||
|
||||
if self.config.depth_supervision_weight > 0.0 or \
|
||||
self.config.alpha_supervision_weight > 0.0 or \
|
||||
self.config.alpha_supervision_weight_under > 0.0:
|
||||
# This works for Metric3D depth, not every model
|
||||
batch_depth_original = batch_depth
|
||||
if self.config.use_bilateral_grid_for_geometry:
|
||||
batch_depth_original = self._downscale_if_required(batch["depth"].to(device))
|
||||
batch_alpha = (batch_depth_original < torch.amax(
|
||||
batch_depth_original, dim=(1,2,3), keepdims=True).detach().item())
|
||||
|
||||
if self.config.depth_supervision_weight > 0.0:
|
||||
output_depth = ray_depth_to_linear_depth(outputs["depth"], **outputs["camera_intrins"])
|
||||
|
||||
def normalize_depth(d):
|
||||
d = torch.log(d.clip(min=1e-4))
|
||||
mean_squared = (d*d * batch_alpha).sum((1, 2, 3)) / batch_alpha.sum((1, 2, 3))
|
||||
mean = (d * batch_alpha).sum((1, 2, 3)) / batch_alpha.sum((1, 2, 3))
|
||||
std = torch.sqrt((mean_squared - mean*mean).clip(min=1e-8))
|
||||
return (d - mean.reshape(1, 1, 1, -1)) / std.reshape(1, 1, 1, -1)
|
||||
|
||||
# batch_alpha = batch_alpha.float()
|
||||
# batch_depth_n = normalize_depth(batch_depth)
|
||||
# output_depth_n = normalize_depth(output_depth)
|
||||
# depth_supervision_loss = self.config.depth_supervision_weight * \
|
||||
# (torch.sum(batch_alpha * (batch_depth_n - output_depth_n)**2) / \
|
||||
# batch_alpha.sum()) ** 0.5
|
||||
depth_supervision_loss = self.config.depth_supervision_weight * \
|
||||
pearson_correlation_loss(
|
||||
torch.log(batch_depth[batch_alpha & gt_depth_mask].clip(min=1e-4)),
|
||||
torch.log(output_depth[batch_alpha & gt_depth_mask].clip(min=1e-4))
|
||||
)
|
||||
|
||||
if self.config.alpha_supervision_weight > 0.0 or \
|
||||
self.config.alpha_supervision_weight_under > 0.0:
|
||||
|
||||
def alpha_loss_fun(x, y):
|
||||
return F.binary_cross_entropy(torch.fmax(x, y), y, reduction="mean")
|
||||
|
||||
batch_alpha = batch_alpha.float()
|
||||
alpha_supervision_loss = self.config.alpha_supervision_weight * \
|
||||
alpha_loss_fun(outputs["alpha"][gt_depth_mask], batch_alpha[gt_depth_mask])
|
||||
if self.config.alpha_supervision_weight_under > 0.0:
|
||||
alpha_supervision_loss = alpha_supervision_loss + self.config.alpha_supervision_weight_under *\
|
||||
alpha_loss_fun(1.0-outputs["alpha"][gt_depth_mask], 1.0-batch_alpha[gt_depth_mask])
|
||||
|
||||
# normal supervision
|
||||
if self.step > self.config.supervision_warmup and (
|
||||
self.config.normal_supervision_weight > 0.0 or
|
||||
weight_normal_reg > 0.0) and \
|
||||
("normal" in batch or 'normal' in outputs or 'depth_normal' in outputs):
|
||||
if 'normal' in batch:
|
||||
batch_normal = self._downscale_if_required(gt_normal.to(device))
|
||||
weight = 0
|
||||
def normal_loss(x, y):
|
||||
mask = (torch.norm(x, dim=-1) > 0.5) & (torch.norm(y, dim=-1) > 0.5)
|
||||
# return 1.0 - (x*y).sum(-1)[mask].mean()
|
||||
return torch.abs(x-y)[mask].mean() # TODO: why this works much better
|
||||
# with reference image
|
||||
if 'normal' in outputs and 'normal' in batch:
|
||||
normal_supervision_loss = normal_supervision_loss + \
|
||||
normal_loss(batch_normal[gt_normal_mask], outputs["normal"][gt_normal_mask])
|
||||
weight += 1
|
||||
if 'depth_normal' in outputs and 'normal' in batch:
|
||||
normal_supervision_loss = normal_supervision_loss + \
|
||||
normal_loss(batch_normal[gt_normal_mask], outputs["depth_normal"][gt_normal_mask])
|
||||
weight += 1
|
||||
if weight > 0:
|
||||
normal_supervision_loss = normal_supervision_loss * \
|
||||
(self.config.normal_supervision_weight / weight)
|
||||
# with self (2DGS)
|
||||
if 'normal' in outputs and 'depth_normal' in outputs and \
|
||||
self.step >= self.config.reg_warmup_length:
|
||||
normal_reg = weight_normal_reg * \
|
||||
normal_loss(outputs["normal"], outputs["depth_normal"])
|
||||
|
||||
# distortion regularization
|
||||
depth_dist_reg, normal_dist_reg, rgb_dist_reg = 0.0, 0.0, 0.0
|
||||
if self.step >= self.config.reg_warmup_length:
|
||||
if weight_depth_dist_reg > 0.0 and 'depth_distortion' in outputs:
|
||||
depth_dist_reg = weight_depth_dist_reg * outputs['depth_distortion'].mean()
|
||||
if weight_normal_dist_reg > 0.0 and 'normal_distortion' in outputs:
|
||||
normal_dist_reg = weight_normal_dist_reg * outputs['normal_distortion'].mean()
|
||||
if weight_rgb_dist_reg > 0.0 and 'rgb_distortion' in outputs:
|
||||
rgb_dist_reg = weight_rgb_dist_reg * outputs['rgb_distortion'].mean()
|
||||
|
||||
# correct exposure
|
||||
pred_img_e = pred_img
|
||||
exposure_param_reg = 0.0
|
||||
exposure_reg_image = 0.0
|
||||
# correct exposure (deprecated)
|
||||
if self.config.adaptive_exposure_mode is not None and \
|
||||
self.step > self.config.adaptive_exposure_warmup:
|
||||
exposure_reg_image = self.config.exposure_reg_image
|
||||
# mask, note that alpha_mask is defined in "mask out of bound" step
|
||||
alpha_mask = camera_mask
|
||||
if mask is not None and alpha_mask is not None:
|
||||
alpha_mask = alpha_mask & mask
|
||||
# call function
|
||||
pred_img_e, exposure_param_reg = self.exposure_correction(pred_img, gt_img, alpha_mask)
|
||||
raise NotImplementedError("Adaptive exposure is deprecated. Use bilateral grid instead.")
|
||||
|
||||
# image loss
|
||||
image_loss, mse, ssim = self.image_loss(gt_img, pred_img, pred_img_e, exposure_reg_image)
|
||||
# ssim loss
|
||||
ssim = fused_ssim(
|
||||
pred_rgb.permute(0, 3, 1, 2).contiguous(),
|
||||
gt_rgb.permute(0, 3, 1, 2).contiguous(),
|
||||
padding="same",
|
||||
train=True
|
||||
)
|
||||
|
||||
# alpha regularizer
|
||||
alpha_reg = 0.0
|
||||
if self.step >= self.config.reg_warmup_length:
|
||||
alpha = outputs['alpha']
|
||||
alpha_reg = self.alpha_reg(alpha)
|
||||
# call fused kernel to compute loss
|
||||
|
||||
(weight_depth_dist_reg, weight_normal_dist_reg, weight_rgb_dist_reg), weight_normal_reg = \
|
||||
self.get_2dgs_reg_weights()
|
||||
|
||||
losses = _ComputePerPixelLosses.apply(
|
||||
pred_rgb,
|
||||
gt_rgb,
|
||||
pred_depth,
|
||||
gt_depth,
|
||||
pred_normal,
|
||||
pred_depth_normal,
|
||||
gt_normal,
|
||||
pred_alpha,
|
||||
outputs['rgb_distortion'] if 'rgb_distortion' in outputs else None,
|
||||
outputs['depth_distortion'] if 'depth_distortion' in outputs else None,
|
||||
outputs['normal_distortion'] if 'normal_distortion' in outputs else None,
|
||||
gt_alpha,
|
||||
gt_rgb_mask,
|
||||
gt_depth_mask,
|
||||
gt_normal_mask,
|
||||
# gt_alpha_mask,
|
||||
[
|
||||
# RGB supervision
|
||||
1.0 - self.config.ssim_lambda,
|
||||
# depth supervison
|
||||
float(self.step > self.config.supervision_warmup) *
|
||||
self.config.depth_supervision_weight,
|
||||
# normal supervision
|
||||
float(self.step > self.config.supervision_warmup) *
|
||||
self.config.normal_supervision_weight,
|
||||
# alpha supervision (over and under)
|
||||
self.config.alpha_loss_weight,
|
||||
self.config.alpha_loss_weight_under,
|
||||
# normal regularization
|
||||
float(self.step > self.config.reg_warmup_length) *
|
||||
weight_normal_reg,
|
||||
# alpha regularization
|
||||
float(self.step >= self.config.reg_warmup_length) *
|
||||
self.get_alpha_reg_weight(),
|
||||
# distortion regularizations (RGB, depth, normal)
|
||||
float(self.step >= self.config.reg_warmup_length) * weight_rgb_dist_reg,
|
||||
float(self.step >= self.config.reg_warmup_length) * weight_depth_dist_reg,
|
||||
float(self.step >= self.config.reg_warmup_length) * weight_normal_dist_reg,
|
||||
]
|
||||
)
|
||||
(
|
||||
rgb_l1, rgb_psnr,
|
||||
depth_supervision_loss, normal_supervision_loss, alpha_supervision_loss,
|
||||
normal_reg, alpha_reg,
|
||||
rgb_dist_reg, depth_dist_reg, normal_dist_reg
|
||||
) = losses
|
||||
|
||||
# metrics, readable from console during training
|
||||
with torch.no_grad():
|
||||
@@ -436,7 +397,7 @@ class SplatTrainingLosses(torch.nn.Module):
|
||||
self._running_metrics = { 'psnr': [], 'ssim': [] }
|
||||
psnr_list = self._running_metrics['psnr']
|
||||
ssim_list = self._running_metrics['ssim']
|
||||
psnr = -10.0 * math.log10(mse.item())
|
||||
psnr = rgb_psnr.item()
|
||||
ssim = ssim.item()
|
||||
psnr_list.append(psnr)
|
||||
ssim_list.append(ssim)
|
||||
@@ -448,8 +409,7 @@ class SplatTrainingLosses(torch.nn.Module):
|
||||
|
||||
loss_dict = {
|
||||
# [C] RGB and alpha
|
||||
"image_loss": image_loss,
|
||||
"alpha_loss": alpha_loss,
|
||||
"image_loss": rgb_l1 + self.config.ssim_lambda * (1.0 - ssim),
|
||||
"psnr": float(psnr),
|
||||
"ssim": float(ssim),
|
||||
# [S] supervision
|
||||
@@ -464,7 +424,7 @@ class SplatTrainingLosses(torch.nn.Module):
|
||||
"rgb_dist_reg": rgb_dist_reg,
|
||||
# [E] exposure
|
||||
"tv_loss": 0.0, # see get_per_splat_losses()
|
||||
"exposure_param_reg": exposure_param_reg,
|
||||
# "exposure_param_reg": exposure_param_reg,
|
||||
}
|
||||
|
||||
return loss_dict
|
||||
|
||||
@@ -197,7 +197,10 @@ class SpirulaeDataset(InputDataset):
|
||||
background = torch.ones_like(data['image']) * torch.tensor([(0,0,1), (0,0,255)][image_type == 'uint8']).to(data['image'])
|
||||
data['image'] = torch.where(data['mask'], data['image'], background)
|
||||
data['image'] = resize_image(data['image'][None], 2**int(max(0.5*math.log2(data['image'].numel()/10000), 0.0)))[0]
|
||||
return data
|
||||
return {
|
||||
'image_idx': data['image_idx'],
|
||||
'image': data['image'],
|
||||
}
|
||||
with ThreadPoolExecutor() as executor:
|
||||
for result in tqdm(
|
||||
executor.map(load_data, range(len(self))),
|
||||
|
||||
+33
-37
@@ -148,8 +148,6 @@ class SpirulaeModelConfig(ModelConfig):
|
||||
"""Position standard deviation to initialize random gaussians"""
|
||||
ssim_lambda: float = 0.4
|
||||
"""weight of ssim loss; 0.2 for optimal PSNR, higher for better visual quality"""
|
||||
ssim_warmup: int = 0
|
||||
"""warmup of ssim loss"""
|
||||
use_camera_optimizer: bool = False
|
||||
"""Whether to use camera optimizer
|
||||
Note: this only works well in patch batching mode"""
|
||||
@@ -285,11 +283,13 @@ class SpirulaeModelConfig(ModelConfig):
|
||||
reg_warmup_length: int = 0
|
||||
"""Warmup steps for depth, normal, and alpha regularizers.
|
||||
only apply regularizers after this many steps."""
|
||||
apply_loss_for_mask: bool = True
|
||||
"""Set this to False to use masks to ignore distractors (e.g. people and cars, area outside fisheye circle, over exposure)
|
||||
Set this to True to remove background (e.g. sky, background outside centered object)"""
|
||||
alpha_loss_weight: float = 0.01
|
||||
"""Weight for alpha, if mask is provided.
|
||||
Set this to 0.0 to use masks to ignore distractors (e.g. people and cars, area outside fisheye circle, over exposure)
|
||||
Set this to a positive value to remove background (e.g. sky, background around centered object)
|
||||
See scripts/SAM2-GUI for what I use to generate masks"""
|
||||
"""Loss weight for alpha, applies when rendered alpha is above reference alpha"""
|
||||
alpha_loss_weight_under: float = 0.005
|
||||
"""Loss weight for alpha, applies when rendered alpha is below reference alpha"""
|
||||
mcmc_opacity_reg: float = 0.01 # 0.01 in original paper
|
||||
"""Opacity regularization from MCMC
|
||||
Lower usually gives more accurate geometry"""
|
||||
@@ -325,12 +325,6 @@ class SpirulaeModelConfig(ModelConfig):
|
||||
"""Weight for depth supervision by comparing rendered depth with depth predicted by a foundation model"""
|
||||
normal_supervision_weight: float = 0.01
|
||||
"""Weight for normal supervision by comparing normal from rendered depth with normal from depth predicted by a foundation model"""
|
||||
alpha_supervision_weight: float = 0.01
|
||||
"""Weight for alpha supervision by rendered alpha with alpha predicted by a foundation model
|
||||
Useful for removing floaters from sky for outdoor scenes"""
|
||||
alpha_supervision_weight_under: float = 0.005
|
||||
"""Similar to alpha_supervision_weight, but applies when renderer opacity is lower than reference opacity"""
|
||||
|
||||
|
||||
class SpirulaeModel(Model):
|
||||
"""Template Model."""
|
||||
@@ -972,7 +966,6 @@ class SpirulaeModel(Model):
|
||||
if key in meta:
|
||||
value = meta[key]
|
||||
if value is not None:
|
||||
value = torch.where(alpha > 0.0, value / alpha, value)
|
||||
if not self.training:
|
||||
value = torch.sqrt(value + (1/255)**2) - (1/255)
|
||||
outputs[key] = value
|
||||
@@ -1076,39 +1069,42 @@ class SpirulaeModel(Model):
|
||||
if _max_vals[key] == 0.0:
|
||||
return '~'
|
||||
|
||||
if decimals is None:
|
||||
decimals = int(max(-math.log10(0.001*_max_vals[key]), 0))
|
||||
if decimals is None: # 3 sig figs
|
||||
decimals = int(max(-math.log10(0.001*abs(l)), 0)) if l != 0.0 else 0
|
||||
s = f"{{:.{decimals}f}}".format(l)
|
||||
if s.startswith('0.'):
|
||||
s = s[1:]
|
||||
return s
|
||||
|
||||
mcmc_reg = (self.config.use_mcmc and self.step < self.config.stop_refine_at)
|
||||
opacity_floor = f" {self.strategy.get_opacity_floor(self.step):.2f}".replace('0.', 'o.') \
|
||||
opacity_floor = f"[OpacFloor] {self.strategy.get_opacity_floor(self.step):.3f}".replace('0.', '.') \
|
||||
if self.config.primitive == "opaque_triangle" else ""
|
||||
chunks = [
|
||||
f"[N] {len(self.opacities)} {mem_stats}" + opacity_floor,
|
||||
f"[C] {fmt('image_loss', 1.0)} "
|
||||
f"{fmt('alpha_loss', self.config.alpha_loss_weight)} "
|
||||
f"{fmt('psnr', 1.0, 2)} "
|
||||
f"{fmt('ssim', 1.0, 3)}",
|
||||
f"[S] {fmt('depth_ref_loss', self.config.depth_supervision_weight)} "
|
||||
f"{fmt('normal_ref_loss', self.config.normal_supervision_weight)} "
|
||||
f"{fmt('alpha_ref_loss', self.config.alpha_supervision_weight)}",
|
||||
f"[G] {fmt('normal_reg', self.training_losses.get_2dgs_reg_weights()[1])} "
|
||||
f"{fmt('alpha_reg', self.training_losses.get_alpha_reg_weight())} "
|
||||
f"{fmt('depth_dist_reg', self.training_losses.get_2dgs_reg_weights()[0][0])} "
|
||||
f"{fmt('normal_dist_reg', self.training_losses.get_2dgs_reg_weights()[0][1])} "
|
||||
f"{fmt('rgb_dist_reg', self.training_losses.get_2dgs_reg_weights()[0][2])}",
|
||||
f"[M] {fmt('mcmc_opacity_reg', self.config.mcmc_opacity_reg * mcmc_reg)} "
|
||||
f"{fmt('mcmc_scale_reg', self.config.mcmc_scale_reg * mcmc_reg)}",
|
||||
f"[R] {fmt('erank_reg', max(self.config.erank_reg_s3, self.config.erank_reg))} "
|
||||
f"{fmt('scale_reg', self.config.scale_regularization_weight)}",
|
||||
f"[E] {fmt('tv_loss', 10.0)} "
|
||||
f"{fmt('exposure_param_reg', self.config.exposure_reg_param)}",
|
||||
f"[N] {len(self.opacities)}",
|
||||
f"[Mem] {mem_stats}",
|
||||
f"[Train] loss={fmt('image_loss', 1.0)} "
|
||||
f"psnr={fmt('psnr', 1.0, 2)} "
|
||||
f"ssim={fmt('ssim', 1.0, 3)}",
|
||||
" \n",
|
||||
f"[RefLoss] depth={fmt('depth_ref_loss', self.config.depth_supervision_weight, 3)} "
|
||||
f"normal={fmt('normal_ref_loss', self.config.normal_supervision_weight, 3)} "
|
||||
f"alpha={fmt('alpha_ref_loss', 0.5*(self.config.alpha_loss_weight+self.config.alpha_loss_weight_under), 4)}",
|
||||
f"[DistLoss] depth={fmt('depth_dist_reg', self.training_losses.get_2dgs_reg_weights()[0][0], 3)} "
|
||||
f"normal={fmt('normal_dist_reg', self.training_losses.get_2dgs_reg_weights()[0][1], 3)} "
|
||||
f"rgb={fmt('rgb_dist_reg', self.training_losses.get_2dgs_reg_weights()[0][2], 3)}",
|
||||
" \n",
|
||||
f"[ImReg] normal={fmt('normal_reg', self.training_losses.get_2dgs_reg_weights()[1], 3)} "
|
||||
f"alpha={fmt('alpha_reg', self.training_losses.get_alpha_reg_weight(), 3)}",
|
||||
f"[SplatReg] opac={fmt('mcmc_opacity_reg', self.config.mcmc_opacity_reg * mcmc_reg, 3)} "
|
||||
f"scale={fmt('mcmc_scale_reg', self.config.mcmc_scale_reg * mcmc_reg, 4)} "
|
||||
f"erank={fmt('erank_reg', max(self.config.erank_reg_s3, self.config.erank_reg, 3))} "
|
||||
f"aniso={fmt('scale_reg', self.config.scale_regularization_weight, 3)}",
|
||||
" \n",
|
||||
f"[Reg] bilagrid={fmt('tv_loss', 10.0)}",
|
||||
] + [opacity_floor] * (len(opacity_floor) > 0) + [
|
||||
" \n",
|
||||
]
|
||||
chunks = [c for c in chunks if any(char.isdigit() for char in c)]
|
||||
CONSOLE.print(' '.join(chunks).replace('\n', '') + " ", end="\r")
|
||||
CONSOLE.print(' '.join(chunks).replace('\n ', '\n'), end="\033[F"*4)
|
||||
|
||||
@torch.no_grad()
|
||||
def get_outputs_for_camera(self, camera: Cameras, obb_box: Optional[OrientedBox] = None) -> Dict[str, torch.Tensor]:
|
||||
|
||||
@@ -0,0 +1,400 @@
|
||||
#include "PerPixelLoss.cuh"
|
||||
|
||||
#include <cooperative_groups.h>
|
||||
#include <cooperative_groups/reduce.h>
|
||||
namespace cg = cooperative_groups;
|
||||
|
||||
#define TensorView _Slang_TensorView
|
||||
#include "generated/slang_all.cu"
|
||||
#undef TensorView
|
||||
|
||||
#include "common.cuh"
|
||||
|
||||
|
||||
__global__ void per_pixel_losses_forward_kernel(
|
||||
const size_t num_pixels,
|
||||
const float3* __restrict__ render_rgb,
|
||||
const float3* __restrict__ ref_rgb,
|
||||
const float* __restrict__ render_depth,
|
||||
const float* __restrict__ ref_depth,
|
||||
const float3* __restrict__ render_normal,
|
||||
const float3* __restrict__ depth_normal,
|
||||
const float3* __restrict__ ref_normal,
|
||||
const float* __restrict__ render_alpha,
|
||||
const float3* __restrict__ rgb_dist,
|
||||
const float* __restrict__ depth_dist,
|
||||
const float3* __restrict__ normal_dist,
|
||||
const bool* __restrict__ ref_alpha,
|
||||
const bool* __restrict__ mask,
|
||||
const bool* __restrict__ depth_mask,
|
||||
const bool* __restrict__ normal_mask,
|
||||
float* __restrict__ out_losses
|
||||
) {
|
||||
size_t idx = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
|
||||
FixedArray<float, (uint)RawLossIndex::length> losses;
|
||||
|
||||
bool inside = idx < num_pixels;
|
||||
if (inside) {
|
||||
per_pixel_losses(
|
||||
render_rgb ? render_rgb[idx] : make_float3(0),
|
||||
ref_rgb ? ref_rgb[idx] : make_float3(0),
|
||||
render_depth ? render_depth[idx] : 1.f,
|
||||
ref_depth ? ref_depth[idx] : 1.f,
|
||||
render_normal ? render_normal[idx] : make_float3(0),
|
||||
depth_normal ? depth_normal[idx] : make_float3(0),
|
||||
ref_normal ? ref_normal[idx] : make_float3(0),
|
||||
render_alpha ? render_alpha[idx] : 0.f,
|
||||
rgb_dist ? rgb_dist[idx] : make_float3(0),
|
||||
depth_dist ? depth_dist[idx] : 0.f,
|
||||
normal_dist ? normal_dist[idx] : make_float3(0),
|
||||
ref_alpha ? ref_alpha[idx] : true,
|
||||
mask ? mask[idx] : true,
|
||||
depth_mask ? depth_mask[idx] : true,
|
||||
normal_mask ? normal_mask[idx] : true,
|
||||
&losses
|
||||
);
|
||||
}
|
||||
|
||||
auto block = cg::this_thread_block();
|
||||
cg::thread_block_tile<WARP_SIZE> warp = cg::tiled_partition<WARP_SIZE>(block);
|
||||
uint warp_idx = block.thread_rank() / WARP_SIZE;
|
||||
|
||||
__shared__ float atomic_reduce[WARP_SIZE];
|
||||
|
||||
for (uint i = 0; i < (uint)RawLossIndex::length; i++) {
|
||||
float loss = inside ? losses[i] : 0.0f;
|
||||
float loss_reduced = cg::reduce(warp, loss, cg::plus<float>());
|
||||
if (warp.thread_rank() == 0)
|
||||
atomic_reduce[warp_idx] = loss_reduced;
|
||||
__syncthreads();
|
||||
loss = (warp_idx == 0) ? atomic_reduce[warp.thread_rank()] : 0.0f;
|
||||
if (__ballot_sync(~0u, loss != 0.0f) == 0)
|
||||
continue;
|
||||
loss_reduced = cg::reduce(warp, loss, cg::plus<float>());
|
||||
if (block.thread_rank() == 0 && loss_reduced != 0.0f)
|
||||
atomicAdd(out_losses+i, loss_reduced);
|
||||
}
|
||||
}
|
||||
|
||||
__global__ void per_pixel_losses_backward_kernel(
|
||||
const size_t num_pixels,
|
||||
const float3* __restrict__ render_rgb,
|
||||
const float3* __restrict__ ref_rgb,
|
||||
const float* __restrict__ render_depth,
|
||||
const float* __restrict__ ref_depth,
|
||||
const float3* __restrict__ render_normal,
|
||||
const float3* __restrict__ depth_normal,
|
||||
const float3* __restrict__ ref_normal,
|
||||
const float* __restrict__ render_alpha,
|
||||
const float3* __restrict__ rgb_dist,
|
||||
const float* __restrict__ depth_dist,
|
||||
const float3* __restrict__ normal_dist,
|
||||
const bool* __restrict__ ref_alpha,
|
||||
const bool* __restrict__ mask,
|
||||
const bool* __restrict__ depth_mask,
|
||||
const bool* __restrict__ normal_mask,
|
||||
const float* __restrict__ v_out_losses,
|
||||
float3* __restrict__ v_render_rgb,
|
||||
float3* __restrict__ v_ref_rgb,
|
||||
float* __restrict__ v_render_depth,
|
||||
float* __restrict__ v_ref_depth,
|
||||
float3* __restrict__ v_render_normal,
|
||||
float3* __restrict__ v_depth_normal,
|
||||
float3* __restrict__ v_ref_normal,
|
||||
float* __restrict__ v_render_alpha,
|
||||
float3* __restrict__ v_rgb_dist,
|
||||
float* __restrict__ v_depth_dist,
|
||||
float3* __restrict__ v_normal_dist
|
||||
) {
|
||||
size_t idx = blockIdx.x * blockDim.x + threadIdx.x;
|
||||
|
||||
bool inside = idx < num_pixels;
|
||||
if (!inside) return;
|
||||
|
||||
FixedArray<float, (uint)RawLossIndex::length> v_losses;
|
||||
for (uint i = 0; i < (uint)RawLossIndex::length; i++)
|
||||
v_losses[i] = v_out_losses[i];
|
||||
|
||||
float3 temp_v_render_rgb;
|
||||
float3 temp_v_ref_rgb;
|
||||
float temp_v_render_depth;
|
||||
float temp_v_ref_depth;
|
||||
float3 temp_v_render_normal;
|
||||
float3 temp_v_depth_normal;
|
||||
float3 temp_v_ref_normal;
|
||||
float temp_v_render_alpha;
|
||||
float3 temp_v_rgb_dist;
|
||||
float temp_v_depth_dist;
|
||||
float3 temp_v_normal_dist;
|
||||
|
||||
per_pixel_losses_bwd(
|
||||
render_rgb ? render_rgb[idx] : make_float3(0),
|
||||
ref_rgb ? ref_rgb[idx] : make_float3(0),
|
||||
render_depth ? render_depth[idx] : 1.f,
|
||||
ref_depth ? ref_depth[idx] : 1.f,
|
||||
render_normal ? render_normal[idx] : make_float3(0),
|
||||
depth_normal ? depth_normal[idx] : make_float3(0),
|
||||
ref_normal ? ref_normal[idx] : make_float3(0),
|
||||
render_alpha ? render_alpha[idx] : 0.f,
|
||||
rgb_dist ? rgb_dist[idx] : make_float3(0),
|
||||
depth_dist ? depth_dist[idx] : 0.f,
|
||||
normal_dist ? normal_dist[idx] : make_float3(0),
|
||||
ref_alpha ? ref_alpha[idx] : true,
|
||||
mask ? mask[idx] : true,
|
||||
depth_mask ? depth_mask[idx] : true,
|
||||
normal_mask ? normal_mask[idx] : true,
|
||||
&v_losses,
|
||||
&temp_v_render_rgb,
|
||||
&temp_v_ref_rgb,
|
||||
&temp_v_render_depth,
|
||||
&temp_v_ref_depth,
|
||||
&temp_v_render_normal,
|
||||
&temp_v_depth_normal,
|
||||
&temp_v_ref_normal,
|
||||
&temp_v_render_alpha,
|
||||
&temp_v_rgb_dist,
|
||||
&temp_v_depth_dist,
|
||||
&temp_v_normal_dist
|
||||
);
|
||||
|
||||
if (v_render_rgb) v_render_rgb[idx] = temp_v_render_rgb;
|
||||
if (v_ref_rgb) v_ref_rgb[idx] = temp_v_ref_rgb;
|
||||
if (v_render_depth) v_render_depth[idx] = temp_v_render_depth;
|
||||
if (v_ref_depth) v_ref_depth[idx] = temp_v_ref_depth;
|
||||
if (v_render_normal) v_render_normal[idx] = temp_v_render_normal;
|
||||
if (v_depth_normal) v_depth_normal[idx] = temp_v_depth_normal;
|
||||
if (v_ref_normal) v_ref_normal[idx] = temp_v_ref_normal;
|
||||
if (v_render_alpha) v_render_alpha[idx] = temp_v_render_alpha;
|
||||
if (v_rgb_dist) v_rgb_dist[idx] = temp_v_rgb_dist;
|
||||
if (v_depth_dist) v_depth_dist[idx] = temp_v_depth_dist;
|
||||
if (v_normal_dist) v_normal_dist[idx] = temp_v_normal_dist;
|
||||
}
|
||||
|
||||
__global__ void per_pixel_losses_reduce_forward_kernel(
|
||||
FixedArray<float, (uint)RawLossIndex::length>* __restrict__ raw_losses,
|
||||
size_t num_pixels,
|
||||
FixedArray<float, (uint)LossWeightIndex::length> loss_weights,
|
||||
FixedArray<float, (uint)LossIndex::length>* __restrict__ losses
|
||||
) {
|
||||
per_pixel_losses_reduce(
|
||||
raw_losses, num_pixels, &loss_weights,
|
||||
losses
|
||||
);
|
||||
}
|
||||
|
||||
__global__ void per_pixel_losses_reduce_backward_kernel(
|
||||
FixedArray<float, (uint)RawLossIndex::length>* __restrict__ raw_losses,
|
||||
size_t num_pixels,
|
||||
FixedArray<float, (uint)LossWeightIndex::length> loss_weights,
|
||||
FixedArray<float, (uint)LossIndex::length>* __restrict__ v_losses,
|
||||
FixedArray<float, (uint)RawLossIndex::length>* __restrict__ v_raw_losses
|
||||
) {
|
||||
per_pixel_losses_reduce_bwd(
|
||||
raw_losses, num_pixels, &loss_weights,
|
||||
v_losses, v_raw_losses
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
std::tuple<torch::Tensor, torch::Tensor>
|
||||
compute_per_pixel_losses_forward_tensor(
|
||||
std::optional<at::Tensor> render_rgb,
|
||||
std::optional<at::Tensor> ref_rgb,
|
||||
std::optional<at::Tensor> render_depth,
|
||||
std::optional<at::Tensor> ref_depth,
|
||||
std::optional<at::Tensor> render_normal,
|
||||
std::optional<at::Tensor> depth_normal,
|
||||
std::optional<at::Tensor> ref_normal,
|
||||
std::optional<at::Tensor> render_alpha,
|
||||
std::optional<at::Tensor> rgb_dist,
|
||||
std::optional<at::Tensor> depth_dist,
|
||||
std::optional<at::Tensor> normal_dist,
|
||||
std::optional<at::Tensor> ref_alpha,
|
||||
std::optional<at::Tensor> mask,
|
||||
std::optional<at::Tensor> depth_mask,
|
||||
std::optional<at::Tensor> normal_mask,
|
||||
const std::array<float, (uint)LossIndex::length> loss_weights_0
|
||||
) {
|
||||
long B = -1, H = -1, W = -1;
|
||||
auto check_generic = [&](std::string name, const at::Tensor& tensor) {
|
||||
CHECK_CUDA(tensor);
|
||||
if (tensor.ndimension() != 4)
|
||||
AT_ERROR(name + " must be (B, H, W, C)");
|
||||
if (B == -1)
|
||||
B = tensor.size(0), H = tensor.size(1), W = tensor.size(2);
|
||||
else if (B != tensor.size(0) || H != tensor.size(1) || W != tensor.size(2))
|
||||
AT_ERROR("Tensor shape mismatch with render_rgb (" + name + ")");
|
||||
};
|
||||
auto check_float = [&](std::string name, const std::optional<at::Tensor>& tensor, int ncomp) {
|
||||
if (tensor.has_value()) {
|
||||
check_generic(name, tensor.value());
|
||||
if (tensor.value().size(-1) != ncomp)
|
||||
AT_ERROR("Last dimension of " + name + " must be " + std::to_string(ncomp));
|
||||
}
|
||||
};
|
||||
|
||||
if (!render_rgb.has_value())
|
||||
AT_ERROR("render_rgb must be provided");
|
||||
check_float("render_rgb", render_rgb, 3);
|
||||
check_float("ref_rgb", ref_rgb, 3);
|
||||
check_float("render_depth", render_depth, 1);
|
||||
check_float("ref_depth", ref_depth, 1);
|
||||
check_float("render_normal", render_normal, 3);
|
||||
check_float("depth_normal", depth_normal, 3);
|
||||
check_float("ref_normal", ref_normal, 3);
|
||||
check_float("render_alpha", render_alpha, 1);
|
||||
check_float("rgb_dist", rgb_dist, 3);
|
||||
check_float("depth_dist", depth_dist, 1);
|
||||
check_float("normal_dist", normal_dist, 3);
|
||||
check_float("ref_alpha", ref_alpha, 1);
|
||||
check_float("mask", mask, 1);
|
||||
check_float("depth_mask", depth_mask, 1);
|
||||
check_float("normal_mask", normal_mask, 1);
|
||||
|
||||
size_t num_pixels = render_rgb.value().numel() / 3;
|
||||
|
||||
FixedArray<float, (uint)LossWeightIndex::length> loss_weights =
|
||||
*reinterpret_cast<const FixedArray<float, (uint)LossWeightIndex::length>*>(loss_weights_0.data());
|
||||
|
||||
torch::Tensor raw_losses = torch::zeros({(uint)RawLossIndex::length}, render_rgb.value().options());
|
||||
torch::Tensor losses = torch::zeros({(uint)LossIndex::length}, render_rgb.value().options());
|
||||
|
||||
per_pixel_losses_forward_kernel<<<_LAUNCH_ARGS_1D(num_pixels, WARP_SIZE*WARP_SIZE)>>>(
|
||||
num_pixels,
|
||||
render_rgb.has_value() ? (float3*)render_rgb.value().contiguous().data_ptr<float>() : nullptr,
|
||||
ref_rgb.has_value() ? (float3*)ref_rgb.value().contiguous().data_ptr<float>() : nullptr,
|
||||
render_depth.has_value() ? (float*)render_depth.value().contiguous().data_ptr<float>() : nullptr,
|
||||
ref_depth.has_value() ? (float*)ref_depth.value().contiguous().data_ptr<float>() : nullptr,
|
||||
render_normal.has_value() ? (float3*)render_normal.value().contiguous().data_ptr<float>() : nullptr,
|
||||
depth_normal.has_value() ? (float3*)depth_normal.value().contiguous().data_ptr<float>() : nullptr,
|
||||
ref_normal.has_value() ? (float3*)ref_normal.value().contiguous().data_ptr<float>() : nullptr,
|
||||
render_alpha.has_value() ? (float*)render_alpha.value().contiguous().data_ptr<float>() : nullptr,
|
||||
rgb_dist.has_value() ? (float3*)rgb_dist.value().contiguous().data_ptr<float>() : nullptr,
|
||||
depth_dist.has_value() ? (float*)depth_dist.value().contiguous().data_ptr<float>() : nullptr,
|
||||
normal_dist.has_value() ? (float3*)normal_dist.value().contiguous().data_ptr<float>() : nullptr,
|
||||
ref_alpha.has_value() ? (bool*)ref_alpha.value().contiguous().data_ptr<bool>() : nullptr,
|
||||
mask.has_value() ? (bool*)mask.value().contiguous().data_ptr<bool>() : nullptr,
|
||||
depth_mask.has_value() ? (bool*)depth_mask.value().contiguous().data_ptr<bool>() : nullptr,
|
||||
normal_mask.has_value() ? (bool*)normal_mask.value().contiguous().data_ptr<bool>() : nullptr,
|
||||
raw_losses.data_ptr<float>()
|
||||
);
|
||||
|
||||
per_pixel_losses_reduce_forward_kernel<<<1, 1>>>(
|
||||
(FixedArray<float, (uint)RawLossIndex::length>*)raw_losses.data_ptr<float>(),
|
||||
num_pixels,
|
||||
loss_weights,
|
||||
(FixedArray<float, (uint)LossIndex::length>*)losses.data_ptr<float>()
|
||||
);
|
||||
|
||||
return std::make_tuple(losses, raw_losses);
|
||||
}
|
||||
|
||||
|
||||
std::tuple<
|
||||
std::optional<at::Tensor>, // render_rgb
|
||||
std::optional<at::Tensor>, // ref_rgb
|
||||
std::optional<at::Tensor>, // render_depth
|
||||
std::optional<at::Tensor>, // ref_depth
|
||||
std::optional<at::Tensor>, // render_normal
|
||||
std::optional<at::Tensor>, // depth_normal
|
||||
std::optional<at::Tensor>, // ref_normal
|
||||
std::optional<at::Tensor>, // render_alpha
|
||||
std::optional<at::Tensor>, // rgb_dist
|
||||
std::optional<at::Tensor>, // depth_dist
|
||||
std::optional<at::Tensor> // normal_dist
|
||||
> compute_per_pixel_losses_backward_tensor(
|
||||
std::optional<at::Tensor> render_rgb,
|
||||
std::optional<at::Tensor> ref_rgb,
|
||||
std::optional<at::Tensor> render_depth,
|
||||
std::optional<at::Tensor> ref_depth,
|
||||
std::optional<at::Tensor> render_normal,
|
||||
std::optional<at::Tensor> depth_normal,
|
||||
std::optional<at::Tensor> ref_normal,
|
||||
std::optional<at::Tensor> render_alpha,
|
||||
std::optional<at::Tensor> rgb_dist,
|
||||
std::optional<at::Tensor> depth_dist,
|
||||
std::optional<at::Tensor> normal_dist,
|
||||
std::optional<at::Tensor> ref_alpha,
|
||||
std::optional<at::Tensor> mask,
|
||||
std::optional<at::Tensor> depth_mask,
|
||||
std::optional<at::Tensor> normal_mask,
|
||||
at::Tensor raw_losses,
|
||||
const std::array<float, (uint)LossIndex::length> loss_weights_0,
|
||||
at::Tensor v_losses
|
||||
) {
|
||||
|
||||
size_t num_pixels = render_rgb.value().numel() / 3;
|
||||
|
||||
std::optional<at::Tensor> v_render_rgb = render_rgb.has_value() ? (std::optional<at::Tensor>)at::empty_like(render_rgb.value()) : std::nullopt;
|
||||
std::optional<at::Tensor> v_ref_rgb = ref_rgb.has_value() ? (std::optional<at::Tensor>)at::empty_like(ref_rgb.value()) : std::nullopt;
|
||||
std::optional<at::Tensor> v_render_depth = render_depth.has_value() ? (std::optional<at::Tensor>)at::empty_like(render_depth.value()) : std::nullopt;
|
||||
std::optional<at::Tensor> v_ref_depth = ref_depth.has_value() ? (std::optional<at::Tensor>)at::empty_like(ref_depth.value()) : std::nullopt;
|
||||
std::optional<at::Tensor> v_render_normal = render_normal.has_value() ? (std::optional<at::Tensor>)at::empty_like(render_normal.value()) : std::nullopt;
|
||||
std::optional<at::Tensor> v_depth_normal = depth_normal.has_value() ? (std::optional<at::Tensor>)at::empty_like(depth_normal.value()) : std::nullopt;
|
||||
std::optional<at::Tensor> v_ref_normal = ref_normal.has_value() ? (std::optional<at::Tensor>)at::empty_like(ref_normal.value()) : std::nullopt;
|
||||
std::optional<at::Tensor> v_render_alpha = render_alpha.has_value() ? (std::optional<at::Tensor>)at::empty_like(render_alpha.value()) : std::nullopt;
|
||||
std::optional<at::Tensor> v_rgb_dist = rgb_dist.has_value() ? (std::optional<at::Tensor>)at::empty_like(rgb_dist.value()) : std::nullopt;
|
||||
std::optional<at::Tensor> v_depth_dist = depth_dist.has_value() ? (std::optional<at::Tensor>)at::empty_like(depth_dist.value()) : std::nullopt;
|
||||
std::optional<at::Tensor> v_normal_dist = normal_dist.has_value() ? (std::optional<at::Tensor>)at::empty_like(normal_dist.value()) : std::nullopt;
|
||||
|
||||
FixedArray<float, (uint)LossWeightIndex::length> loss_weights =
|
||||
*reinterpret_cast<const FixedArray<float, (uint)LossWeightIndex::length>*>(loss_weights_0.data());
|
||||
|
||||
torch::Tensor v_raw_losses = torch::empty({(uint)RawLossIndex::length}, render_rgb.value().options());
|
||||
|
||||
per_pixel_losses_reduce_backward_kernel<<<1, 1>>>(
|
||||
(FixedArray<float, (uint)RawLossIndex::length>*)raw_losses.data_ptr<float>(),
|
||||
num_pixels,
|
||||
loss_weights,
|
||||
(FixedArray<float, (uint)LossIndex::length>*)v_losses.data_ptr<float>(),
|
||||
(FixedArray<float, (uint)RawLossIndex::length>*)v_raw_losses.data_ptr<float>()
|
||||
);
|
||||
|
||||
per_pixel_losses_backward_kernel<<<_LAUNCH_ARGS_1D(num_pixels, WARP_SIZE*WARP_SIZE)>>>(
|
||||
num_pixels,
|
||||
render_rgb.has_value() ? (float3*)render_rgb.value().contiguous().data_ptr<float>() : nullptr,
|
||||
ref_rgb.has_value() ? (float3*)ref_rgb.value().contiguous().data_ptr<float>() : nullptr,
|
||||
render_depth.has_value() ? (float*)render_depth.value().contiguous().data_ptr<float>() : nullptr,
|
||||
ref_depth.has_value() ? (float*)ref_depth.value().contiguous().data_ptr<float>() : nullptr,
|
||||
render_normal.has_value() ? (float3*)render_normal.value().contiguous().data_ptr<float>() : nullptr,
|
||||
depth_normal.has_value() ? (float3*)depth_normal.value().contiguous().data_ptr<float>() : nullptr,
|
||||
ref_normal.has_value() ? (float3*)ref_normal.value().contiguous().data_ptr<float>() : nullptr,
|
||||
render_alpha.has_value() ? (float*)render_alpha.value().contiguous().data_ptr<float>() : nullptr,
|
||||
rgb_dist.has_value() ? (float3*)rgb_dist.value().contiguous().data_ptr<float>() : nullptr,
|
||||
depth_dist.has_value() ? (float*)depth_dist.value().contiguous().data_ptr<float>() : nullptr,
|
||||
normal_dist.has_value() ? (float3*)normal_dist.value().contiguous().data_ptr<float>() : nullptr,
|
||||
ref_alpha.has_value() ? (bool*)ref_alpha.value().contiguous().data_ptr<bool>() : nullptr,
|
||||
mask.has_value() ? (bool*)mask.value().contiguous().data_ptr<bool>() : nullptr,
|
||||
depth_mask.has_value() ? (bool*)depth_mask.value().contiguous().data_ptr<bool>() : nullptr,
|
||||
normal_mask.has_value() ? (bool*)normal_mask.value().contiguous().data_ptr<bool>() : nullptr,
|
||||
v_raw_losses.data_ptr<float>(),
|
||||
v_render_rgb.has_value() ? (float3*)v_render_rgb.value().contiguous().data_ptr<float>() : nullptr,
|
||||
v_ref_rgb.has_value() ? (float3*)v_ref_rgb.value().contiguous().data_ptr<float>() : nullptr,
|
||||
v_render_depth.has_value() ? (float*)v_render_depth.value().contiguous().data_ptr<float>() : nullptr,
|
||||
v_ref_depth.has_value() ? (float*)v_ref_depth.value().contiguous().data_ptr<float>() : nullptr,
|
||||
v_render_normal.has_value() ? (float3*)v_render_normal.value().contiguous().data_ptr<float>() : nullptr,
|
||||
v_depth_normal.has_value() ? (float3*)v_depth_normal.value().contiguous().data_ptr<float>() : nullptr,
|
||||
v_ref_normal.has_value() ? (float3*)v_ref_normal.value().contiguous().data_ptr<float>() : nullptr,
|
||||
v_render_alpha.has_value() ? (float*)v_render_alpha.value().contiguous().data_ptr<float>() : nullptr,
|
||||
v_rgb_dist.has_value() ? (float3*)v_rgb_dist.value().contiguous().data_ptr<float>() : nullptr,
|
||||
v_depth_dist.has_value() ? (float*)v_depth_dist.value().contiguous().data_ptr<float>() : nullptr,
|
||||
v_normal_dist.has_value() ? (float3*)v_normal_dist.value().contiguous().data_ptr<float>() : nullptr
|
||||
);
|
||||
|
||||
return std::make_tuple(
|
||||
v_render_rgb,
|
||||
v_ref_rgb,
|
||||
v_render_depth,
|
||||
v_ref_depth,
|
||||
v_render_normal,
|
||||
v_depth_normal,
|
||||
v_ref_normal,
|
||||
v_render_alpha,
|
||||
v_rgb_dist,
|
||||
v_depth_dist,
|
||||
v_normal_dist
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,116 @@
|
||||
#pragma once
|
||||
|
||||
#include <torch/types.h>
|
||||
|
||||
|
||||
enum class RawLossIndex {
|
||||
RgbL1,
|
||||
RgbL2,
|
||||
DepthSupX,
|
||||
DepthSupY,
|
||||
DepthSupXX,
|
||||
DepthSupYY,
|
||||
DepthSupXY,
|
||||
RenderNormalSup,
|
||||
DepthNormalSup,
|
||||
AlphaSup,
|
||||
AlphaSupUnder,
|
||||
NormalReg,
|
||||
AlphaReg,
|
||||
RgbDistReg,
|
||||
DepthDistReg,
|
||||
NormalDistReg,
|
||||
MaskTotal,
|
||||
DepthMaskTotal,
|
||||
RenderNormalMaskTotal,
|
||||
DepthNormalMaskTotal,
|
||||
NormalRegMaskTotal,
|
||||
length
|
||||
};
|
||||
|
||||
enum class LossWeightIndex {
|
||||
RgbSup,
|
||||
DepthSup,
|
||||
NormalSup,
|
||||
AlphaSup,
|
||||
AlphaSupUnder,
|
||||
NormalReg,
|
||||
AlphaReg,
|
||||
RgbDistReg,
|
||||
DepthDistReg,
|
||||
NormalDistReg,
|
||||
length
|
||||
};
|
||||
|
||||
enum class LossIndex {
|
||||
RgbL1,
|
||||
RgbPSNR,
|
||||
DepthSup,
|
||||
NormalSup,
|
||||
AlphaSup,
|
||||
NormalReg,
|
||||
AlphaReg,
|
||||
RgbDistReg,
|
||||
DepthDistReg,
|
||||
NormalDistReg,
|
||||
length
|
||||
};
|
||||
|
||||
|
||||
/* == AUTO HEADER GENERATOR - DO NOT EDIT THIS LINE OR ANYTHING BELOW THIS LINE == */
|
||||
|
||||
|
||||
|
||||
std::tuple<torch::Tensor, torch::Tensor>
|
||||
compute_per_pixel_losses_forward_tensor(
|
||||
std::optional<at::Tensor> render_rgb,
|
||||
std::optional<at::Tensor> ref_rgb,
|
||||
std::optional<at::Tensor> render_depth,
|
||||
std::optional<at::Tensor> ref_depth,
|
||||
std::optional<at::Tensor> render_normal,
|
||||
std::optional<at::Tensor> depth_normal,
|
||||
std::optional<at::Tensor> ref_normal,
|
||||
std::optional<at::Tensor> render_alpha,
|
||||
std::optional<at::Tensor> rgb_dist,
|
||||
std::optional<at::Tensor> depth_dist,
|
||||
std::optional<at::Tensor> normal_dist,
|
||||
std::optional<at::Tensor> ref_alpha,
|
||||
std::optional<at::Tensor> mask,
|
||||
std::optional<at::Tensor> depth_mask,
|
||||
std::optional<at::Tensor> normal_mask,
|
||||
const std::array<float, (uint)LossIndex::length> loss_weights_0
|
||||
);
|
||||
|
||||
|
||||
std::tuple<
|
||||
std::optional<at::Tensor>, // render_rgb
|
||||
std::optional<at::Tensor>, // ref_rgb
|
||||
std::optional<at::Tensor>, // render_depth
|
||||
std::optional<at::Tensor>, // ref_depth
|
||||
std::optional<at::Tensor>, // render_normal
|
||||
std::optional<at::Tensor>, // depth_normal
|
||||
std::optional<at::Tensor>, // ref_normal
|
||||
std::optional<at::Tensor>, // render_alpha
|
||||
std::optional<at::Tensor>, // rgb_dist
|
||||
std::optional<at::Tensor>, // depth_dist
|
||||
std::optional<at::Tensor> // normal_dist
|
||||
> compute_per_pixel_losses_backward_tensor(
|
||||
std::optional<at::Tensor> render_rgb,
|
||||
std::optional<at::Tensor> ref_rgb,
|
||||
std::optional<at::Tensor> render_depth,
|
||||
std::optional<at::Tensor> ref_depth,
|
||||
std::optional<at::Tensor> render_normal,
|
||||
std::optional<at::Tensor> depth_normal,
|
||||
std::optional<at::Tensor> ref_normal,
|
||||
std::optional<at::Tensor> render_alpha,
|
||||
std::optional<at::Tensor> rgb_dist,
|
||||
std::optional<at::Tensor> depth_dist,
|
||||
std::optional<at::Tensor> normal_dist,
|
||||
std::optional<at::Tensor> ref_alpha,
|
||||
std::optional<at::Tensor> mask,
|
||||
std::optional<at::Tensor> depth_mask,
|
||||
std::optional<at::Tensor> normal_mask,
|
||||
at::Tensor raw_losses,
|
||||
const std::array<float, (uint)LossIndex::length> loss_weights_0,
|
||||
at::Tensor v_losses
|
||||
);
|
||||
@@ -5,6 +5,7 @@
|
||||
#include "SphericalHarmonics.cuh"
|
||||
#include "BackgroundSphericalHarmonics.cuh"
|
||||
#include "PerSplatLoss.cuh"
|
||||
#include "PerPixelLoss.cuh"
|
||||
#include "PixelWise.cuh"
|
||||
#include "SplatTileIntersector.cuh"
|
||||
#include "Projection.cuh"
|
||||
@@ -41,6 +42,10 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.def("compute_per_splat_losses_forward", &compute_per_splat_losses_forward_tensor);
|
||||
m.def("compute_per_splat_losses_backward", &compute_per_splat_losses_backward_tensor);
|
||||
|
||||
// PerPixelLoss.cuh
|
||||
m.def("compute_per_pixel_losses_forward", &compute_per_pixel_losses_forward_tensor);
|
||||
m.def("compute_per_pixel_losses_backward", &compute_per_pixel_losses_backward_tensor);
|
||||
|
||||
// PixelWise.cuh
|
||||
m.def("blend_background_forward", &blend_background_forward_tensor);
|
||||
m.def("blend_background_backward", &blend_background_backward_tensor);
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,3 +1,4 @@
|
||||
import per_splat_losses;
|
||||
import per_pixel_losses;
|
||||
import pixel_wise;
|
||||
import depth_to_normal;
|
||||
|
||||
@@ -0,0 +1,327 @@
|
||||
// All per-Pixel losses in one kernel
|
||||
|
||||
[ForceInline]
|
||||
[Differentiable]
|
||||
float mean3(float3 a) {
|
||||
return (a.x + a.y + a.z) * (1.0 / 3.0);
|
||||
}
|
||||
|
||||
[ForceInline]
|
||||
[Differentiable]
|
||||
float l1_loss(float3 a, float3 b) {
|
||||
return mean3(abs(b - a));
|
||||
}
|
||||
|
||||
[ForceInline]
|
||||
[Differentiable]
|
||||
float l2_loss(float3 a, float3 b) {
|
||||
return dot(b - a, b - a) * (1.0 / 3.0);
|
||||
}
|
||||
|
||||
[ForceInline]
|
||||
[Differentiable]
|
||||
float3 normalize_normal(float3 normal, inout bool mask) {
|
||||
float norm2 = dot(normal, normal);
|
||||
if (norm2 == 0.0) {
|
||||
mask = false;
|
||||
return float3(0.0);
|
||||
}
|
||||
return normal * rsqrt(norm2);
|
||||
}
|
||||
|
||||
[ForceInline]
|
||||
[Differentiable]
|
||||
float normal_loss(float3 normal, float3 normal_ref) {
|
||||
// return 1.0 - dot(normal, normal_ref);
|
||||
return l1_loss(normal, normal_ref);
|
||||
}
|
||||
|
||||
[ForceInline]
|
||||
[Differentiable]
|
||||
float bce_loss(float val, float ref) {
|
||||
const float eps = 1e-6;
|
||||
return -lerp(log(max(1.0 - val, eps)), log(max(val, eps)), ref);
|
||||
}
|
||||
|
||||
[ForceInline]
|
||||
[Differentiable]
|
||||
float alpha_loss(float x, float y) {
|
||||
return bce_loss(max(x, y), y);
|
||||
}
|
||||
|
||||
enum RawLossIndex {
|
||||
RgbL1,
|
||||
RgbL2,
|
||||
DepthSupX,
|
||||
DepthSupY,
|
||||
DepthSupXX,
|
||||
DepthSupYY,
|
||||
DepthSupXY,
|
||||
RenderNormalSup,
|
||||
DepthNormalSup,
|
||||
AlphaSup,
|
||||
AlphaSupUnder,
|
||||
NormalReg,
|
||||
AlphaReg,
|
||||
RgbDistReg,
|
||||
DepthDistReg,
|
||||
NormalDistReg,
|
||||
MaskTotal,
|
||||
DepthMaskTotal,
|
||||
RenderNormalMaskTotal,
|
||||
DepthNormalMaskTotal,
|
||||
NormalRegMaskTotal,
|
||||
length
|
||||
};
|
||||
|
||||
|
||||
[Differentiable]
|
||||
[CudaDeviceExport]
|
||||
float[RawLossIndex::length] per_pixel_losses(
|
||||
float3 render_rgb,
|
||||
float3 ref_rgb,
|
||||
float render_depth,
|
||||
float ref_depth,
|
||||
float3 render_normal,
|
||||
float3 depth_normal,
|
||||
float3 ref_normal,
|
||||
float render_alpha,
|
||||
float3 rgb_dist,
|
||||
float depth_dist,
|
||||
float3 normal_dist,
|
||||
bool ref_alpha,
|
||||
bool mask,
|
||||
bool depth_mask,
|
||||
bool normal_mask
|
||||
) {
|
||||
float[RawLossIndex::length] losses;
|
||||
|
||||
// RGB - L1 and L2
|
||||
losses[RawLossIndex::RgbL1] = float(mask) * l1_loss(render_rgb, ref_rgb);
|
||||
losses[RawLossIndex::RgbL2] = float(mask) * l2_loss(render_rgb, ref_rgb);
|
||||
|
||||
// Depth - Pearson correlation in log space
|
||||
depth_mask &= mask;
|
||||
render_depth = float(depth_mask) * log(max(render_depth, 1e-4));
|
||||
ref_depth = float(depth_mask) * log(max(ref_depth, 1e-4));
|
||||
losses[RawLossIndex::DepthSupX] = render_depth;
|
||||
losses[RawLossIndex::DepthSupY] = ref_depth;
|
||||
losses[RawLossIndex::DepthSupXX] = render_depth * render_depth;
|
||||
losses[RawLossIndex::DepthSupYY] = ref_depth * ref_depth;
|
||||
losses[RawLossIndex::DepthSupXY] = render_depth * ref_depth;
|
||||
|
||||
// Normal - Supervision and regularization - Pairwise loss between different normals
|
||||
normal_mask &= mask;
|
||||
bool render_normal_mask = true,
|
||||
depth_normal_mask = true;
|
||||
render_normal = normalize_normal(render_normal, render_normal_mask);
|
||||
depth_normal = normalize_normal(depth_normal, depth_normal_mask);
|
||||
ref_normal = normalize_normal(ref_normal, normal_mask);
|
||||
losses[RawLossIndex::RenderNormalSup] =
|
||||
float(render_normal_mask & normal_mask) * normal_loss(render_normal, ref_normal);
|
||||
losses[RawLossIndex::DepthNormalSup] =
|
||||
float(depth_normal_mask & normal_mask) * normal_loss(depth_normal, ref_normal);
|
||||
losses[RawLossIndex::NormalReg] =
|
||||
float(render_normal_mask & depth_normal_mask) * normal_loss(render_normal, depth_normal);
|
||||
|
||||
// Alpha - loss with mask + regularization to go to 0 or 1
|
||||
render_alpha = clamp(render_alpha, 0.0f, 1.0f);
|
||||
losses[RawLossIndex::AlphaSup] = alpha_loss(render_alpha, float(ref_alpha));
|
||||
losses[RawLossIndex::AlphaSupUnder] = alpha_loss(1.0 - render_alpha, 1.0 - float(ref_alpha));
|
||||
losses[RawLossIndex::AlphaReg] = 4.0f * render_alpha * (1.0f - render_alpha); // TODO: support push-to-1 mode
|
||||
|
||||
// Distortion regularization
|
||||
losses[RawLossIndex::RgbDistReg] = mean3(rgb_dist) / max(render_alpha, 1e-12);
|
||||
losses[RawLossIndex::DepthDistReg] = depth_dist / max(render_alpha, 1e-12);
|
||||
losses[RawLossIndex::NormalDistReg] = mean3(normal_dist) / max(render_alpha, 1e-12);
|
||||
|
||||
// Mask total
|
||||
losses[RawLossIndex::MaskTotal] = float(mask);
|
||||
losses[RawLossIndex::DepthMaskTotal] = float(depth_mask);
|
||||
losses[RawLossIndex::RenderNormalMaskTotal] = float(render_normal_mask & normal_mask);
|
||||
losses[RawLossIndex::DepthNormalMaskTotal] = float(depth_normal_mask & normal_mask);
|
||||
losses[RawLossIndex::NormalRegMaskTotal] = float(render_normal_mask & depth_normal_mask);
|
||||
|
||||
return losses;
|
||||
}
|
||||
|
||||
[CudaDeviceExport]
|
||||
void per_pixel_losses_bwd(
|
||||
float3 render_rgb,
|
||||
float3 ref_rgb,
|
||||
float render_depth,
|
||||
float ref_depth,
|
||||
float3 render_normal,
|
||||
float3 depth_normal,
|
||||
float3 ref_normal,
|
||||
float render_alpha,
|
||||
float3 rgb_dist,
|
||||
float depth_dist,
|
||||
float3 normal_dist,
|
||||
bool mask,
|
||||
bool depth_mask,
|
||||
bool normal_mask,
|
||||
bool alpha_mask,
|
||||
float[RawLossIndex::length] v_losses,
|
||||
out float3 v_render_rgb,
|
||||
out float3 v_ref_rgb,
|
||||
out float v_render_depth,
|
||||
out float v_ref_depth,
|
||||
out float3 v_render_normal,
|
||||
out float3 v_depth_normal,
|
||||
out float3 v_ref_normal,
|
||||
out float v_render_alpha,
|
||||
out float3 v_rgb_dist,
|
||||
out float v_depth_dist,
|
||||
out float3 v_normal_dist
|
||||
) {
|
||||
DifferentialPair<float3> dp_render_rgb = diffPair(render_rgb);
|
||||
DifferentialPair<float3> dp_ref_rgb = diffPair(ref_rgb);
|
||||
DifferentialPair<float> dp_render_depth = diffPair(render_depth);
|
||||
DifferentialPair<float> dp_ref_depth = diffPair(ref_depth);
|
||||
DifferentialPair<float3> dp_render_normal = diffPair(render_normal);
|
||||
DifferentialPair<float3> dp_depth_normal = diffPair(depth_normal);
|
||||
DifferentialPair<float3> dp_ref_normal = diffPair(ref_normal);
|
||||
DifferentialPair<float> dp_render_alpha = diffPair(render_alpha);
|
||||
DifferentialPair<float3> dp_rgb_dist = diffPair(rgb_dist);
|
||||
DifferentialPair<float> dp_depth_dist = diffPair(depth_dist);
|
||||
DifferentialPair<float3> dp_normal_dist = diffPair(normal_dist);
|
||||
bwd_diff(per_pixel_losses)(
|
||||
dp_render_rgb,
|
||||
dp_ref_rgb,
|
||||
dp_render_depth,
|
||||
dp_ref_depth,
|
||||
dp_render_normal,
|
||||
dp_depth_normal,
|
||||
dp_ref_normal,
|
||||
dp_render_alpha,
|
||||
dp_rgb_dist,
|
||||
dp_depth_dist,
|
||||
dp_normal_dist,
|
||||
mask,
|
||||
depth_mask,
|
||||
normal_mask,
|
||||
alpha_mask,
|
||||
v_losses
|
||||
);
|
||||
v_render_rgb = dp_render_rgb.d;
|
||||
v_ref_rgb = dp_ref_rgb.d;
|
||||
v_render_depth = dp_render_depth.d;
|
||||
v_ref_depth = dp_ref_depth.d;
|
||||
v_render_normal = dp_render_normal.d;
|
||||
v_depth_normal = dp_depth_normal.d;
|
||||
v_ref_normal = dp_ref_normal.d;
|
||||
v_render_alpha = dp_render_alpha.d;
|
||||
v_rgb_dist = dp_rgb_dist.d;
|
||||
v_depth_dist = dp_depth_dist.d;
|
||||
v_normal_dist = dp_normal_dist.d;
|
||||
}
|
||||
|
||||
enum LossWeightIndex {
|
||||
RgbSup,
|
||||
DepthSup,
|
||||
NormalSup,
|
||||
AlphaSup,
|
||||
AlphaSupUnder,
|
||||
NormalReg,
|
||||
AlphaReg,
|
||||
RgbDistReg,
|
||||
DepthDistReg,
|
||||
NormalDistReg,
|
||||
length
|
||||
};
|
||||
|
||||
enum LossIndex {
|
||||
RgbL1,
|
||||
RgbPSNR,
|
||||
DepthSup,
|
||||
NormalSup,
|
||||
AlphaSup,
|
||||
NormalReg,
|
||||
AlphaReg,
|
||||
RgbDistReg,
|
||||
DepthDistReg,
|
||||
NormalDistReg,
|
||||
length
|
||||
};
|
||||
|
||||
[Differentiable]
|
||||
[CudaDeviceExport]
|
||||
float[LossIndex::length] per_pixel_losses_reduce(
|
||||
float[RawLossIndex::length] raw_losses,
|
||||
no_diff float num_pixels,
|
||||
no_diff float[LossWeightIndex::length] weights
|
||||
) {
|
||||
float[LossIndex::length] losses;
|
||||
|
||||
// image loss
|
||||
losses[LossIndex::RgbL1] = weights[LossWeightIndex::RgbSup] *
|
||||
raw_losses[RawLossIndex::RgbL1] / max(raw_losses[RawLossIndex::MaskTotal], 1);
|
||||
losses[LossIndex::RgbPSNR] = -10.0 * log10(
|
||||
raw_losses[RawLossIndex::RgbL2] / max(raw_losses[RawLossIndex::MaskTotal], 1));
|
||||
|
||||
// depth supervision loss - Pearson correlation
|
||||
losses[LossIndex::DepthSup] = raw_losses[RawLossIndex::DepthMaskTotal] > 0.0 ?
|
||||
weights[LossWeightIndex::DepthSup] * (1.0 -
|
||||
(
|
||||
raw_losses[RawLossIndex::DepthSupXY] -
|
||||
raw_losses[RawLossIndex::DepthSupX] * raw_losses[RawLossIndex::DepthSupY]
|
||||
/ raw_losses[RawLossIndex::DepthMaskTotal]
|
||||
) / sqrt(max(1e-12,
|
||||
(raw_losses[RawLossIndex::DepthSupXX] -
|
||||
raw_losses[RawLossIndex::DepthSupX] * raw_losses[RawLossIndex::DepthSupX]
|
||||
/ raw_losses[RawLossIndex::DepthMaskTotal]
|
||||
) * (raw_losses[RawLossIndex::DepthSupYY] -
|
||||
raw_losses[RawLossIndex::DepthSupY] * raw_losses[RawLossIndex::DepthSupY]
|
||||
/ raw_losses[RawLossIndex::DepthMaskTotal]
|
||||
)
|
||||
+ 1.0)) // numerical issue
|
||||
) : 0.0;
|
||||
|
||||
// normal supervision loss
|
||||
losses[LossIndex::NormalSup] = weights[LossWeightIndex::NormalSup] * (
|
||||
raw_losses[RawLossIndex::RenderNormalSup] / max(raw_losses[RawLossIndex::RenderNormalMaskTotal], 1) +
|
||||
raw_losses[RawLossIndex::DepthNormalSup] / max(raw_losses[RawLossIndex::DepthNormalMaskTotal], 1)
|
||||
) / max(int(raw_losses[RawLossIndex::RenderNormalMaskTotal] > 0.5) +
|
||||
int(raw_losses[RawLossIndex::DepthNormalMaskTotal] > 0.5), 1);
|
||||
|
||||
// alpha supervision loss
|
||||
losses[LossIndex::AlphaSup] = (
|
||||
weights[LossWeightIndex::AlphaSup] * raw_losses[RawLossIndex::AlphaSup] +
|
||||
weights[LossWeightIndex::AlphaSupUnder] * raw_losses[RawLossIndex::AlphaSupUnder]
|
||||
) / max(num_pixels, 1);
|
||||
|
||||
// normal regularization
|
||||
losses[LossIndex::NormalReg] = weights[LossWeightIndex::NormalReg] *
|
||||
raw_losses[RawLossIndex::NormalReg] / max(raw_losses[RawLossIndex::NormalRegMaskTotal], 1);
|
||||
|
||||
// alpha regularization
|
||||
// TODO: push-to-1 mode for random background
|
||||
losses[LossIndex::AlphaReg] = weights[LossWeightIndex::AlphaReg] *
|
||||
raw_losses[RawLossIndex::AlphaReg] / max(num_pixels, 1);
|
||||
|
||||
// distortion regularizations
|
||||
losses[LossIndex::RgbDistReg] = weights[LossWeightIndex::RgbDistReg] *
|
||||
raw_losses[RawLossIndex::RgbDistReg] / max(num_pixels, 1);
|
||||
losses[LossIndex::DepthDistReg] = weights[LossWeightIndex::DepthDistReg] *
|
||||
raw_losses[RawLossIndex::DepthDistReg] / max(num_pixels, 1);
|
||||
losses[LossIndex::NormalDistReg] = weights[LossWeightIndex::NormalDistReg] *
|
||||
raw_losses[RawLossIndex::NormalDistReg] / max(num_pixels, 1);
|
||||
|
||||
return losses;
|
||||
}
|
||||
|
||||
[CudaDeviceExport]
|
||||
float[RawLossIndex::length] per_pixel_losses_reduce_bwd(
|
||||
const float[RawLossIndex::length] raw_losses,
|
||||
no_diff float num_pixels,
|
||||
const no_diff float[LossWeightIndex::length] weights,
|
||||
float[LossIndex::length] v_losses
|
||||
) {
|
||||
DifferentialPair<float[RawLossIndex::length]> dp_raw_losses = diffPair(raw_losses);
|
||||
bwd_diff(per_pixel_losses_reduce)(
|
||||
dp_raw_losses, num_pixels, weights,
|
||||
v_losses
|
||||
);
|
||||
return dp_raw_losses.d;
|
||||
}
|
||||
@@ -32,8 +32,8 @@ class Config:
|
||||
depth_distortion_uv_degree = -1
|
||||
depth_supervision_weight = 0.0
|
||||
normal_supervision_weight = 0.0
|
||||
alpha_supervision_weight = 0.0
|
||||
alpha_supervision_weight_under = 0.0
|
||||
alpha_loss_weight = 0.0
|
||||
alpha_loss_weight_under = 0.0
|
||||
adaptive_exposure_mode = ""
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user