File size: 2,955 Bytes
2ebd374
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
"""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)