mirror of
https://github.com/harry7557558/spirula-studio.git
synced 2026-10-02 02:44:54 +08:00
investigate and fix splat quality regression
This commit is contained in:
@@ -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;
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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),
|
||||
]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user