WireFrameDETR / src /multi_view_merger.py
StarAtNyte1's picture
Add src/multi_view_merger.py
e1d80e7 verified
Raw
History Blame Contribute Delete
5.28 kB
"""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