"""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