mirror of
https://github.com/harry7557558/spirula-studio.git
synced 2026-10-02 02:44:54 +08:00
regularization (forward)
This commit is contained in:
@@ -1,3 +1,5 @@
|
||||
.vscode/
|
||||
|
||||
outputs/
|
||||
|
||||
# Byte-compiled / optimized / DLL files
|
||||
|
||||
+45
-27
@@ -108,7 +108,7 @@ class SpirulaeModelConfig(ModelConfig):
|
||||
"""period of steps where gaussians are culled and densified"""
|
||||
resolution_schedule: int = 3000
|
||||
"""training starts at 1/d resolution, every n steps this is doubled"""
|
||||
background_color: Literal["random", "black", "white"] = "random"
|
||||
background_color: Literal["random", "black", "white"] = "white"
|
||||
"""Whether to randomize the background color."""
|
||||
num_downscales: int = 2
|
||||
"""at the beginning, resolution is 1/2^d, where d is this number"""
|
||||
@@ -394,6 +394,8 @@ class SpirulaeModel(Model):
|
||||
|
||||
def after_train(self, step: int):
|
||||
assert step == self.step
|
||||
if self.max_2Dsize is None:
|
||||
self.max_2Dsize = torch.zeros_like(self.radii, dtype=torch.float32)
|
||||
# to save some training time, we no longer need to update those stats post refinement
|
||||
if self.step >= self.config.stop_split_at:
|
||||
return
|
||||
@@ -414,8 +416,6 @@ class SpirulaeModel(Model):
|
||||
self.xys_grad_norm[visible_mask] = grads[visible_mask] + self.xys_grad_norm[visible_mask]
|
||||
|
||||
# update the max screen size, as a ratio of number of pixels
|
||||
if self.max_2Dsize is None:
|
||||
self.max_2Dsize = torch.zeros_like(self.radii, dtype=torch.float32)
|
||||
# newradii = self.radii.detach()[self.depth_sort_i][visible_mask]
|
||||
newradii = self.radii.detach()[visible_mask]
|
||||
self.max_2Dsize[visible_mask] = torch.maximum(
|
||||
@@ -755,15 +755,8 @@ class SpirulaeModel(Model):
|
||||
|
||||
# print(self.config.sh_degree, rgbs.shape)
|
||||
|
||||
def kernel(r):
|
||||
f1 = 1.0-1.5*r*r*(1.0-0.5*r)
|
||||
f2 = 0.25*(2.0-r)**3
|
||||
f = f1 + (f2-f1) * (0.5+0.5*torch.sign(r-1.0))
|
||||
return torch.fmax(f, torch.zeros_like(f))
|
||||
kernel_radius = 2.0
|
||||
|
||||
BLOCK_WIDTH = 16 # this controls the tile size of rasterization, 16 is a good default
|
||||
self.xys, depths, self.radii, conics, comp, num_tiles_hit, cov3d = project_gaussians( # type: ignore
|
||||
self.xys, depths, depth_grads, self.radii, conics, comp, num_tiles_hit, cov3d = project_gaussians( # type: ignore
|
||||
means_crop,
|
||||
torch.exp(scales_crop),
|
||||
1,
|
||||
@@ -794,19 +787,12 @@ class SpirulaeModel(Model):
|
||||
else:
|
||||
raise ValueError("Unknown rasterize_mode: %s", self.config.rasterize_mode)
|
||||
|
||||
# sort_i = torch.argsort(depths)
|
||||
# self.xys, depths = self.xys[sort_i], depths[sort_i] # (2,), 1
|
||||
# self.radii, conics = self.radii[sort_i], conics[sort_i] # 1, (3,)
|
||||
# comp, num_tiles_hit = comp[sort_i], num_tiles_hit[sort_i] # 1, 1
|
||||
# rgbs, opacities = rgbs[sort_i], opacities[sort_i]
|
||||
# with torch.no_grad():
|
||||
# self.depth_sort_i = torch.zeros_like(sort_i)
|
||||
# self.depth_sort_i[sort_i.clone()] = torch.arange(0, len(depths), device=sort_i.device)
|
||||
# print((depths[1:] >= depths[:-1]).sum().item(), (depths[1:] > -1e4).sum().item())
|
||||
# depth_grads_normalized = depth_grads / torch.norm(depth_grads, dim=1, keepdim=True)
|
||||
|
||||
rgb, alpha = rasterize_gaussians( # type: ignore
|
||||
rgb, reg_depth, reg_normal, alpha = rasterize_gaussians( # type: ignore
|
||||
self.xys,
|
||||
depths,
|
||||
depth_grads,
|
||||
self.radii,
|
||||
conics,
|
||||
num_tiles_hit, # type: ignore
|
||||
@@ -815,30 +801,62 @@ class SpirulaeModel(Model):
|
||||
H,
|
||||
W,
|
||||
BLOCK_WIDTH,
|
||||
background=background,
|
||||
return_alpha=True,
|
||||
background=background
|
||||
) # type: ignore
|
||||
alpha = alpha[..., None]
|
||||
rgb = torch.clamp(rgb, max=1.0) # type: ignore
|
||||
depth_im = None
|
||||
depth_grad_im = None
|
||||
if self.config.output_depth_during_training or not self.training:
|
||||
depth_im = rasterize_gaussians( # type: ignore
|
||||
self.xys,
|
||||
depths,
|
||||
depth_grads,
|
||||
self.radii,
|
||||
conics,
|
||||
num_tiles_hit, # type: ignore
|
||||
depths[:, None].repeat(1, 3),
|
||||
torch.concatenate((depths[:, None], depth_grads), axis=1),
|
||||
opacities,
|
||||
H,
|
||||
W,
|
||||
BLOCK_WIDTH,
|
||||
background=torch.zeros(3, device=self.device),
|
||||
)[..., 0:1] # type: ignore
|
||||
)[0] # type: ignore
|
||||
depth_im, depth_grad_im = depth_im[..., 0:1], depth_im[..., 1:3]
|
||||
depth_im = torch.where(alpha > 0, depth_im / alpha, depth_im.detach().max())
|
||||
# depth_grad_im = torch.where(alpha > 0, depth_grad_im / alpha, depth_grad_im.detach().max())
|
||||
depth_grad_norm_im = torch.norm(depth_grad_im, dim=2, keepdim=True)
|
||||
depth_grad_hue_im = torch.atan2(depth_grad_im[...,1:2], depth_grad_im[...,0:1])
|
||||
depth_grad_im = torch.clip(
|
||||
torch.tanh(depth_grad_norm_im * 0.05*torch.numel(depth_im)**0.5) * \
|
||||
(torch.concat((
|
||||
torch.cos(depth_grad_hue_im),
|
||||
torch.cos(depth_grad_hue_im-2.0*np.pi/3.0),
|
||||
torch.cos(depth_grad_hue_im+2.0*np.pi/3.0)
|
||||
), axis=2)*0.5+0.5) * 2.0, 0.0, 1.0)
|
||||
|
||||
# print(rgb.shape, alpha.shape)
|
||||
return {"rgb": rgb, "depth": depth_im, "accumulation": alpha, "background": background} # type: ignore
|
||||
with torch.no_grad():
|
||||
# print(torch.amin(depth_grad_norm_im).item(), torch.mean(depth_grad_norm_im).item(), torch.amax(depth_grad_norm_im).item())
|
||||
pass
|
||||
|
||||
with torch.no_grad():
|
||||
# print(torch.isnan(reg_depth.view(-1)).float().mean().item(),
|
||||
# torch.isnan(reg_normal.view(-1)).float().mean().item(),
|
||||
# reg_depth_nonzero_mean = (reg_depth!=0.0).float().mean().item(),
|
||||
# reg_normal_nonzero_mean = (reg_normal!=0.0).float().mean().item())
|
||||
# print(rgb.shape, alpha.shape, reg_depth.shape, reg_normal.shape)
|
||||
# print(torch.amin(reg_depth).item(), torch.mean(reg_depth).item(), torch.amax(reg_depth).item())
|
||||
pass
|
||||
|
||||
return {
|
||||
"rgb": rgb,
|
||||
"depth": depth_im,
|
||||
"accumulation": alpha,
|
||||
"depth_grad": depth_grad_im,
|
||||
"reg_depth": reg_depth.unsqueeze(2),
|
||||
"reg_normal": reg_normal.unsqueeze(2),
|
||||
"background": background
|
||||
} # type: ignore
|
||||
|
||||
def get_gt_img(self, image: torch.Tensor):
|
||||
"""Compute groundtruth image with iteration dependent downscale factor for evaluation purpose
|
||||
|
||||
@@ -105,8 +105,10 @@ __global__ void rasterize_backward_kernel(
|
||||
conic_batch[tr] = conics[g_id];
|
||||
rgbs_batch[tr] = rgbs[g_id];
|
||||
}
|
||||
|
||||
// wait for other threads to collect the gaussians in batch
|
||||
block.sync();
|
||||
|
||||
// process gaussians in the current batch for this pixel
|
||||
// 0 index is the furthest back gaussian in the batch
|
||||
for (int t = max(0,batch_end - warp_bin_final); t < batch_size; ++t) {
|
||||
@@ -224,6 +226,7 @@ __global__ void project_gaussians_backward_kernel(
|
||||
const float* __restrict__ compensation,
|
||||
const float2* __restrict__ v_xy,
|
||||
const float* __restrict__ v_depth,
|
||||
const float2* __restrict__ v_depth_grad,
|
||||
const float3* __restrict__ v_conic,
|
||||
const float* __restrict__ v_compensation,
|
||||
float3* __restrict__ v_cov2d,
|
||||
@@ -385,7 +388,7 @@ __device__ void scale_rot_to_cov3d_vjp(
|
||||
);
|
||||
glm::mat3 R = quat_to_rotmat(quat);
|
||||
glm::mat3 S = scale_to_mat(
|
||||
make_float3(scale.x, scale.y, 0.0f), glob_scale);
|
||||
{ scale.x, scale.y, 0.0f }, glob_scale);
|
||||
glm::mat3 M = R * S;
|
||||
// https://math.stackexchange.com/a/3850121
|
||||
// for D = W * X, G = df/dD
|
||||
|
||||
@@ -20,6 +20,7 @@ __global__ void project_gaussians_backward_kernel(
|
||||
const float* __restrict__ compensation,
|
||||
const float2* __restrict__ v_xy,
|
||||
const float* __restrict__ v_depth,
|
||||
const float2* __restrict__ v_depth_grad,
|
||||
const float3* __restrict__ v_conic,
|
||||
const float* __restrict__ v_compensation,
|
||||
float3* __restrict__ v_cov2d,
|
||||
|
||||
@@ -158,6 +158,7 @@ std::tuple<
|
||||
torch::Tensor,
|
||||
torch::Tensor,
|
||||
torch::Tensor,
|
||||
torch::Tensor,
|
||||
torch::Tensor>
|
||||
project_gaussians_forward_tensor(
|
||||
const int num_points,
|
||||
@@ -194,6 +195,8 @@ project_gaussians_forward_tensor(
|
||||
torch::zeros({num_points, 2}, means3d.options().dtype(torch::kFloat32));
|
||||
torch::Tensor depths_d =
|
||||
torch::zeros({num_points}, means3d.options().dtype(torch::kFloat32));
|
||||
torch::Tensor depth_grads_d =
|
||||
torch::zeros({num_points, 2}, means3d.options().dtype(torch::kFloat32));
|
||||
torch::Tensor radii_d =
|
||||
torch::zeros({num_points}, means3d.options().dtype(torch::kInt32));
|
||||
torch::Tensor conics_d =
|
||||
@@ -221,6 +224,7 @@ project_gaussians_forward_tensor(
|
||||
cov3d_d.contiguous().data_ptr<float>(),
|
||||
(float2 *)xys_d.contiguous().data_ptr<float>(),
|
||||
depths_d.contiguous().data_ptr<float>(),
|
||||
(float2 *)depth_grads_d.contiguous().data_ptr<float>(),
|
||||
radii_d.contiguous().data_ptr<int>(),
|
||||
(float3 *)conics_d.contiguous().data_ptr<float>(),
|
||||
compensation_d.contiguous().data_ptr<float>(),
|
||||
@@ -228,7 +232,8 @@ project_gaussians_forward_tensor(
|
||||
);
|
||||
|
||||
return std::make_tuple(
|
||||
cov3d_d, xys_d, depths_d, radii_d, conics_d, compensation_d, num_tiles_hit_d
|
||||
cov3d_d, xys_d, depths_d, depth_grads_d,
|
||||
radii_d, conics_d, compensation_d, num_tiles_hit_d
|
||||
);
|
||||
}
|
||||
|
||||
@@ -257,6 +262,7 @@ project_gaussians_backward_tensor(
|
||||
torch::Tensor &compensation,
|
||||
torch::Tensor &v_xy,
|
||||
torch::Tensor &v_depth,
|
||||
torch::Tensor &v_depth_grad,
|
||||
torch::Tensor &v_conic,
|
||||
torch::Tensor &v_compensation
|
||||
){
|
||||
@@ -298,6 +304,7 @@ project_gaussians_backward_tensor(
|
||||
(float *)compensation.contiguous().data_ptr<float>(),
|
||||
(float2 *)v_xy.contiguous().data_ptr<float>(),
|
||||
v_depth.contiguous().data_ptr<float>(),
|
||||
(float2 *)v_depth_grad.contiguous().data_ptr<float>(),
|
||||
(float3 *)v_conic.contiguous().data_ptr<float>(),
|
||||
(float *)v_compensation.contiguous().data_ptr<float>(),
|
||||
// Outputs.
|
||||
@@ -375,7 +382,8 @@ torch::Tensor get_tile_bin_edges_tensor(
|
||||
return tile_bins;
|
||||
}
|
||||
|
||||
std::tuple<torch::Tensor, torch::Tensor, torch::Tensor>
|
||||
std::tuple<torch::Tensor, torch::Tensor,
|
||||
torch::Tensor, torch::Tensor, torch::Tensor>
|
||||
rasterize_forward_tensor(
|
||||
const std::tuple<int, int, int> tile_bounds,
|
||||
const std::tuple<int, int, int> block,
|
||||
@@ -383,6 +391,8 @@ rasterize_forward_tensor(
|
||||
const torch::Tensor &gaussian_ids_sorted,
|
||||
const torch::Tensor &tile_bins,
|
||||
const torch::Tensor &xys,
|
||||
const torch::Tensor &depths,
|
||||
const torch::Tensor &depth_grads,
|
||||
const torch::Tensor &conics,
|
||||
const torch::Tensor &colors,
|
||||
const torch::Tensor &opacities,
|
||||
@@ -419,6 +429,12 @@ rasterize_forward_tensor(
|
||||
torch::Tensor out_img = torch::zeros(
|
||||
{img_height, img_width, channels}, xys.options().dtype(torch::kFloat32)
|
||||
);
|
||||
torch::Tensor out_reg_depth = torch::zeros(
|
||||
{img_height, img_width}, xys.options().dtype(torch::kFloat32)
|
||||
);
|
||||
torch::Tensor out_reg_normal = torch::zeros(
|
||||
{img_height, img_width}, xys.options().dtype(torch::kFloat32)
|
||||
);
|
||||
torch::Tensor final_Ts = torch::zeros(
|
||||
{img_height, img_width}, xys.options().dtype(torch::kFloat32)
|
||||
);
|
||||
@@ -432,16 +448,21 @@ rasterize_forward_tensor(
|
||||
gaussian_ids_sorted.contiguous().data_ptr<int32_t>(),
|
||||
(int2 *)tile_bins.contiguous().data_ptr<int>(),
|
||||
(float2 *)xys.contiguous().data_ptr<float>(),
|
||||
depths.contiguous().data_ptr<float>(),
|
||||
(float2 *)depth_grads.contiguous().data_ptr<float>(),
|
||||
(float3 *)conics.contiguous().data_ptr<float>(),
|
||||
(float3 *)colors.contiguous().data_ptr<float>(),
|
||||
opacities.contiguous().data_ptr<float>(),
|
||||
final_Ts.contiguous().data_ptr<float>(),
|
||||
final_idx.contiguous().data_ptr<int>(),
|
||||
(float3 *)out_img.contiguous().data_ptr<float>(),
|
||||
out_reg_depth.contiguous().data_ptr<float>(),
|
||||
out_reg_normal.contiguous().data_ptr<float>(),
|
||||
*(float3 *)background.contiguous().data_ptr<float>()
|
||||
);
|
||||
|
||||
return std::make_tuple(out_img, final_Ts, final_idx);
|
||||
return std::make_tuple(out_img, out_reg_depth, out_reg_normal,
|
||||
final_Ts, final_idx);
|
||||
}
|
||||
|
||||
|
||||
@@ -460,6 +481,8 @@ std::
|
||||
const torch::Tensor &gaussians_ids_sorted,
|
||||
const torch::Tensor &tile_bins,
|
||||
const torch::Tensor &xys,
|
||||
const torch::Tensor &depths,
|
||||
const torch::Tensor &depth_grads,
|
||||
const torch::Tensor &conics,
|
||||
const torch::Tensor &colors,
|
||||
const torch::Tensor &opacities,
|
||||
|
||||
@@ -46,6 +46,7 @@ std::tuple<
|
||||
torch::Tensor,
|
||||
torch::Tensor,
|
||||
torch::Tensor,
|
||||
torch::Tensor,
|
||||
torch::Tensor>
|
||||
project_gaussians_forward_tensor(
|
||||
const int num_points,
|
||||
@@ -89,6 +90,7 @@ project_gaussians_backward_tensor(
|
||||
torch::Tensor &compensation,
|
||||
torch::Tensor &v_xy,
|
||||
torch::Tensor &v_depth,
|
||||
torch::Tensor &v_depth_grad,
|
||||
torch::Tensor &v_conic,
|
||||
torch::Tensor &v_compensation
|
||||
);
|
||||
@@ -112,6 +114,8 @@ torch::Tensor get_tile_bin_edges_tensor(
|
||||
);
|
||||
|
||||
std::tuple<
|
||||
torch::Tensor,
|
||||
torch::Tensor,
|
||||
torch::Tensor,
|
||||
torch::Tensor,
|
||||
torch::Tensor
|
||||
@@ -122,6 +126,8 @@ std::tuple<
|
||||
const torch::Tensor &gaussian_ids_sorted,
|
||||
const torch::Tensor &tile_bins,
|
||||
const torch::Tensor &xys,
|
||||
const torch::Tensor &depths,
|
||||
const torch::Tensor &depth_grads,
|
||||
const torch::Tensor &conics,
|
||||
const torch::Tensor &colors,
|
||||
const torch::Tensor &opacities,
|
||||
@@ -143,6 +149,8 @@ std::
|
||||
const torch::Tensor &gaussians_ids_sorted,
|
||||
const torch::Tensor &tile_bins,
|
||||
const torch::Tensor &xys,
|
||||
const torch::Tensor &depths,
|
||||
const torch::Tensor &depth_grads,
|
||||
const torch::Tensor &conics,
|
||||
const torch::Tensor &colors,
|
||||
const torch::Tensor &opacities,
|
||||
|
||||
@@ -25,6 +25,7 @@ __global__ void project_gaussians_forward_kernel(
|
||||
float* __restrict__ covs3d,
|
||||
float2* __restrict__ xys,
|
||||
float* __restrict__ depths,
|
||||
float2* __restrict__ depth_grads,
|
||||
int* __restrict__ radii,
|
||||
float3* __restrict__ conics,
|
||||
float* __restrict__ compensation,
|
||||
@@ -48,7 +49,7 @@ __global__ void project_gaussians_forward_kernel(
|
||||
// printf("p_view %d %.2f %.2f %.2f\n", idx, p_view.x, p_view.y, p_view.z);
|
||||
|
||||
// compute the projected covariance
|
||||
float3 scale = make_float3(scales[idx].x, scales[idx].y, 0.0f);
|
||||
float3 scale = { scales[idx].x, scales[idx].y, 0.0f };
|
||||
float4 quat = quats[idx];
|
||||
// printf("%d scale %.2f %.2f %.2f\n", idx, scale.x, scale.y, scale.z);
|
||||
// printf("%d quat %.2f %.2f %.2f %.2f\n", idx, quat.w, quat.x, quat.y,
|
||||
@@ -89,8 +90,12 @@ __global__ void project_gaussians_forward_kernel(
|
||||
return;
|
||||
}
|
||||
|
||||
// compute the depth gradient
|
||||
float2 depth_grad = projected_depth_grad(viewmat, fx, fy, quat, p_view);
|
||||
|
||||
num_tiles_hit[idx] = tile_area;
|
||||
depths[idx] = p_view.z;
|
||||
depth_grads[idx] = {depth_grad.x, depth_grad.y};
|
||||
radii[idx] = (int)radius;
|
||||
xys[idx] = center;
|
||||
compensation[idx] = comp;
|
||||
@@ -176,12 +181,16 @@ __global__ void rasterize_forward(
|
||||
const int32_t* __restrict__ gaussian_ids_sorted,
|
||||
const int2* __restrict__ tile_bins,
|
||||
const float2* __restrict__ xys,
|
||||
const float* __restrict__ depths,
|
||||
const float2* __restrict__ depth_grads,
|
||||
const float3* __restrict__ conics,
|
||||
const float3* __restrict__ colors,
|
||||
const float* __restrict__ opacities,
|
||||
float* __restrict__ final_Ts,
|
||||
int* __restrict__ final_index,
|
||||
float3* __restrict__ out_img,
|
||||
float* __restrict__ out_reg_depth,
|
||||
float* __restrict__ out_reg_normal,
|
||||
const float3& __restrict__ background
|
||||
) {
|
||||
// each thread draws one pixel, but also timeshares caching gaussians in a
|
||||
@@ -195,8 +204,7 @@ __global__ void rasterize_forward(
|
||||
unsigned j =
|
||||
block.group_index().x * block.group_dim().x + block.thread_index().x;
|
||||
|
||||
float px = (float)j + 0.5;
|
||||
float py = (float)i + 0.5;
|
||||
float2 p = { (float)j + 0.5f, (float)i + 0.5f };
|
||||
int32_t pix_id = i * img_size.x + j;
|
||||
|
||||
// return if out of bounds
|
||||
@@ -214,9 +222,9 @@ __global__ void rasterize_forward(
|
||||
__shared__ int32_t id_batch[MAX_BLOCK_SIZE];
|
||||
__shared__ float3 xy_opacity_batch[MAX_BLOCK_SIZE];
|
||||
__shared__ float3 conic_batch[MAX_BLOCK_SIZE];
|
||||
__shared__ float3 depth_grad_batch[MAX_BLOCK_SIZE];
|
||||
|
||||
// current visibility left to render
|
||||
float T = 1.f;
|
||||
// index of most recent gaussian to write to this thread's pixel
|
||||
int cur_idx = 0;
|
||||
|
||||
@@ -224,14 +232,16 @@ __global__ void rasterize_forward(
|
||||
// each thread loads one gaussian at a time before rasterizing its
|
||||
// designated pixel
|
||||
int tr = block.thread_rank();
|
||||
float3 pix_out = {0.f, 0.f, 0.f};
|
||||
float T = 1.f; // current/total visibility
|
||||
float sum_vis = 0.f; // sum of visibilities
|
||||
float2 sum_depth_grad = {0.f, 0.f}; // sum of "normals"
|
||||
float3 pix_out = {0.f, 0.f, 0.f}; // output radiance
|
||||
for (int b = 0; b < num_batches; ++b) {
|
||||
// resync all threads before beginning next batch
|
||||
// end early if entire tile is done
|
||||
if (__syncthreads_count(done) >= block_size) {
|
||||
break;
|
||||
}
|
||||
|
||||
// each thread fetch 1 gaussian from front to back
|
||||
// index of gaussian to load
|
||||
int batch_start = range.x + block_size * b;
|
||||
@@ -243,44 +253,44 @@ __global__ void rasterize_forward(
|
||||
const float opac = opacities[g_id];
|
||||
xy_opacity_batch[tr] = {xy.x, xy.y, opac};
|
||||
conic_batch[tr] = conics[g_id];
|
||||
const float2 depth_grad = depth_grads[g_id];
|
||||
depth_grad_batch[tr] = {depths[g_id], depth_grad.x, depth_grad.y};
|
||||
}
|
||||
|
||||
// wait for other threads to collect the gaussians in batch
|
||||
block.sync();
|
||||
|
||||
// process gaussians in the current batch for this pixel
|
||||
int batch_size = min(block_size, range.y - batch_start);
|
||||
for (int t = 0; (t < batch_size) && !done; ++t) {
|
||||
const float3 conic = conic_batch[t];
|
||||
const float3 xy_opac = xy_opacity_batch[t];
|
||||
const float opac = xy_opac.z;
|
||||
const float2 delta = {xy_opac.x - px, xy_opac.y - py};
|
||||
const float sigma = 0.5f * (conic.x * delta.x * delta.x +
|
||||
conic.z * delta.y * delta.y) +
|
||||
conic.y * delta.x * delta.y;
|
||||
const float alpha = min(0.999f, opac * __expf(-sigma));
|
||||
if (sigma < 0.f || alpha < 1.f / 255.f) {
|
||||
float alpha;
|
||||
if (!get_alpha(conic_batch[t], xy_opacity_batch[t], p, alpha))
|
||||
continue;
|
||||
}
|
||||
|
||||
const float next_T = T * (1.f - alpha);
|
||||
if (next_T <= 1e-4f) { // this pixel is done
|
||||
// we want to render the last gaussian that contributes and note
|
||||
// that here idx > range.x so we don't underflow
|
||||
done = true;
|
||||
break;
|
||||
}
|
||||
|
||||
int32_t g = id_batch[t];
|
||||
const float vis = alpha * T;
|
||||
const float3 c = colors[g];
|
||||
const float3 depth_grad_batch_t = depth_grad_batch[t];
|
||||
const float2 depth_grad = {depth_grad_batch_t.y, depth_grad_batch_t.z};
|
||||
pix_out.x = pix_out.x + c.x * vis;
|
||||
pix_out.y = pix_out.y + c.y * vis;
|
||||
pix_out.z = pix_out.z + c.z * vis;
|
||||
sum_vis += vis;
|
||||
sum_depth_grad.x = sum_depth_grad.x + vis * depth_grad.x;
|
||||
sum_depth_grad.y = sum_depth_grad.y + vis * depth_grad.y;
|
||||
T = next_T;
|
||||
cur_idx = batch_start + t;
|
||||
if (T <= 1e-4f) {
|
||||
done = true;
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
float sum_depth_grad_norm = hypot(sum_depth_grad.x, sum_depth_grad.y) + 1e-6f;
|
||||
float2 mean_depth_grad = {
|
||||
sum_depth_grad.x / sum_depth_grad_norm,
|
||||
sum_depth_grad.y / sum_depth_grad_norm
|
||||
};
|
||||
|
||||
if (inside) {
|
||||
// add background
|
||||
@@ -293,6 +303,81 @@ __global__ void rasterize_forward(
|
||||
final_color.z = pix_out.z + T * background.z;
|
||||
out_img[pix_id] = final_color;
|
||||
}
|
||||
|
||||
// calculate regularization weights
|
||||
done = !inside;
|
||||
float reg_depth = 0.f;
|
||||
float reg_normal = 0.f;
|
||||
T = 1.0f;
|
||||
float cur_vis = 0.0f;
|
||||
for (int b = 0; b < num_batches; ++b) {
|
||||
// resync all threads before beginning next batch
|
||||
// end early if entire tile is done
|
||||
if (__syncthreads_count(done) >= block_size) {
|
||||
break;
|
||||
}
|
||||
// each thread fetch 1 gaussian from front to back
|
||||
// index of gaussian to load
|
||||
int batch_start = range.x + block_size * b;
|
||||
int idx = batch_start + tr;
|
||||
if (idx < range.y) {
|
||||
int32_t g_id = gaussian_ids_sorted[idx];
|
||||
id_batch[tr] = g_id;
|
||||
const float2 xy = xys[g_id];
|
||||
const float opac = opacities[g_id];
|
||||
xy_opacity_batch[tr] = {xy.x, xy.y, opac};
|
||||
conic_batch[tr] = conics[g_id];
|
||||
const float2 depth_grad = depth_grads[g_id];
|
||||
depth_grad_batch[tr] = {depths[g_id], depth_grad.x, depth_grad.y};
|
||||
}
|
||||
// wait for other threads to collect the gaussians in batch
|
||||
block.sync();
|
||||
|
||||
// process gaussians in the current batch for this pixel
|
||||
int batch_size = min(block_size, range.y - batch_start);
|
||||
for (int t = 0; (t < batch_size) && !done; ++t) {
|
||||
float alpha;
|
||||
if (!get_alpha(conic_batch[t], xy_opacity_batch[t], p, alpha))
|
||||
continue;
|
||||
const float next_T = T * (1.f - alpha);
|
||||
|
||||
const float vis = alpha * T;
|
||||
float cur_vis_next = cur_vis + vis;
|
||||
|
||||
const float3 depth_grad_batch_t = depth_grad_batch[t];
|
||||
const float depth = depth_grad_batch_t.x;
|
||||
const float2 depth_grad = {depth_grad_batch_t.y, depth_grad_batch_t.z};
|
||||
float depth_grad_norm = hypot(depth_grad.x, depth_grad.y) + 1e-6f;
|
||||
float2 depth_grad_normalized = {
|
||||
depth_grad.x / depth_grad_norm,
|
||||
depth_grad.y / depth_grad_norm
|
||||
};
|
||||
|
||||
// depth regularization:
|
||||
// 1/2 \sum[i,j] w[i] w[j] |z[j]-z[i]| =
|
||||
// \sum[i] w[i] z[i] ( \sum [j<i] w[j] - \sum[j>i] w[j] )
|
||||
reg_depth += vis * depth * (cur_vis - (sum_vis-cur_vis_next));
|
||||
|
||||
// normal regularization
|
||||
reg_normal += vis * (1.0f - (
|
||||
depth_grad_normalized.x * mean_depth_grad.x +
|
||||
depth_grad_normalized.y * mean_depth_grad.y
|
||||
));
|
||||
|
||||
cur_vis = cur_vis_next;
|
||||
T = next_T;
|
||||
cur_idx = batch_start + t;
|
||||
if (T <= 1e-4f) {
|
||||
done = true;
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if (inside) {
|
||||
out_reg_depth[pix_id] = reg_depth;
|
||||
out_reg_normal[pix_id] = reg_normal;
|
||||
}
|
||||
}
|
||||
|
||||
// device helper to approximate projected 2d cov from 3d mean and cov
|
||||
@@ -371,6 +456,9 @@ __device__ void project_cov3d_ewa(
|
||||
cov2d.z = c11 + 0.3f;
|
||||
float det_blur = cov2d.x * cov2d.z - cov2d.y * cov2d.y;
|
||||
compensation = std::sqrt(std::max(0.f, det_orig / det_blur));
|
||||
|
||||
// depth to pixel gradient
|
||||
// TO-DO
|
||||
}
|
||||
|
||||
// device helper to get 3D covariance from scale and quat parameters
|
||||
@@ -395,3 +483,24 @@ __device__ void scale_rot_to_cov3d(
|
||||
cov3d[4] = tmp[1][2];
|
||||
cov3d[5] = tmp[2][2];
|
||||
}
|
||||
|
||||
// device helper to get screen space depth gradient
|
||||
__device__ float2 projected_depth_grad(
|
||||
const float* viewmat, const float fx, const float fy,
|
||||
const float4 quat, const float3 p_view
|
||||
) {
|
||||
glm::mat3 R = glm::transpose(glm::mat3(
|
||||
viewmat[0], viewmat[1], viewmat[2],
|
||||
viewmat[4], viewmat[5], viewmat[6],
|
||||
viewmat[8], viewmat[9], viewmat[10]
|
||||
)) * quat_to_rotmat(quat);
|
||||
glm::vec3 n = glm::vec3(R[2][0], R[2][1], R[2][2]);
|
||||
glm::vec3 p = glm::vec3(p_view.x, p_view.y, p_view.z);
|
||||
glm::mat3 invJ = glm::mat3(
|
||||
p.z/fx, 0.0f, 0.0f,
|
||||
0.0f, p.z/fy, 0.0f,
|
||||
p.x/p.z, p.y/p.z, 1.0f
|
||||
);
|
||||
n = glm::transpose(invJ) * n;
|
||||
return { -n.x/n.z, -n.y/n.z };
|
||||
}
|
||||
|
||||
@@ -18,6 +18,7 @@ __global__ void project_gaussians_forward_kernel(
|
||||
float* __restrict__ covs3d,
|
||||
float2* __restrict__ xys,
|
||||
float* __restrict__ depths,
|
||||
float2* __restrict__ depth_grads,
|
||||
int* __restrict__ radii,
|
||||
float3* __restrict__ conics,
|
||||
float* __restrict__ compensation,
|
||||
@@ -31,12 +32,16 @@ __global__ void rasterize_forward(
|
||||
const int32_t* __restrict__ gaussian_ids_sorted,
|
||||
const int2* __restrict__ tile_bins,
|
||||
const float2* __restrict__ xys,
|
||||
const float* __restrict__ depths,
|
||||
const float2* __restrict__ depth_grads,
|
||||
const float3* __restrict__ conics,
|
||||
const float3* __restrict__ colors,
|
||||
const float* __restrict__ opacities,
|
||||
float* __restrict__ final_Ts,
|
||||
int* __restrict__ final_index,
|
||||
float3* __restrict__ out_img,
|
||||
float* __restrict__ out_reg_depth,
|
||||
float* __restrict__ out_reg_normal,
|
||||
const float3& __restrict__ background
|
||||
);
|
||||
|
||||
@@ -58,6 +63,11 @@ __device__ void scale_rot_to_cov3d(
|
||||
const float3 scale, const float glob_scale, const float4 quat, float *cov3d
|
||||
);
|
||||
|
||||
__device__ float2 projected_depth_grad(
|
||||
const float* viewmat, const float fx, const float fy,
|
||||
const float4 quat, const float3 p_view
|
||||
);
|
||||
|
||||
__global__ void map_gaussian_to_intersects(
|
||||
const int num_points,
|
||||
const float2* __restrict__ xys,
|
||||
|
||||
@@ -5,7 +5,7 @@
|
||||
#include <iostream>
|
||||
|
||||
|
||||
#if 0
|
||||
#if 1
|
||||
|
||||
// "Gaussian" kernel
|
||||
inline __device__ float visibility_kernel(const float r2) {
|
||||
@@ -51,6 +51,20 @@ inline __device__ float visibility_kernel_radius() {
|
||||
#endif
|
||||
|
||||
|
||||
inline __device__ bool get_alpha(
|
||||
const float3 conic, const float3 xy_opac, const float2 p,
|
||||
float &alpha
|
||||
) {
|
||||
const float opac = xy_opac.z;
|
||||
const float2 delta = {xy_opac.x - p.x, xy_opac.y - p.y};
|
||||
const float r2 = 0.5f * (conic.x * delta.x * delta.x +
|
||||
conic.z * delta.y * delta.y) +
|
||||
conic.y * delta.x * delta.y;
|
||||
alpha = min(0.999f, opac * visibility_kernel(r2));
|
||||
return r2 >= 0.f && alpha >= 1.f / 255.f;
|
||||
}
|
||||
|
||||
|
||||
inline __device__ void get_bbox(
|
||||
const float2 center,
|
||||
const float2 dims,
|
||||
@@ -173,11 +187,11 @@ inline __device__ float4 transform_4x4(const float *mat, const float3 p) {
|
||||
}
|
||||
|
||||
inline __device__ float2 project_pix(
|
||||
const float2 fxfy, const float3 p_view, const float2 pp
|
||||
const float2 f, const float3 p_view, const float2 c
|
||||
) {
|
||||
float rw = 1.f / (p_view.z + 1e-6f);
|
||||
float2 p_proj = { p_view.x * rw, p_view.y * rw };
|
||||
float2 p_pix = { p_proj.x * fxfy.x + pp.x, p_proj.y * fxfy.y + pp.y };
|
||||
float2 p_pix = { p_proj.x * f.x + c.x, p_proj.y * f.y + c.y };
|
||||
return p_pix;
|
||||
}
|
||||
|
||||
|
||||
@@ -49,7 +49,8 @@ def project_gaussians(
|
||||
A tuple of {Tensor, Tensor, Tensor, Tensor, Tensor, Tensor, Tensor}:
|
||||
|
||||
- **xys** (Tensor): x,y locations of 2D gaussian projections.
|
||||
- **depths** (Tensor): z depth of gaussians.
|
||||
- **depths** (Tensor): z depth of gaussians at the center.
|
||||
- **depth_grads** (Tensor): xy gradient of z depth of gaussians.
|
||||
- **radii** (Tensor): radii of 2D gaussian projections.
|
||||
- **conics** (Tensor): conic parameters for 2D gaussian.
|
||||
- **compensation** (Tensor): the density compensation for blurring 2D kernel
|
||||
@@ -103,6 +104,7 @@ class _ProjectGaussians(Function):
|
||||
cov3d,
|
||||
xys,
|
||||
depths,
|
||||
depth_grads,
|
||||
radii,
|
||||
conics,
|
||||
compensation,
|
||||
@@ -146,13 +148,14 @@ class _ProjectGaussians(Function):
|
||||
compensation,
|
||||
)
|
||||
|
||||
return (xys, depths, radii, conics, compensation, num_tiles_hit, cov3d)
|
||||
return (xys, depths, depth_grads, radii, conics, compensation, num_tiles_hit, cov3d)
|
||||
|
||||
@staticmethod
|
||||
def backward(
|
||||
ctx,
|
||||
v_xys,
|
||||
v_depths,
|
||||
v_depth_grads,
|
||||
v_radii,
|
||||
v_conics,
|
||||
v_compensation,
|
||||
@@ -189,6 +192,7 @@ class _ProjectGaussians(Function):
|
||||
compensation,
|
||||
v_xys,
|
||||
v_depths,
|
||||
v_depth_grads,
|
||||
v_conics,
|
||||
v_compensation,
|
||||
)
|
||||
|
||||
@@ -15,6 +15,7 @@ from .utils import bin_and_sort_gaussians, compute_cumulative_intersects
|
||||
def rasterize_gaussians(
|
||||
xys: Float[Tensor, "*batch 2"],
|
||||
depths: Float[Tensor, "*batch 1"],
|
||||
depth_grads: Float[Tensor, "*batch 2"],
|
||||
radii: Float[Tensor, "*batch 1"],
|
||||
conics: Float[Tensor, "*batch 3"],
|
||||
num_tiles_hit: Int[Tensor, "*batch 1"],
|
||||
@@ -74,6 +75,7 @@ def rasterize_gaussians(
|
||||
return _RasterizeGaussians.apply(
|
||||
xys.contiguous(),
|
||||
depths.contiguous(),
|
||||
depth_grads.contiguous(),
|
||||
radii.contiguous(),
|
||||
conics.contiguous(),
|
||||
num_tiles_hit.contiguous(),
|
||||
@@ -95,6 +97,7 @@ class _RasterizeGaussians(Function):
|
||||
ctx,
|
||||
xys: Float[Tensor, "*batch 2"],
|
||||
depths: Float[Tensor, "*batch 1"],
|
||||
depth_grads: Float[Tensor, "*batch 2"],
|
||||
radii: Float[Tensor, "*batch 1"],
|
||||
conics: Float[Tensor, "*batch 3"],
|
||||
num_tiles_hit: Int[Tensor, "*batch 1"],
|
||||
@@ -122,6 +125,8 @@ class _RasterizeGaussians(Function):
|
||||
torch.ones(img_height, img_width, colors.shape[-1], device=xys.device)
|
||||
* background
|
||||
)
|
||||
out_reg_depth = torch.zeros(img_height, img_width, device=xys.device)
|
||||
out_reg_normal = torch.zeros(img_height, img_width, device=xys.device)
|
||||
gaussian_ids_sorted = torch.zeros(0, 1, device=xys.device)
|
||||
tile_bins = torch.zeros(0, 2, device=xys.device)
|
||||
final_Ts = torch.zeros(img_height, img_width, device=xys.device)
|
||||
@@ -144,15 +149,16 @@ class _RasterizeGaussians(Function):
|
||||
block_width,
|
||||
)
|
||||
assert colors.shape[-1] == 3
|
||||
rasterize_fn = _C.rasterize_forward
|
||||
|
||||
out_img, final_Ts, final_idx = rasterize_fn(
|
||||
out_img, out_reg_depth, out_reg_normal, final_Ts, final_idx = _C.rasterize_forward(
|
||||
tile_bounds,
|
||||
block,
|
||||
img_size,
|
||||
gaussian_ids_sorted,
|
||||
tile_bins,
|
||||
xys,
|
||||
depths,
|
||||
depth_grads,
|
||||
conics,
|
||||
colors,
|
||||
opacity,
|
||||
@@ -167,6 +173,8 @@ class _RasterizeGaussians(Function):
|
||||
gaussian_ids_sorted,
|
||||
tile_bins,
|
||||
xys,
|
||||
depths,
|
||||
depth_grads,
|
||||
conics,
|
||||
colors,
|
||||
opacity,
|
||||
@@ -175,14 +183,11 @@ class _RasterizeGaussians(Function):
|
||||
final_idx,
|
||||
)
|
||||
|
||||
if return_alpha:
|
||||
out_alpha = 1 - final_Ts
|
||||
return out_img, out_alpha
|
||||
else:
|
||||
return out_img
|
||||
out_alpha = 1.0 - final_Ts
|
||||
return out_img, out_reg_depth, out_reg_normal, out_alpha
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, v_out_img, v_out_alpha=None):
|
||||
def backward(ctx, v_out_img, v_out_reg_depth, v_out_reg_normal, v_out_alpha=None):
|
||||
img_height = ctx.img_height
|
||||
img_width = ctx.img_width
|
||||
num_intersects = ctx.num_intersects
|
||||
@@ -194,6 +199,8 @@ class _RasterizeGaussians(Function):
|
||||
gaussian_ids_sorted,
|
||||
tile_bins,
|
||||
xys,
|
||||
depths,
|
||||
depth_grads,
|
||||
conics,
|
||||
colors,
|
||||
opacity,
|
||||
@@ -211,14 +218,15 @@ class _RasterizeGaussians(Function):
|
||||
|
||||
else:
|
||||
assert colors.shape[-1] == 3
|
||||
rasterize_fn = _C.rasterize_backward
|
||||
v_xy, v_xy_abs, v_conic, v_colors, v_opacity = rasterize_fn(
|
||||
v_xy, v_xy_abs, v_conic, v_colors, v_opacity = _C.rasterize_backward(
|
||||
img_height,
|
||||
img_width,
|
||||
ctx.block_width,
|
||||
gaussian_ids_sorted,
|
||||
tile_bins,
|
||||
xys,
|
||||
depths,
|
||||
depth_grads,
|
||||
conics,
|
||||
colors,
|
||||
opacity,
|
||||
@@ -237,6 +245,7 @@ class _RasterizeGaussians(Function):
|
||||
return (
|
||||
v_xy, # xys
|
||||
None, # depths
|
||||
None, # depth_grads
|
||||
None, # radii
|
||||
v_conic, # conics
|
||||
None, # num_tiles_hit
|
||||
|
||||
Reference in New Issue
Block a user