regularization (forward)

This commit is contained in:
harry7557558
2024-05-08 20:13:13 -04:00
parent 28c3fa4bbe
commit ead5e839eb
11 changed files with 271 additions and 70 deletions
+2
View File
@@ -1,3 +1,5 @@
.vscode/
outputs/
# Byte-compiled / optimized / DLL files
+45 -27
View File
@@ -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
+4 -1
View File
@@ -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,
+26 -3
View File
@@ -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,
+133 -24
View File
@@ -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,
+17 -3
View File
@@ -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;
}
+6 -2
View File
@@ -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,
)
+19 -10
View File
@@ -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