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