File size: 6,955 Bytes
6befb78 99819e3 6befb78 99819e3 | 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 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 | """Learned vertex detector for S23DR 2026.
Train a CNN to predict 2D vertex heatmaps from gestalt + depth images.
Uses ground-truth 3D wireframe vertices projected to 2D as supervision.
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
import torchvision.models as models
import numpy as np
from typing import Tuple, List, Optional
class VertexHeatmapNet(nn.Module):
def __init__(self, in_channels=7, num_classes=2, pretrained_backbone=True):
super().__init__()
backbone = models.resnet18(weights='IMAGENET1K_V1' if pretrained_backbone else None)
self.conv1 = nn.Conv2d(in_channels, 64, kernel_size=7, stride=2, padding=3, bias=False)
if pretrained_backbone:
with torch.no_grad():
pretrained_weight = backbone.conv1.weight
new_weight = torch.zeros(64, in_channels, 7, 7)
new_weight[:, :3, :, :] = pretrained_weight
avg_weight = pretrained_weight.mean(dim=1, keepdim=True)
for i in range(3, in_channels):
new_weight[:, i:i+1, :, :] = avg_weight
self.conv1.weight = nn.Parameter(new_weight)
self.bn1 = backbone.bn1
self.relu = backbone.relu
self.maxpool = backbone.maxpool
self.layer1 = backbone.layer1
self.layer2 = backbone.layer2
self.layer3 = backbone.layer3
self.layer4 = backbone.layer4
self.up4 = nn.Sequential(nn.ConvTranspose2d(512, 256, 4, stride=2, padding=1), nn.BatchNorm2d(256), nn.ReLU(inplace=True))
self.up3 = nn.Sequential(nn.ConvTranspose2d(512, 128, 4, stride=2, padding=1), nn.BatchNorm2d(128), nn.ReLU(inplace=True))
self.up2 = nn.Sequential(nn.ConvTranspose2d(256, 64, 4, stride=2, padding=1), nn.BatchNorm2d(64), nn.ReLU(inplace=True))
self.up1 = nn.Sequential(nn.ConvTranspose2d(128, 64, 4, stride=2, padding=1), nn.BatchNorm2d(64), nn.ReLU(inplace=True))
self.head = nn.Sequential(nn.Conv2d(64, 32, 3, padding=1), nn.ReLU(inplace=True), nn.Conv2d(32, num_classes, 1))
def forward(self, x):
x = self.conv1(x); x = self.bn1(x); x0 = self.relu(x); x = self.maxpool(x0)
x1 = self.layer1(x); x2 = self.layer2(x1); x3 = self.layer3(x2); x4 = self.layer4(x3)
d4 = self.up4(x4); d3 = self.up3(torch.cat([d4, x3], dim=1))
d2 = self.up2(torch.cat([d3, x2], dim=1)); d1 = self.up1(torch.cat([d2, x1], dim=1))
return self.head(d1)
def create_vertex_heatmap(vertices_2d, vertex_types, height, width, sigma=3.0):
heatmap = np.zeros((2, height, width), dtype=np.float32)
type_to_channel = {'apex': 0, 'eave_end_point': 1}
for (u, v), vtype in zip(vertices_2d, vertex_types):
ch = type_to_channel.get(vtype, 0)
u_int, v_int = int(round(u)), int(round(v))
if u_int < 0 or u_int >= width or v_int < 0 or v_int >= height:
continue
radius = int(3 * sigma)
for dy in range(-radius, radius + 1):
for dx in range(-radius, radius + 1):
py, px = v_int + dy, u_int + dx
if 0 <= py < height and 0 <= px < width:
val = np.exp(-(dx*dx + dy*dy) / (2 * sigma * sigma))
heatmap[ch, py, px] = max(heatmap[ch, py, px], val)
return heatmap
def prepare_input_tensor(gestalt_img, depth_img, ade_img, target_size=(192, 256)):
H, W = target_size
gest = np.array(gestalt_img.resize((W, H))).astype(np.float32) / 255.0
if gest.ndim == 2: gest = np.stack([gest]*3, axis=-1)
depth = np.array(depth_img.resize((W, H))).astype(np.float32) / 1000.0
depth = np.clip(depth / 50.0, 0, 1)
if depth.ndim == 2: depth = depth[:, :, np.newaxis]
ade = np.array(ade_img.resize((W, H))).astype(np.float32) / 255.0
if ade.ndim == 2: ade = np.stack([ade]*3, axis=-1)
combined = np.concatenate([gest, depth, ade], axis=-1)
return torch.from_numpy(combined).permute(2, 0, 1)
def extract_vertices_from_heatmap(heatmap, threshold=0.3, nms_radius=5):
from scipy.ndimage import maximum_filter
vertices, types = [], []
type_names = ['apex', 'eave_end_point']
for ch in range(heatmap.shape[0]):
hm = heatmap[ch]
local_max = maximum_filter(hm, size=2*nms_radius+1)
peaks = (hm == local_max) & (hm >= threshold)
ys, xs = np.where(peaks)
for y, x in zip(ys, xs):
vertices.append([x, y]); types.append(type_names[ch])
if not vertices: return np.zeros((0, 2)), []
return np.array(vertices), types
class VertexDetectorTrainer:
def __init__(self, model, lr=1e-4, device='cuda'):
self.model = model.to(device); self.device = device
self.optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4)
self.scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(self.optimizer, T_max=100, eta_min=1e-6)
def focal_loss(self, pred, target, alpha=2.0, beta=4.0):
pred = torch.clamp(torch.sigmoid(pred), 1e-6, 1 - 1e-6)
pos_mask = (target >= 0.99); neg_mask = ~pos_mask
pos_loss = -((1 - pred) ** alpha) * torch.log(pred) * pos_mask.float()
neg_loss = -((1 - target) ** beta) * (pred ** alpha) * torch.log(1 - pred) * neg_mask.float()
return (pos_loss.sum() + neg_loss.sum()) / pos_mask.float().sum().clamp(min=1)
def train_step(self, input_tensor, target_heatmap):
self.model.train(); self.optimizer.zero_grad()
input_tensor = input_tensor.to(self.device); target_heatmap = target_heatmap.to(self.device)
pred = self.model(input_tensor)
if pred.shape[-2:] != target_heatmap.shape[-2:]:
target_heatmap = F.interpolate(target_heatmap, size=pred.shape[-2:], mode='bilinear', align_corners=False)
loss = self.focal_loss(pred, target_heatmap)
loss.backward(); torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1.0); self.optimizer.step()
return loss.item()
@torch.no_grad()
def predict(self, input_tensor):
self.model.eval()
if input_tensor.ndim == 3: input_tensor = input_tensor.unsqueeze(0)
return torch.sigmoid(self.model(input_tensor.to(self.device)))[0].cpu().numpy()
def save(self, path):
torch.save({
'model_state_dict': self.model.state_dict(),
'optimizer_state_dict': self.optimizer.state_dict(),
'scheduler_state_dict': self.scheduler.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'])
self.optimizer.load_state_dict(ckpt['optimizer_state_dict'])
if 'scheduler_state_dict' in ckpt:
self.scheduler.load_state_dict(ckpt['scheduler_state_dict'])
|