WireFrameDETR / src /edge_existence.py
StarAtNyte1's picture
Update code
99819e3 verified
Raw
History Blame Contribute Delete
3.84 kB
"""Edge existence classifier for S23DR 2026.
Given merged 3D vertices, predict which vertex pairs form real wireframe edges.
Enables cross-view edge prediction that gestalt-adjacency cannot provide.
"""
import numpy as np
import torch
import torch.nn as nn
class EdgeExistenceModel(nn.Module):
def __init__(self, input_dim=32, hidden_dim=64):
super().__init__()
self.net = nn.Sequential(
nn.Linear(input_dim, hidden_dim),
nn.BatchNorm1d(hidden_dim),
nn.ReLU(),
nn.Dropout(0.2),
nn.Linear(hidden_dim, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, 1),
)
def forward(self, x):
return self.net(x).squeeze(-1)
class EdgeExistenceTrainer:
def __init__(self, device='cuda', lr=1e-3, pos_weight=10.0):
self.device = device
self.model = EdgeExistenceModel().to(device)
self.optimizer = torch.optim.AdamW(self.model.parameters(), lr=lr, weight_decay=1e-4)
pos_weight_tensor = torch.tensor([pos_weight], device=device)
self.criterion = nn.BCEWithLogitsLoss(pos_weight=pos_weight_tensor)
def train_step(self, features, labels):
self.model.train()
self.optimizer.zero_grad()
logits = self.model(features.to(self.device))
loss = self.criterion(logits, labels.float().to(self.device))
loss.backward()
self.optimizer.step()
return loss.item()
@torch.no_grad()
def predict_proba(self, features):
self.model.eval()
if isinstance(features, np.ndarray):
features = torch.from_numpy(features).float()
return torch.sigmoid(self.model(features.to(self.device))).cpu().numpy()
def predict(self, features, threshold=0.4):
return self.predict_proba(features) >= threshold
def save(self, path):
import os
os.makedirs(os.path.dirname(path), exist_ok=True)
torch.save({'model_state_dict': self.model.state_dict()}, path)
def load(self, path):
ckpt = torch.load(path, map_location=self.device, weights_only=True)
self.model.load_state_dict(ckpt['model_state_dict'])
def generate_candidate_pairs(vertices, dist_thresh=8.0):
"""Return all vertex index pairs within dist_thresh metres."""
if len(vertices) < 2:
return []
from scipy.spatial import cKDTree
tree = cKDTree(vertices)
pairs = list(tree.query_pairs(dist_thresh))
return pairs
def label_candidate_pairs(pred_vertices, candidate_pairs, gt_vertices, gt_edges, match_th=0.5):
"""Label each candidate pair 1 (positive) or 0 (negative).
Positive: midpoint within match_th of a GT edge midpoint AND both endpoints
within match_th of a GT vertex.
"""
if len(gt_vertices) == 0 or len(gt_edges) == 0 or len(candidate_pairs) == 0:
return np.zeros(len(candidate_pairs), dtype=np.float32)
from scipy.spatial import cKDTree
gt_v = np.array(gt_vertices)
gt_tree = cKDTree(gt_v)
valid_gt_edges = [(a, b) for a, b in gt_edges if a < len(gt_v) and b < len(gt_v)]
if not valid_gt_edges:
return np.zeros(len(candidate_pairs), dtype=np.float32)
gt_mids = np.array([(gt_v[a] + gt_v[b]) / 2 for a, b in valid_gt_edges])
mid_tree = cKDTree(gt_mids)
labels = np.zeros(len(candidate_pairs), dtype=np.float32)
for i, (a, b) in enumerate(candidate_pairs):
if a >= len(pred_vertices) or b >= len(pred_vertices):
continue
midpoint = (pred_vertices[a] + pred_vertices[b]) / 2
mid_dist, _ = mid_tree.query(midpoint)
if mid_dist > match_th:
continue
da, _ = gt_tree.query(pred_vertices[a])
db, _ = gt_tree.query(pred_vertices[b])
if da < match_th and db < match_th:
labels[i] = 1.0
return labels