fused per pixel loss function

This commit is contained in:
Harry Chen
2025-11-01 20:02:45 -04:00
parent 0f4e05414d
commit 328de10470
12 changed files with 2778 additions and 885 deletions
+2 -1
View File
@@ -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")
+2 -2
View File
@@ -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):
+200 -240
View File
@@ -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
+4 -1
View File
@@ -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
View File
@@ -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
View File
@@ -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;
}
+2 -2
View File
@@ -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 = ""