WireFrameDETR / src /depth_refiner.py
StarAtNyte1's picture
Add src/depth_refiner.py
2ebd374 verified
Raw
History Blame Contribute Delete
2.96 kB
"""Depth refinement module for S23DR 2026.
Learn a per-scene scale + shift correction for MoGe v2 depth maps
using COLMAP sparse depth as supervision.
"""
import numpy as np
import torch
import torch.nn as nn
from typing import Tuple, Optional
def ransac_scale_fit(dense_depth, sparse_depth, house_mask=None, n_iterations=100, inlier_threshold=0.5, min_samples=3):
mask = (sparse_depth > 0) & (dense_depth > 0) & (sparse_depth < 50) & (dense_depth < 50)
if house_mask is not None: mask = mask & house_mask
X, Y = dense_depth[mask], sparse_depth[mask]
if len(X) < min_samples:
return (np.median(Y / X), 0.0) if len(X) > 0 else (1.0, 0.0)
best_scale, best_shift, best_inliers = 1.0, 0.0, 0
rng = np.random.RandomState(42)
for _ in range(n_iterations):
idx = rng.choice(len(X), min(min_samples, len(X)), replace=False)
A = np.column_stack([X[idx], np.ones(len(idx))])
try: params, _, _, _ = np.linalg.lstsq(A, Y[idx], rcond=None); scale, shift = params[0], params[1]
except: continue
if scale <= 0: continue
inliers = (np.abs(scale * X + shift - Y) < inlier_threshold).sum()
if inliers > best_inliers: best_inliers = inliers; best_scale = scale; best_shift = shift
inlier_mask = np.abs(best_scale * X + best_shift - Y) < inlier_threshold
if inlier_mask.sum() >= min_samples:
A = np.column_stack([X[inlier_mask], np.ones(inlier_mask.sum())])
try: params, _, _, _ = np.linalg.lstsq(A, Y[inlier_mask], rcond=None); best_scale, best_shift = params[0], params[1]
except: pass
return best_scale, best_shift
def fit_depth_to_colmap(dense_depth, colmap_rec, img_id, ade_seg=None, method='ransac'):
from hoho2025.example_solutions import get_sparse_depth, get_house_mask
sparse_depth, found, col_img, proj_pts = get_sparse_depth(colmap_rec, img_id, dense_depth)
if not found: return dense_depth, False
house_mask = get_house_mask(ade_seg) if ade_seg is not None else None
if method == 'ransac':
scale, shift = ransac_scale_fit(dense_depth, sparse_depth, house_mask)
fitted = scale * dense_depth + shift
else:
from hoho2025.example_solutions import fit_scale_robust_median
_, fitted = fit_scale_robust_median(dense_depth, sparse_depth, house_mask)
return np.clip(fitted, 0, None), True
class DepthRefinerNet(nn.Module):
def __init__(self, in_channels=5):
super().__init__()
self.net = nn.Sequential(nn.Conv2d(in_channels, 32, 3, padding=1), nn.ReLU(), nn.Conv2d(32, 32, 3, padding=1), nn.ReLU(), nn.Conv2d(32, 16, 3, padding=1), nn.ReLU(), nn.Conv2d(16, 2, 1))
def forward(self, x, dense_depth):
correction = self.net(x)
scale = torch.sigmoid(correction[:, 0:1]) * 4.0
shift = correction[:, 1:2] * 2.0
return torch.clamp(scale * dense_depth + shift, min=0)