File size: 5,278 Bytes
e1d80e7
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
"""Multi-view merger and graph optimizer for S23DR 2026."""

import numpy as np
from collections import defaultdict
from typing import Dict, List, Tuple, Optional
from scipy.spatial import KDTree


def merge_vertices_weighted(vert_edge_per_image, th=0.5, depth_confidence=None):
    all_verts, all_weights, connections_3d, types = [], [], [], []
    cur_start = 0
    for img_idx, data in vert_edge_per_image.items():
        if len(data) < 3: continue
        vertices_2d, connections, vertices_3d = data[:3]
        if len(vertices_3d) == 0: continue
        weight = depth_confidence.get(img_idx, 1.0) if depth_confidence else 1.0
        types.extend([v.get('type', 'apex') if isinstance(v, dict) else 'apex' for v in vertices_2d])
        all_verts.append(vertices_3d); all_weights.extend([weight] * len(vertices_3d))
        for (x, y) in connections: connections_3d.append((x + cur_start, y + cur_start))
        cur_start += len(vertices_3d)
    if not all_verts: return np.zeros((2, 3)), [(0, 1)]
    all_verts = np.concatenate(all_verts, axis=0); all_weights = np.array(all_weights)
    types = np.array([1 if t == 'apex' else 0 for t in types])
    tree = KDTree(all_verts); parent = list(range(len(all_verts)))
    def find(x):
        while parent[x] != x: parent[x] = parent[parent[x]]; x = parent[x]
        return x
    def union(a, b):
        ra, rb = find(a), find(b)
        if ra != rb: parent[ra] = rb
    for i, j in tree.query_pairs(r=th):
        if types[i] == types[j]: union(i, j)
    groups = defaultdict(list)
    for i in range(len(all_verts)): groups[find(i)].append(i)
    new_verts, old_to_new = [], {}
    for group_id, (root, members) in enumerate(groups.items()):
        w = all_weights[members]; pos = all_verts[members]
        new_verts.append(np.average(pos, weights=w, axis=0) if w.sum() > 0 else pos.mean(axis=0))
        for m in members: old_to_new[m] = group_id
    new_verts = np.array(new_verts)
    new_conns = set()
    for a, b in connections_3d:
        na, nb = old_to_new.get(a), old_to_new.get(b)
        if na is not None and nb is not None and na != nb: new_conns.add(tuple(sorted((na, nb))))
    return new_verts, list(new_conns)


def snap_to_planes(vertices, connections, n_planes=5, distance_threshold=0.3, min_inliers=4):
    if len(vertices) < 4: return vertices
    snapped = vertices.copy(); remaining = set(range(len(vertices)))
    for plane_iter in range(n_planes):
        if len(remaining) < min_inliers: break
        remaining_idx = list(remaining); pts = vertices[remaining_idx]
        best_plane, best_inliers = None, []
        rng = np.random.RandomState(42 + plane_iter)
        for _ in range(200):
            if len(pts) < 3: break
            sample_idx = rng.choice(len(pts), 3, replace=False)
            p1, p2, p3 = pts[sample_idx]
            normal = np.cross(p2 - p1, p3 - p1); norm = np.linalg.norm(normal)
            if norm < 1e-8: continue
            normal /= norm; d = -np.dot(normal, p1)
            inliers = np.where(np.abs(pts @ normal + d) < distance_threshold)[0]
            if len(inliers) > len(best_inliers): best_inliers = inliers; best_plane = (normal, d)
        if best_plane is None or len(best_inliers) < min_inliers: break
        normal, d = best_plane
        for local_idx in best_inliers:
            global_idx = remaining_idx[local_idx]; pt = vertices[global_idx]
            snapped[global_idx] = pt - (np.dot(normal, pt) + d) * normal; remaining.discard(global_idx)
    return snapped


def regularize_edges(vertices, connections, horizontal_threshold=0.15, vertical_threshold=0.15):
    reg_verts = vertices.copy()
    for a, b in connections:
        if a >= len(vertices) or b >= len(vertices): continue
        direction = reg_verts[b] - reg_verts[a]; length = np.linalg.norm(direction)
        if length < 1e-6: continue
        if abs(direction[2] / length) < np.sin(horizontal_threshold):
            avg_z = (reg_verts[a][2] + reg_verts[b][2]) / 2; reg_verts[a][2] = avg_z; reg_verts[b][2] = avg_z
    return reg_verts


def remove_redundant_edges(vertices, connections, min_edge_length=0.1, max_degree=8):
    edges = list(set(tuple(sorted(c)) for c in connections))
    filtered = [(a, b, np.linalg.norm(vertices[a] - vertices[b])) for a, b in edges if a < len(vertices) and b < len(vertices) and np.linalg.norm(vertices[a] - vertices[b]) >= min_edge_length]
    if not filtered: return [(0, 1)] if len(vertices) >= 2 else []
    filtered.sort(key=lambda x: x[2]); degree = defaultdict(int); final = []
    for a, b, l in filtered:
        if degree[a] < max_degree and degree[b] < max_degree: final.append((a, b)); degree[a] += 1; degree[b] += 1
    return final


def full_graph_optimization(vertices, connections, colmap_rec=None, do_plane_snap=True, do_regularize=True, do_prune=True):
    if len(vertices) < 2 or len(connections) < 1: return vertices, connections
    if do_prune: connections = remove_redundant_edges(vertices, connections)
    if do_plane_snap and len(vertices) >= 4: vertices = snap_to_planes(vertices, connections)
    if do_regularize: vertices = regularize_edges(vertices, connections)
    return vertices, connections