investigate and fix splat quality regression

This commit is contained in:
Harry Chen
2026-06-07 21:42:40 -04:00
parent f7aa6758da
commit f8bfd123a5
20 changed files with 222 additions and 117 deletions
+9 -5
View File
@@ -520,11 +520,15 @@ class Trainer:
input_intrins_list, input_dist_coeffs_list,
train_indices, val_indices)
# When the warp path is active, the bilagrid table is sized to the
# POST-split camera count -- override the model's num_train_data
# accordingly (the model's own pre-init guess uses a uniform K).
if any_warp:
self.model.num_train_data = int(n_post)
# The bilagrid (+ PPISP, + camera-optimizer) tables are sized off
# `self.model.num_train_data`. The model's pre-init uses a uniform
# K=6 upper bound for the C++ datamanager path since it has no
# per-camera K info at __init__ time; resolve it to the real
# post-split count now. For non-warp datasets (n_post == N), this
# shrinks the table back to N -- without this, bilagrid `n_grids`
# is 6x too large and the TV-loss normalization (~1/N_grids) makes
# the TV regularizer 6x weaker than intended.
self.model.num_train_data = int(n_post)
# Build a POST-split Cameras object for the viewer. For the warp
# path each input image is exposed as K=5 (fisheye/equisolid) or
@@ -367,12 +367,15 @@ __global__ void bilagrid_depth_uniform_sample_backward_v1_kernel_depth(
float x = (float)wi / (float)(w-1) * (float)(W-1);
float y = (float)hi / (float)(h-1) * (float)(H-1);
#endif
float z = (sr / (sr + 1.0f)) * (L-1);
// Clamp gz to [0,1] -- matches forward and bilagrid-grad bwd branch.
float gz_raw = sr / (sr + 1.0f);
float gz = fminf(fmaxf(gz_raw, 0.0f), 1.0f);
bool gz_in_range = (gz_raw >= 0.0f && gz_raw <= 1.0f);
float z = gz * (L-1);
int x0 = floorf(x), y0 = floorf(y), z0 = floorf(z);
int x1 = min(x0+1, W-1);
int y1 = min(y0+1, H-1);
int z1 = z0 + 1;
z0 = min(max(z0,0), L-1); z1 = min(max(z1,0), L-1);
int z1 = min(z0+1, L-1);
// fractional parts
float fx = x - x0, fy = y - y0, fz = z - z0;
@@ -446,6 +449,8 @@ __global__ void bilagrid_depth_uniform_sample_backward_v1_kernel_depth(
}
gz_grad += dwdz[corner] * (L-1) * trilerp;
}
// Zero gz_grad outside the [0,1] clamp range -- the clamp's vjp.
if (!gz_in_range) gz_grad = 0.0f;
vr += gz_grad / ((sr+1.0f) * (sr+1.0f));
vr *= scalar;
v_depth[g_off] = isfinite(vr) ? vr : 0.0f;
@@ -80,6 +80,9 @@ __global__ void bilagrid_depth_uniform_sample_forward_kernel(
float gy = (float)hi / (float)(h-1);
#endif
float gz = sr / (sr + 1.0f);
// Clamp gz to [0,1] -- equivalent to clipping z to [0, L-1]. Matches the
// backward branch (which clamps), so fz here is in [0,1] same as bwd's.
gz = fminf(fmaxf(gz, 0.0f), 1.0f);
float x = gx * (W - 1);
float y = gy * (H - 1);
float z = gz * (L - 1);
@@ -90,9 +93,7 @@ __global__ void bilagrid_depth_uniform_sample_forward_kernel(
int z0 = (int)floorf(z);
int x1 = min(x0+1, W-1);
int y1 = min(y0+1, H-1);
int z1 = z0 + 1;
z0 = min(max(z0, 0), L-1);
z1 = min(max(z1, 0), L-1);
int z1 = min(z0+1, L-1);
// interpolation parameters
float fx = x - (float)x0;
@@ -356,12 +356,17 @@ __global__ void bilagrid_loglinear_uniform_sample_backward_v1_kernel_rgb(
float x = (float)wi / (float)(w-1) * (float)(W-1);
float y = (float)hi / (float)(h-1) * (float)(H-1);
#endif
float z = (kC2G_r * sr + kC2G_g * sg + kC2G_b * sb) * (L-1);
// Clamp gz to [0,1] -- matches forward and the bilagrid-grad bwd branch.
// Track whether gz was in range so we can zero gz_grad below (dz/dgz = 0
// outside [0,1] -> the chain through z back to rgb is zero there).
float gz_raw = kC2G_r * sr + kC2G_g * sg + kC2G_b * sb;
float gz = fminf(fmaxf(gz_raw, 0.0f), 1.0f);
bool gz_in_range = (gz_raw >= 0.0f && gz_raw <= 1.0f);
float z = gz * (L-1);
int x0 = floorf(x), y0 = floorf(y), z0 = floorf(z);
int x1 = min(x0+1, W-1);
int y1 = min(y0+1, H-1);
int z1 = z0 + 1;
z0 = min(max(z0,0), L-1); z1 = min(max(z1,0), L-1);
int z1 = min(z0+1, L-1);
// fractional parts
float fx = x - x0, fy = y - y0, fz = z - z0;
@@ -441,6 +446,10 @@ __global__ void bilagrid_loglinear_uniform_sample_backward_v1_kernel_rgb(
}
gz_grad += dwdz[corner] * (L-1) * trilerp;
}
// Zero gz_grad outside the [0,1] clamp range -- the clamp's vjp is the
// indicator that gz_raw was strictly interior. Matches the forward's
// gz clamp.
if (!gz_in_range) gz_grad = 0.0f;
vr += kC2G_r * gz_grad;
vg += kC2G_g * gz_grad;
vb += kC2G_b * gz_grad;
@@ -71,6 +71,9 @@ __global__ void bilagrid_loglinear_uniform_sample_forward_kernel(
float gy = (float)hi / (float)(h-1);
#endif
float gz = kC2G_r * sr + kC2G_g * sg + kC2G_b * sb;
// Clamp gz to [0,1] -- equivalent to clipping z to [0, L-1]. Matches the
// backward branch (which clamps), so fz here is in [0,1] same as bwd's.
gz = fminf(fmaxf(gz, 0.0f), 1.0f);
float x = gx * (W - 1);
float y = gy * (H - 1);
float z = gz * (L - 1);
@@ -81,9 +84,7 @@ __global__ void bilagrid_loglinear_uniform_sample_forward_kernel(
int z0 = (int)floorf(z);
int x1 = min(x0+1, W-1);
int y1 = min(y0+1, H-1);
int z1 = z0 + 1;
z0 = min(max(z0, 0), L-1);
z1 = min(max(z1, 0), L-1);
int z1 = min(z0+1, L-1);
// interpolation parameters
float fx = x - (float)x0;
@@ -342,11 +342,14 @@ __global__ void bilagrid_normal_uniform_sample_backward_v1_kernel_normal(
#else
int g_off = ((ni * h + hi) * w + wi) * 3;
#endif
// Keep sr/sg/sb as the RAW normal (no in-place normalization) and build a
// unit_normal explicitly where needed. Matches the forward kernel's style
// and the bilagrid-grad bwd branch, which both apply `* inv_norm` lazily.
float sr = normal_in[g_off+0];
float sg = normal_in[g_off+1];
float sb = normal_in[g_off+2];
float inv_norm = rsqrtf(sr*sr + sg*sg + sb*sb + 1e-20f);
sr *= inv_norm, sg *= inv_norm, sb *= inv_norm;
float3 unit_normal = {sr*inv_norm, sg*inv_norm, sb*inv_norm};
float dr = v_normal_out[g_off+0];
float dg = v_normal_out[g_off+1];
float db = v_normal_out[g_off+2];
@@ -360,13 +363,16 @@ __global__ void bilagrid_normal_uniform_sample_backward_v1_kernel_normal(
float x = (float)wi / (float)(w-1) * (float)(W-1);
float y = (float)hi / (float)(h-1) * (float)(H-1);
#endif
// float z = (acosf(fminf(fmaxf(sb, -1.0f), 1.0f)) * (1.0f / (float)M_PI)) * (L-1);
float z = (0.5f + 0.5f * sb) * (L-1);
// float z = (acosf(fminf(fmaxf(sb*inv_norm, -1.0f), 1.0f)) * (1.0f / (float)M_PI)) * (L-1);
// Clamp gz to [0,1] -- matches forward and bilagrid-grad bwd branch.
float gz_raw = 0.5f + 0.5f * sb * inv_norm;
float gz = fminf(fmaxf(gz_raw, 0.0f), 1.0f);
bool gz_in_range = (gz_raw >= 0.0f && gz_raw <= 1.0f);
float z = gz * (L-1);
int x0 = floorf(x), y0 = floorf(y), z0 = floorf(z);
int x1 = min(x0+1, W-1);
int y1 = min(y0+1, H-1);
int z1 = z0 + 1;
z0 = min(max(z0,0), L-1); z1 = min(max(z1,0), L-1);
int z1 = min(z0+1, L-1);
float fx = x-x0, fy = y-y0, fz = z-z0;
@@ -407,11 +413,11 @@ __global__ void bilagrid_normal_uniform_sample_backward_v1_kernel_normal(
(ci == 0 ? axis_angle.x : ci == 1 ? axis_angle.y : axis_angle.z) = val;
}
// apply normal
float3 normal = {sr, sg, sb};
// apply normal -- rotate the UNIT normal, mirroring the forward.
float3 grad_axis_angle = {0.0f, 0.0f, 0.0f};
float3 grad_normal = {0.0f, 0.0f, 0.0f};
axis_angle_rotate_bwd(axis_angle, normal, {dr, dg, db}, grad_axis_angle, grad_normal);
float3 grad_unit_normal = {0.0f, 0.0f, 0.0f};
axis_angle_rotate_bwd(axis_angle, unit_normal, {dr, dg, db},
grad_axis_angle, grad_unit_normal);
// spatial derivatives for coords
float dwdz[8] = {
@@ -441,10 +447,18 @@ __global__ void bilagrid_normal_uniform_sample_backward_v1_kernel_normal(
}
gz_grad += dwdz[corner] * (L-1) * trilerp;
}
// grad_normal.z += gz_grad * -rsqrtf(fmaxf(1.0f - sb*sb, 1e-20f)) * (1.0f / (float)M_PI);
grad_normal.z += 0.5f * gz_grad;
// Zero gz_grad outside the [0,1] clamp range -- the clamp's vjp.
if (!gz_in_range) gz_grad = 0.0f;
// d(gz)/d(unit_normal.z) = 0.5 since gz = 0.5 + 0.5*unit_normal.z.
grad_unit_normal.z += 0.5f * gz_grad;
grad_normal = mul3(add3(grad_normal, mul3(normal, -dot3(normal, grad_normal))), inv_norm);
// Convert grad_unit_normal -> grad_raw_normal via the vjp of n_raw -> unit:
// d_unit/d_raw_j = (delta_ij - u_i*u_j) / |n|
// so grad_raw = (grad_unit - unit_normal * dot(unit_normal, grad_unit)) * inv_norm.
float3 grad_normal = mul3(
add3(grad_unit_normal,
mul3(unit_normal, -dot3(unit_normal, grad_unit_normal))),
inv_norm);
v_normal_in[g_off+0] = isfinite(grad_normal.x) ? grad_normal.x : 0.0f;
v_normal_in[g_off+1] = isfinite(grad_normal.y) ? grad_normal.y : 0.0f;
v_normal_in[g_off+2] = isfinite(grad_normal.z) ? grad_normal.z : 0.0f;
@@ -74,6 +74,9 @@ __global__ void bilagrid_normal_uniform_sample_forward_kernel(
#endif
// float gz = acosf(fminf(fmaxf(sb * inv_norm, -1.0f), 1.0f)) * (1.0f / (float)M_PI);
float gz = 0.5f + 0.5f * sb * inv_norm;
// Clamp gz to [0,1] -- equivalent to clipping z to [0, L-1]. Matches the
// backward branch (which clamps), so fz here is in [0,1] same as bwd's.
gz = fminf(fmaxf(gz, 0.0f), 1.0f);
float x = gx * (W - 1);
float y = gy * (H - 1);
float z = gz * (L - 1);
@@ -84,9 +87,7 @@ __global__ void bilagrid_normal_uniform_sample_forward_kernel(
int z0 = (int)floorf(z);
int x1 = min(x0+1, W-1);
int y1 = min(y0+1, H-1);
int z1 = z0 + 1;
z0 = min(max(z0, 0), L-1);
z1 = min(max(z1, 0), L-1);
int z1 = min(z0+1, L-1);
// interpolation parameters
float fx = x - (float)x0;
@@ -75,7 +75,10 @@ __global__ void bilagrid_ppisp_sample_backward_kernel(
float sr = rgb_in[3*g_off+0], sg = rgb_in[3*g_off+1], sb = rgb_in[3*g_off+2];
float gx = coords[2*g_off+0];
float gy = coords[2*g_off+1];
float gz = kC2G_r * sr + kC2G_g * sg + kC2G_b * sb;
// Clamp gz to [0,1] -- matches forward and V1 bilagrid-grad bwd branch.
float gz_raw = kC2G_r * sr + kC2G_g * sg + kC2G_b * sb;
float gz = fminf(fmaxf(gz_raw, 0.0f), 1.0f);
bool gz_in_range = (gz_raw >= 0.0f && gz_raw <= 1.0f);
float x = gx * (W - 1);
float y = gy * (H - 1);
float z = gz * (L - 1);
@@ -250,7 +253,8 @@ __global__ void bilagrid_ppisp_sample_backward_kernel(
v_coords[2*g_off+0] = gx_grad * (float)(x0 != x && x1 != x);
v_coords[2*g_off+1] = gy_grad * (float)(y0 != y && y1 != y);
#endif
gz_grad *= (float)(z0 != z && z1 != z);
// Zero gz_grad outside the [0,1] clamp range -- the clamp's vjp.
if (!gz_in_range) gz_grad = 0.0f;
v_rgb_in[3*g_off+0] = vr + kC2G_r * gz_grad;
v_rgb_in[3*g_off+1] = vg + kC2G_g * gz_grad;;
v_rgb_in[3*g_off+2] = vb + kC2G_b * gz_grad;;
@@ -48,6 +48,9 @@ __global__ void bilagrid_ppisp_sample_forward_kernel(
float gx = coords[2*g_offset+0];
float gy = coords[2*g_offset+1];
float gz = kC2G_r * sr + kC2G_g * sg + kC2G_b * sb;
// Clamp gz to [0,1] -- equivalent to clipping z to [0, L-1]. Matches the
// backward branch (which clamps), so fz here is in [0,1] same as bwd's.
gz = fminf(fmaxf(gz, 0.0f), 1.0f);
float x = gx * (W - 1);
float y = gy * (H - 1);
float z = gz * (L - 1);
@@ -377,12 +377,15 @@ __global__ void bilagrid_ppisp_uniform_sample_backward_v1_kernel_rgb(
float x = (float)wi / (float)(w-1) * (float)(W-1);
float y = (float)hi / (float)(h-1) * (float)(H-1);
#endif
float z = (kC2G_r * sr + kC2G_g * sg + kC2G_b * sb) * (L-1);
// Clamp gz to [0,1] -- matches forward and bilagrid-grad bwd branch.
float gz_raw = kC2G_r * sr + kC2G_g * sg + kC2G_b * sb;
float gz = fminf(fmaxf(gz_raw, 0.0f), 1.0f);
bool gz_in_range = (gz_raw >= 0.0f && gz_raw <= 1.0f);
float z = gz * (L-1);
int x0 = floorf(x), y0 = floorf(y), z0 = floorf(z);
int x1 = min(x0+1, W-1);
int y1 = min(y0+1, H-1);
int z1 = z0 + 1;
z0 = min(max(z0,0), L-1); z1 = min(max(z1,0), L-1);
int z1 = min(z0+1, L-1);
float fx = x-x0, fy = y-y0, fz = z-z0;
@@ -477,6 +480,8 @@ __global__ void bilagrid_ppisp_uniform_sample_backward_v1_kernel_rgb(
}
gz_grad += dwdz[corner] * (L-1) * trilerp;
}
// Zero gz_grad outside the [0,1] clamp range -- the clamp's vjp.
if (!gz_in_range) gz_grad = 0.0f;
grad_rgb.x += kC2G_r * gz_grad;
grad_rgb.y += kC2G_g * gz_grad;
grad_rgb.z += kC2G_b * gz_grad;
@@ -74,6 +74,9 @@ __global__ void bilagrid_ppisp_uniform_sample_forward_kernel(
float gy = (float)hi / (float)(h-1);
#endif
float gz = kC2G_r * sr + kC2G_g * sg + kC2G_b * sb;
// Clamp gz to [0,1] -- equivalent to clipping z to [0, L-1]. Matches the
// backward branch (which clamps), so fz here is in [0,1] same as bwd's.
gz = fminf(fmaxf(gz, 0.0f), 1.0f);
float x = gx * (W - 1);
float y = gy * (H - 1);
float z = gz * (L - 1);
@@ -84,9 +87,7 @@ __global__ void bilagrid_ppisp_uniform_sample_forward_kernel(
int z0 = (int)floorf(z);
int x1 = min(x0+1, W-1);
int y1 = min(y0+1, H-1);
int z1 = z0 + 1;
z0 = min(max(z0, 0), L-1);
z1 = min(max(z1, 0), L-1);
int z1 = min(z0+1, L-1);
// interpolation parameters
float fx = x - (float)x0;
@@ -41,7 +41,10 @@ __global__ void bilagrid_sample_backward_kernel(
float sr = rgb[3*g_off+0], sg = rgb[3*g_off+1], sb = rgb[3*g_off+2];
float gx = coords[2*g_off+0];
float gy = coords[2*g_off+1];
float gz = kC2G_r * sr + kC2G_g * sg + kC2G_b * sb;
// Clamp gz to [0,1] -- matches forward and V1 bilagrid-grad bwd branch.
float gz_raw = kC2G_r * sr + kC2G_g * sg + kC2G_b * sb;
float gz = fminf(fmaxf(gz_raw, 0.0f), 1.0f);
bool gz_in_range = (gz_raw >= 0.0f && gz_raw <= 1.0f);
float x = gx * (W - 1);
float y = gy * (H - 1);
float z = gz * (L - 1);
@@ -170,7 +173,8 @@ __global__ void bilagrid_sample_backward_kernel(
v_coords[2*g_off+0] = gx_grad * (float)(x0 != x && x1 != x);
v_coords[2*g_off+1] = gy_grad * (float)(y0 != y && y1 != y);
#endif
gz_grad *= (float)(z0 != z && z1 != z);
// Zero gz_grad outside the [0,1] clamp range -- the clamp's vjp.
if (!gz_in_range) gz_grad = 0.0f;
v_rgb[3*g_off+0] = vr + kC2G_r * gz_grad;
v_rgb[3*g_off+1] = vg + kC2G_g * gz_grad;;
v_rgb[3*g_off+2] = vb + kC2G_b * gz_grad;;
@@ -2,9 +2,9 @@
template<int C, bool inplace>
__global__ void tv_loss_backward_kernel(
BilagridReader bilagrid, // [N,C,L,H,W]
BilagridReader bilagrid, // [N,L,H,W,C] (channel-last)
const float v_tv_loss, // scalar gradient dL/d(tv_loss)
float* __restrict__ v_bilagrid, // [N,C,L,H,W]
float* __restrict__ v_bilagrid, // [N,L,H,W,C]
int N, int L, int H, int W
) {
int wi = blockIdx.x * blockDim.x + threadIdx.x;
@@ -20,35 +20,43 @@ __global__ void tv_loss_backward_kernel(
float sy = s / (float)(L * (H - 1) * W);
float sz = s / (float)((L - 1) * H * W);
// Channel-last strides.
const int sw = C;
const int sh = W * C;
const int sl = H * W * C;
const int sn = L * H * W * C;
const int cell_base = ni * sn + li * sl + hi * sh + wi * sw;
for (int ci = 0; ci < C; ci++) {
int cell_idx = (((ni * C + ci) * L + li) * H + hi) * W + wi;
int cell_idx = cell_base + ci;
float half_grad = 0.0f;
float val = bilagrid[cell_idx];
if (wi > 0) {
float val0 = bilagrid[cell_idx - 1];
float val0 = bilagrid[cell_idx - sw];
half_grad += (val - val0) * sx;
}
if (wi < W - 1) {
float val0 = bilagrid[cell_idx + 1];
float val0 = bilagrid[cell_idx + sw];
half_grad += (val - val0) * sx;
}
if (hi > 0) {
float val0 = bilagrid[cell_idx - W];
float val0 = bilagrid[cell_idx - sh];
half_grad += (val - val0) * sy;
}
if (hi < H - 1) {
float val0 = bilagrid[cell_idx + W];
float val0 = bilagrid[cell_idx + sh];
half_grad += (val - val0) * sy;
}
if (li > 0) {
float val0 = bilagrid[cell_idx - W*H];
float val0 = bilagrid[cell_idx - sl];
half_grad += (val - val0) * sz;
}
if (li < L - 1) {
float val0 = bilagrid[cell_idx + W*H];
float val0 = bilagrid[cell_idx + sl];
half_grad += (val - val0) * sz;
}
@@ -97,7 +105,7 @@ void tv_loss_backward(
template<int C, bool inplace>
__global__ void channel_mean_backward_kernel(
const float* __restrict__ v_channel_mean, // [C]
float* __restrict__ v_bilagrid, // [N,C,L,H,W]
float* __restrict__ v_bilagrid, // [N,L,H,W,C] (channel-last)
int N, int L, int H, int W
) {
int wi = blockIdx.x * blockDim.x + threadIdx.x;
@@ -108,12 +116,14 @@ __global__ void channel_mean_backward_kernel(
int li = idx % L; idx /= L;
int ni = idx;
const int cell_base = ((((ni * L) + li) * H + hi) * W + wi) * C;
#pragma unroll
for (int ci = 0; ci < C; ci++) {
float grad = v_channel_mean[ci] / (N*L*H*W);
int cell_idx = (((ni * C + ci) * L + li) * H + hi) * W + wi;
int cell_idx = cell_base + ci;
if (inplace)
v_bilagrid[cell_idx] += grad;
else
@@ -9,7 +9,7 @@ namespace cg = cooperative_groups;
template<int C>
__global__ void tv_loss_forward_kernel(
BilagridReader bilagrid, // [N,C,L,H,W]
BilagridReader bilagrid, // [N,L,H,W,C] (channel-last)
float* __restrict__ tv_loss,
int N, int L, int H, int W
) {
@@ -22,29 +22,35 @@ __global__ void tv_loss_forward_kernel(
// int ci = idx % C; idx /= C;
int ni = idx;
// Channel-last strides: ci=1, wi=C, hi=W*C, li=H*W*C, ni=L*H*W*C.
const int sw = C;
const int sh = W * C;
const int sl = H * W * C;
const int sn = L * H * W * C;
float tv_sum = 0.0f;
if (inside) {
const int cell_base = ni * sn + li * sl + hi * sh + wi * sw;
#pragma unroll
for (int ci = 0; ci < C; ci++) {
int base = (ni*C+ci)*L*H*W;
int cell_idx = base + (li*H+hi)*W+wi;
int cell_idx = cell_base + ci;
float val = bilagrid[cell_idx];
if (wi > 0) {
float val0 = bilagrid[cell_idx - 1];
float val0 = bilagrid[cell_idx - sw];
float l2 = (val-val0) * (val-val0);
tv_sum += l2 / (L*H*(W-1));
}
if (hi > 0) {
float val0 = bilagrid[cell_idx - W];
float val0 = bilagrid[cell_idx - sh];
float l2 = (val-val0) * (val-val0);
tv_sum += l2 / (L*(H-1)*W);
}
if (li > 0) {
float val0 = bilagrid[cell_idx - W*H];
float val0 = bilagrid[cell_idx - sl];
float l2 = (val-val0) * (val-val0);
tv_sum += l2 / ((L-1)*H*W);
}
@@ -123,7 +129,7 @@ void tv_loss_forward(
template<int C>
__global__ void channel_mean_forward_kernel(
BilagridReader bilagrid, // [N,C,L,H,W]
BilagridReader bilagrid, // [N,L,H,W,C] (channel-last)
float* __restrict__ channel_mean, // [C]
int N, int L, int H, int W
) {
@@ -140,15 +146,15 @@ __global__ void channel_mean_forward_kernel(
__shared__ float sharedData[blockSize];
int tid = (threadIdx.z * blockDim.y + threadIdx.y) * blockDim.x + threadIdx.x;
const int cell_base = ((((ni * L) + li) * H + hi) * W + wi) * C;
#pragma unroll
for (int ci = 0; ci < C; ci++) {
float val = 0.0f;
if (inside) {
int base = (ni*C+ci)*L*H*W;
int cell_idx = base + (li*H+hi)*W+wi;
val = bilagrid[cell_idx];
val = bilagrid[cell_base + ci];
}
sharedData[tid] = val;
@@ -316,12 +316,15 @@ __global__ void bilagrid_uniform_sample_backward_v1_kernel_rgb(
float x = (float)wi / (float)(w-1) * (float)(W-1);
float y = (float)hi / (float)(h-1) * (float)(H-1);
#endif
float z = (kC2G_r * sr + kC2G_g * sg + kC2G_b * sb) * (L-1);
// Clamp gz to [0,1] -- matches forward and bilagrid-grad bwd branch.
float gz_raw = kC2G_r * sr + kC2G_g * sg + kC2G_b * sb;
float gz = fminf(fmaxf(gz_raw, 0.0f), 1.0f);
bool gz_in_range = (gz_raw >= 0.0f && gz_raw <= 1.0f);
float z = gz * (L-1);
int x0 = floorf(x), y0 = floorf(y), z0 = floorf(z);
int x1 = min(x0+1, W-1);
int y1 = min(y0+1, H-1);
int z1 = z0 + 1;
z0 = min(max(z0,0), L-1); z1 = min(max(z1,0), L-1);
int z1 = min(z0+1, L-1);
// fractional parts
float fx = x - x0, fy = y - y0, fz = z - z0;
@@ -395,6 +398,8 @@ __global__ void bilagrid_uniform_sample_backward_v1_kernel_rgb(
}
gz_grad += dwdz[corner] * (L-1) * trilerp;
}
// Zero gz_grad outside the [0,1] clamp range -- the clamp's vjp.
if (!gz_in_range) gz_grad = 0.0f;
vr += kC2G_r * gz_grad;
vg += kC2G_g * gz_grad;
vb += kC2G_b * gz_grad;
@@ -111,14 +111,17 @@ __global__ void bilagrid_uniform_sample_backward_v2_kernel(
float x = (float)wi / (float)(w-1) * (float)(W-1);
float y = (float)hi / (float)(h-1) * (float)(H-1);
#endif
float z = (kC2G_r * sr + kC2G_g * sg + kC2G_b * sb) * (L-1);
// Clamp gz to [0,1] -- matches forward and V1 bilagrid-grad bwd branch.
float gz_raw = kC2G_r * sr + kC2G_g * sg + kC2G_b * sb;
float gz = fminf(fmaxf(gz_raw, 0.0f), 1.0f);
bool gz_in_range = (gz_raw >= 0.0f && gz_raw <= 1.0f);
float z = gz * (L-1);
// floor + ceil, clamped
int x0 = floorf(x), y0 = floorf(y), z0 = floorf(z);
int x1 = min(x0+1, W-1);
int y1 = min(y0+1, H-1);
int z1 = z0 + 1;
z0 = min(max(z0,0), L-1); z1 = min(max(z1,0), L-1);
int z1 = min(z0+1, L-1);
// fractional parts
float fx = x - x0, fy = y - y0, fz = z - z0;
@@ -215,8 +218,9 @@ __global__ void bilagrid_uniform_sample_backward_v2_kernel(
}
#endif
// save gradient, with discontinuity masking
gz_grad *= (float)(z0 != z && z1 != z);
// Zero gz_grad outside the [0,1] clamp range -- the clamp's vjp.
// (Replaces the prior z-discontinuity heuristic; same intent.)
if (!gz_in_range) gz_grad = 0.0f;
if (inside) {
v_rgb[g_off+0] = vr + kC2G_r * gz_grad;
v_rgb[g_off+1] = vg + kC2G_g * gz_grad;
@@ -71,6 +71,9 @@ __global__ void bilagrid_uniform_sample_forward_kernel(
float gy = (float)hi / (float)(h-1);
#endif
float gz = kC2G_r * sr + kC2G_g * sg + kC2G_b * sb;
// Clamp gz to [0,1] -- equivalent to clipping z to [0, L-1]. Matches the
// backward branch (which clamps), so fz here is in [0,1] same as bwd's.
gz = fminf(fmaxf(gz, 0.0f), 1.0f);
float x = gx * (W - 1);
float y = gy * (H - 1);
float z = gz * (L - 1);
@@ -81,9 +84,7 @@ __global__ void bilagrid_uniform_sample_forward_kernel(
int z0 = (int)floorf(z);
int x1 = min(x0+1, W-1);
int y1 = min(y0+1, H-1);
int z1 = z0 + 1;
z0 = min(max(z0, 0), L-1);
z1 = min(max(z1, 0), L-1);
int z1 = min(z0+1, L-1);
// interpolation parameters
float fx = x - (float)x0;
+62 -35
View File
@@ -16,6 +16,7 @@
#include <numeric>
#include <random>
#include <stdexcept>
#include <set>
#include <unordered_map>
@@ -141,41 +142,12 @@ private:
// All decoders write directly into a pre-allocated batch slot `dst` (pointer
// to the start of row 0 of this image's slot in the [B,H,W,C] buffer).
//
// `expected_h`, `expected_w` are the IndexGroup's promised dimensions; if the
// decoded image disagrees we throw (the IndexGroup invariant was violated).
void decode_rgb_into(const std::string& path,
int expected_h, int expected_w,
PixelDType dtype,
uint8_t* dst)
{
int w, h, ch;
if (dtype == PixelDType::UINT16) {
stbi_us* img = stbi_load_16(path.c_str(), &w, &h, &ch, 3);
if (!img) throw std::runtime_error("DataManager: failed to load 16-bit RGB '" + path + "': " + stbi_failure_reason());
if (w != expected_w || h != expected_h) {
stbi_image_free(img);
throw std::runtime_error("DataManager: rgb shape mismatch for '" + path + "', "
+ std::to_string(w) + "x" + std::to_string(h) + " != expected "
+ std::to_string(expected_w) + "x" + std::to_string(expected_h));
}
std::memcpy(dst, img, (size_t)w * h * 3 * sizeof(stbi_us));
stbi_image_free(img);
} else if (dtype == PixelDType::UINT8) {
stbi_uc* img = stbi_load(path.c_str(), &w, &h, &ch, 3);
if (!img) throw std::runtime_error("DataManager: failed to load 8-bit RGB '" + path + "': " + stbi_failure_reason());
if (w != expected_w || h != expected_h) {
stbi_image_free(img);
throw std::runtime_error("DataManager: rgb shape mismatch for '" + path + "', "
+ std::to_string(w) + "x" + std::to_string(h) + " != expected "
+ std::to_string(expected_w) + "x" + std::to_string(expected_h));
}
std::memcpy(dst, img, (size_t)w * h * 3);
stbi_image_free(img);
} else {
throw std::runtime_error("DataManager: float RGB inputs not supported in stb_image path");
}
}
// `expected_h`, `expected_w` are the IndexGroup's promised dimensions. If the
// decoded image disagrees, the decoder warns (once per IndexGroup, keyed on
// the expected (W,H) pair) and bilinearly resizes the image into the slot,
// matching how mask / depth / normal already handle intra-group shape
// drift. Matches the gsplat / nerfstudio convention of accepting off-by-one
// downscale dims (e.g. Mip-NeRF 360 images_(2|4) round vs. floor).
// ---- CPU resize helpers ---------------------------------------------------
//
@@ -347,6 +319,61 @@ inline void apply_mask_boundary_offset_in_place(uint8_t* mask, int h, int w,
// bilinear for depth / normal). The 1x1 mask case is a degenerate
// nearest-neighbor broadcast and falls out for free.
// Emit a one-shot warning the first time a given (kind, expected_w,
// expected_h) tuple sees an on-disk size mismatch. Keyed on the IndexGroup's
// promised dims, so each group fires at most one warning per modality even
// though the worker pool decodes many files concurrently.
static void _warn_rgb_dim_mismatch_once(
const std::string& path,
int actual_w, int actual_h,
int expected_w, int expected_h)
{
static std::mutex mu;
static std::set<std::pair<int, int>> seen;
std::lock_guard<std::mutex> lk(mu);
auto key = std::make_pair(expected_w, expected_h);
if (seen.insert(key).second) {
std::fprintf(stderr,
"DataManager: rgb shape mismatch for '%s': %dx%d vs camera %dx%d. "
"Resizing on-disk image to match camera dims. "
"Suppressing further warnings for this group.\n",
path.c_str(),
actual_w, actual_h, expected_w, expected_h);
std::fflush(stderr);
}
}
void decode_rgb_into(const std::string& path,
int expected_h, int expected_w,
PixelDType dtype,
uint8_t* dst)
{
int w, h, ch;
if (dtype == PixelDType::UINT16) {
stbi_us* img = stbi_load_16(path.c_str(), &w, &h, &ch, 3);
if (!img) throw std::runtime_error("DataManager: failed to load 16-bit RGB '" + path + "': " + stbi_failure_reason());
if (w == expected_w && h == expected_h) {
std::memcpy(dst, img, (size_t)w * h * 3 * sizeof(stbi_us));
} else {
_warn_rgb_dim_mismatch_once(path, w, h, expected_w, expected_h);
cpu_bilinear_resize<stbi_us, 3>(img, h, w, (stbi_us*)dst, expected_h, expected_w);
}
stbi_image_free(img);
} else if (dtype == PixelDType::UINT8) {
stbi_uc* img = stbi_load(path.c_str(), &w, &h, &ch, 3);
if (!img) throw std::runtime_error("DataManager: failed to load 8-bit RGB '" + path + "': " + stbi_failure_reason());
if (w == expected_w && h == expected_h) {
std::memcpy(dst, img, (size_t)w * h * 3);
} else {
_warn_rgb_dim_mismatch_once(path, w, h, expected_w, expected_h);
cpu_bilinear_resize<stbi_uc, 3>(img, h, w, dst, expected_h, expected_w);
}
stbi_image_free(img);
} else {
throw std::runtime_error("DataManager: float RGB inputs not supported in stb_image path");
}
}
void decode_mask_into(const std::string& path,
int dst_h, int dst_w,
float boundary_offset_frac,
+2 -2
View File
@@ -599,8 +599,8 @@ __global__ void densify_update_weight_kernel(
weight *= sigmoid(opacs[idx]);
if (accum_weight_scalar != nullptr)
weight *= accum_weight_scalar[0];
if (weight == 0.0f)
return;
// if (weight == 0.0f)
// return;
float2 accum = accum_buffer[idx];
if (score_mode == (int)DensifyScoreMode::Max) {
+5 -5
View File
@@ -149,12 +149,12 @@ def bench_360_v2(preset: str, path_to_360_v2: Path, output_prefix: Path,
exit(0)
scenes = [
# ("bicycle", 4),
("bicycle", 4),
("garden", 4),
# ("stump", 4),
# ("bonsai", 2),
# ("counter", 2),
# ("kitchen", 2),
("stump", 4),
("bonsai", 2),
("counter", 2),
("kitchen", 2),
("room", 2),
]