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'])