add option to offset masks

This commit is contained in:
Harry Chen
2026-06-06 16:05:07 -04:00
parent ede592494f
commit 75e2b395a2
6 changed files with 128 additions and 9 deletions
+2 -2
View File
@@ -81,8 +81,8 @@ endforeach()
set(CMAKE_CXX_STANDARD 17)
set(CMAKE_CXX_STANDARD_REQUIRED ON)
# set(SPLAT_CXX_FLAGS "-O3")
set(SPLAT_CXX_FLAGS "-g")
set(SPLAT_CXX_FLAGS "-O3")
# set(SPLAT_CXX_FLAGS "-g")
if(NOT WIN32)
list(APPEND SPLAT_CXX_FLAGS "-Wno-sign-compare")
endif()
+8
View File
@@ -43,6 +43,14 @@ class SpirulaeSplatDataManagerConfig:
load_normals: bool = True
"""Whether to load normal maps, if exist"""
mask_boundary_offset: float = 0
"""Signed boundary offset applied to binarized masks at decode time, as a
fraction of sqrt(W*H) of the decoded mask. Positive dilates (grows)
foreground, negative erodes (shrinks). Runs on CPU during data loading
via separable Felzenszwalb-Huttenlocher squared-Euclidean DT (exact,
O(N) per row + col). Only honored by the C++ data manager path
(``use_cpp_data_manager=True``)."""
warp_to_pinhole: bool = False
"""Whether to split an image into 5 undistorted pinhole images.
Can sometimes give better quality and compatibility for dataset captured by fisheye/360 cameras."""
+1
View File
@@ -482,6 +482,7 @@ class Trainer:
c_cfg.train_batch_size = train_bs
c_cfg.val_batch_size = val_bs
c_cfg.warp_to_pinhole = bool(dm_cfg.warp_to_pinhole)
c_cfg.mask_boundary_offset = float(dm_cfg.mask_boundary_offset)
# Input intrins / dist_coeffs are needed by the wide warp kernel
# (fisheye / equisolid). Pass them even when only equirectangular
+107 -6
View File
@@ -247,6 +247,101 @@ inline void cpu_nearest_resize_u8(const uint8_t* src, int sh, int sw,
}
// Felzenszwalb-Huttenlocher 1D lower-envelope squared-Euclidean DT.
// f[i] = function value (0 at sources, +INF -- or any large/INF value -- at
// non-sources, or a previous row pass's d^2 in the second sweep).
// d[q] = min over i of f[i] + (q-i)^2.
// v / z are scratch (length n / n+1). INF entries are skipped (no parabola).
inline void dt_1d_squared(const double* f, double* d, int n,
int* v, double* z)
{
const double INF_D = std::numeric_limits<double>::infinity();
int first = 0;
while (first < n && std::isinf(f[first])) ++first;
if (first >= n) {
for (int q = 0; q < n; ++q) d[q] = INF_D;
return;
}
int k = 0;
v[0] = first;
z[0] = -INF_D;
z[1] = +INF_D;
for (int q = first + 1; q < n; ++q) {
if (std::isinf(f[q])) continue;
double s;
for (;;) {
int vk = v[k];
double aq = f[q] + (double)q * (double)q;
double av = f[vk] + (double)vk * (double)vk;
s = (aq - av) / (2.0 * (double)(q - vk));
if (s > z[k]) break;
--k;
}
++k;
v[k] = q;
z[k] = s;
z[k+1] = +INF_D;
}
k = 0;
for (int q = 0; q < n; ++q) {
while (z[k+1] < (double)q) ++k;
double dq = (double)(q - v[k]);
d[q] = dq * dq + f[v[k]];
}
}
// 2D squared-Euclidean DT via separable 1D passes (row, then col).
// `src_value` selects which mask value (0 or 1) acts as the source set.
inline void dt2d_squared(const uint8_t* mask, int h, int w, uint8_t src_value,
double* d2_out,
double* tmp_line, double* tmp_out_line,
int* v_buf, double* z_buf)
{
const double INF_D = std::numeric_limits<double>::infinity();
for (int y = 0; y < h; ++y) {
const uint8_t* mrow = mask + (size_t)y * w;
double* drow = d2_out + (size_t)y * w;
for (int x = 0; x < w; ++x)
tmp_line[x] = (mrow[x] == src_value) ? 0.0 : INF_D;
dt_1d_squared(tmp_line, drow, w, v_buf, z_buf);
}
for (int x = 0; x < w; ++x) {
for (int y = 0; y < h; ++y) tmp_line[y] = d2_out[(size_t)y * w + x];
dt_1d_squared(tmp_line, tmp_out_line, h, v_buf, z_buf);
for (int y = 0; y < h; ++y) d2_out[(size_t)y * w + x] = tmp_out_line[y];
}
}
// Shrink (offset_px < 0) or dilate (offset_px > 0) a binary mask in place
// by abs(offset_px) Euclidean pixels. Output: 1 where signed-distance <=
// offset_px, signed-distance = (dist to foreground) - (dist to background).
inline void apply_mask_boundary_offset_in_place(uint8_t* mask, int h, int w,
float offset_px)
{
if (offset_px == 0.0f || h <= 0 || w <= 0 || (h == 1 && w == 1)) return;
int mx = std::max(h, w);
std::vector<double> d2_fg((size_t)h * w);
std::vector<double> d2_bg((size_t)h * w);
std::vector<double> line_in((size_t)mx);
std::vector<double> line_out((size_t)mx);
std::vector<int> v_buf((size_t)mx);
std::vector<double> z_buf((size_t)mx + 1);
dt2d_squared(mask, h, w, /*src_value=*/1, d2_fg.data(),
line_in.data(), line_out.data(), v_buf.data(), z_buf.data());
dt2d_squared(mask, h, w, /*src_value=*/0, d2_bg.data(),
line_in.data(), line_out.data(), v_buf.data(), z_buf.data());
const double off = (double)offset_px;
for (size_t i = 0; i < (size_t)h * w; ++i) {
double sd = std::sqrt(d2_fg[i]) - std::sqrt(d2_bg[i]);
mask[i] = (sd <= off) ? 1 : 0;
}
}
// Modality decoders. `dst_h` / `dst_w` are the BATCH-slot shape (= group
// shape); when the file is smaller it is upsampled (nearest for mask,
// bilinear for depth / normal). The 1x1 mask case is a degenerate
@@ -254,6 +349,7 @@ inline void cpu_nearest_resize_u8(const uint8_t* src, int sh, int sw,
void decode_mask_into(const std::string& path,
int dst_h, int dst_w,
float boundary_offset_frac,
uint8_t* dst)
{
int w, h, ch;
@@ -271,6 +367,14 @@ void decode_mask_into(const std::string& path,
cpu_nearest_resize_u8(img, h, w, dst, dst_h, dst_w);
}
stbi_image_free(img);
// Apply signed boundary offset (dilate/erode) at the decoded resolution.
// offset_px = fraction * sqrt(dst_W * dst_H).
if (boundary_offset_frac != 0.0f) {
float offset_px = boundary_offset_frac
* std::sqrt((float)dst_w * (float)dst_h);
apply_mask_boundary_offset_in_place(dst, dst_h, dst_w, offset_px);
}
}
void decode_depth_into(const std::string& path,
@@ -910,6 +1014,7 @@ void DataManagerImpl::preload_cpu_cache() {
int32_t mh = _mask_h_per[i], mw = _mask_w_per[i];
_mask_cache[i].assign((size_t)mw * mh, 0);
decode_mask_into(_mask_filenames[i], mh, mw,
_cfg.mask_boundary_offset,
_mask_cache[i].data());
}
if (has_depths() && !_depth_filenames[i].empty()) {
@@ -1170,11 +1275,6 @@ void DataManagerImpl::start_disk_pipeline() {
_q_ready_train = std::make_unique<BoundedQueue<std::shared_ptr<DecodedBatch>>>((size_t)prefetch);
_q_ready_val = std::make_unique<BoundedQueue<std::shared_ptr<DecodedBatch>>>((size_t)prefetch);
// Dedicated RGB workers + (optionally) dedicated mask workers — the
// user asked for "rgb+mask" to share a thread group conceptually, but
// a clean modality split (one pool per queue) is more robust than
// tagging jobs through a shared queue. The mask pool shares the
// workers_rgb knob: spawn ceil(N/2) mask workers up to a floor of 1.
int n_rgb_workers = std::max(1, _cfg.workers_rgb);
for (int t = 0; t < n_rgb_workers; ++t)
_workers.emplace_back(&DataManagerImpl::worker_loop_rgb, this);
@@ -1240,7 +1340,8 @@ void DataManagerImpl::worker_loop_mask() {
int H = b.mask_height, W = b.mask_width;
size_t row = (size_t)H * W;
uint8_t* dst = b.mask_buffer.data() + (size_t)job.slot * row;
decode_mask_into(_mask_filenames[job.ds_index], H, W, dst);
decode_mask_into(_mask_filenames[job.ds_index], H, W,
_cfg.mask_boundary_offset, dst);
if (job.remaining->fetch_sub(1) == 1) {
job.batch->build_views();
job.ready_q->push(job.batch);
+8 -1
View File
@@ -75,7 +75,7 @@ struct DataManagerConfig {
// Worker count per modality pool (DISK mode only). Defaults to a small
// fixed number; tune up for fast disks / many CPU cores.
int workers_rgb = 8;
int workers_rgb = 16;
int workers_depth = 8;
int workers_normal = 8;
@@ -83,6 +83,13 @@ struct DataManagerConfig {
// Hard upper bound on RAM consumed by ready / partially-ready batches.
int prefetch_batches = 4;
// Signed boundary offset applied to binarized masks at decode time,
// expressed as a fraction of sqrt(W*H) of the decoded mask. Positive ->
// dilate (grow) foreground; negative -> erode (shrink) foreground; zero
// disables. Implemented via separable Felzenszwalb-Huttenlocher squared
// Euclidean distance transform on CPU (O(N) per row+col, exact).
float mask_boundary_offset = 0.0f;
// When true: cameras with model FISHEYE or EQUISOLID get split into 5
// cubemap-face pinhole sub-cameras at training time. PINHOLE cameras
// pass through unchanged. EQUIRECTANGULAR cameras are ALWAYS split into
+2
View File
@@ -268,6 +268,8 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
.def_readwrite("workers_depth", &DataManagerConfig::workers_depth)
.def_readwrite("workers_normal", &DataManagerConfig::workers_normal)
.def_readwrite("prefetch_batches", &DataManagerConfig::prefetch_batches)
.def_readwrite("mask_boundary_offset",
&DataManagerConfig::mask_boundary_offset)
.def_readwrite("warp_to_pinhole", &DataManagerConfig::warp_to_pinhole);
py::enum_<CameraModelType>(m, "CameraModelType")