| import copy |
|
|
| import torch |
| import torch.nn as nn |
|
|
| from .cdn import build_cdn_attn_mask, CDNModule, cdn_loss |
|
|
|
|
| |
| |
| |
|
|
| class SinCos3DPE(nn.Module): |
| """Sinusoidal 3D positional encoding split evenly across X/Y/Z.""" |
|
|
| def __init__(self, embed_dim, max_temp=10000): |
| super().__init__() |
| self.embed_dim = embed_dim |
| self.max_temp = max_temp |
| self.channels = embed_dim // 3 |
|
|
| def forward(self, xyz): |
| device = xyz.device |
| dim_t = torch.arange(self.channels, dtype=torch.float32, device=device) |
| dim_t = self.max_temp ** (2 * (dim_t // 2) / self.channels) |
| pos = xyz.unsqueeze(-1) * 2 * torch.pi |
| pos_scaled = pos / dim_t |
| pos_sin = pos_scaled[:, :, :, 0::2].sin() |
| pos_cos = pos_scaled[:, :, :, 1::2].cos() |
| return torch.cat([pos_sin, pos_cos], dim=-1).flatten(2) |
|
|
|
|
| |
| |
| |
|
|
| class WireframeDecoderLayer(nn.Module): |
| def __init__(self, embed_dim, nhead=8, dim_feedforward=2048, dropout=0.1): |
| super().__init__() |
| self.self_attn = nn.MultiheadAttention(embed_dim, nhead, dropout=dropout, batch_first=True) |
| self.cross_attn = nn.MultiheadAttention(embed_dim, nhead, dropout=dropout, batch_first=True) |
| self.linear1 = nn.Linear(embed_dim, dim_feedforward) |
| self.dropout = nn.Dropout(dropout) |
| self.linear2 = nn.Linear(dim_feedforward, embed_dim) |
| self.norm1 = nn.LayerNorm(embed_dim) |
| self.norm2 = nn.LayerNorm(embed_dim) |
| self.norm3 = nn.LayerNorm(embed_dim) |
| self.dropout1 = nn.Dropout(dropout) |
| self.dropout2 = nn.Dropout(dropout) |
| self.dropout3 = nn.Dropout(dropout) |
| self.activation = nn.ReLU() |
|
|
| @staticmethod |
| def with_pos_embed(tensor, pos): |
| return tensor if pos is None else tensor + pos |
|
|
| def forward(self, tgt, memory, tgt_pos=None, memory_pos=None, attn_mask=None): |
| q = k = self.with_pos_embed(tgt, tgt_pos) |
| tgt2 = self.self_attn(q, k, value=tgt, attn_mask=attn_mask, need_weights=False)[0] |
| tgt = self.norm1(tgt + self.dropout1(tgt2)) |
| tgt2 = self.cross_attn( |
| query=self.with_pos_embed(tgt, tgt_pos), |
| key=self.with_pos_embed(memory, memory_pos), |
| value=memory, |
| need_weights=False, |
| )[0] |
| tgt = self.norm2(tgt + self.dropout2(tgt2)) |
| tgt2 = self.linear2(self.dropout(self.activation(self.linear1(tgt)))) |
| tgt = self.norm3(tgt + self.dropout3(tgt2)) |
| return tgt |
|
|
|
|
| class WireframeDecoder(nn.Module): |
| def __init__(self, decoder_layer, num_layers): |
| super().__init__() |
| self.layers = nn.ModuleList([copy.deepcopy(decoder_layer) for _ in range(num_layers)]) |
| self.norm = nn.LayerNorm(decoder_layer.linear2.out_features) |
|
|
| def forward(self, tgt, memory, tgt_pos=None, memory_pos=None, attn_mask=None): |
| output = tgt |
| intermediate = [] |
| for layer in self.layers: |
| output = layer(output, memory, tgt_pos=tgt_pos, memory_pos=memory_pos, attn_mask=attn_mask) |
| intermediate.append(self.norm(output)) |
| return intermediate[-1], intermediate |
|
|
|
|
| |
| |
| |
|
|
| class WireframeBase(nn.Module): |
| """Transformer encoder-decoder for 3D point cloud wireframe detection.""" |
|
|
| def __init__( |
| self, |
| input_dim, |
| embed_dim=256, |
| num_queries=100, |
| num_classes=1, |
| num_encoder_layers=4, |
| num_decoder_layers=6, |
| vertex_feat_dim=32, |
| enable_vertex_feat_head=False, |
| ): |
| super().__init__() |
| self.embed_dim = embed_dim |
| self.vertex_feat_dim = vertex_feat_dim |
| self.enable_vertex_feat_head = enable_vertex_feat_head |
|
|
| self.input_proj = nn.Linear(input_dim, embed_dim) |
| self.pos_encoder = SinCos3DPE(embed_dim) |
| self.encoder = nn.TransformerEncoder( |
| nn.TransformerEncoderLayer(d_model=embed_dim, nhead=8, batch_first=True), |
| num_layers=num_encoder_layers, |
| ) |
| self.query_anchor_points = nn.Parameter(torch.rand(num_queries, 6)) |
| self.query_content = nn.Embedding(num_queries, embed_dim) |
|
|
| decoder_layer = WireframeDecoderLayer(embed_dim=embed_dim, nhead=8) |
| self.decoder = WireframeDecoder(decoder_layer, num_layers=num_decoder_layers) |
|
|
| self.class_head = nn.Linear(embed_dim, num_classes + 1) |
| self.coord_head = nn.Sequential(nn.Linear(embed_dim, embed_dim), nn.ReLU(), nn.Linear(embed_dim, 6)) |
| if enable_vertex_feat_head: |
| self.vertex_feat_head = self._build_vertex_feat_head() |
|
|
| def _build_vertex_feat_head(self): |
| return nn.Sequential( |
| nn.Linear(self.embed_dim, self.embed_dim), |
| nn.ReLU(), |
| nn.Linear(self.embed_dim, 2 * self.vertex_feat_dim), |
| ) |
|
|
| def enable_vertex_features(self) -> None: |
| if not hasattr(self, "vertex_feat_head"): |
| self.vertex_feat_head = self._build_vertex_feat_head() |
| ref_param = self.input_proj.weight |
| self.vertex_feat_head.to(device=ref_param.device, dtype=ref_param.dtype) |
| self.enable_vertex_feat_head = True |
|
|
| def _encode(self, x: torch.Tensor) -> torch.Tensor: |
| return self.encoder(x) |
|
|
| def _decode_outputs(self, hs, intermediate_hs, B): |
| outputs = { |
| "pred_logits": self.class_head(hs), |
| "pred_coords": self.coord_head(hs).sigmoid(), |
| "aux_outputs": [ |
| {"pred_logits": self.class_head(h), "pred_coords": self.coord_head(h).sigmoid()} |
| for h in intermediate_hs[:-1] |
| ], |
| } |
| if self.enable_vertex_feat_head: |
| self.enable_vertex_features() |
| outputs["pred_vertex_feats"] = nn.functional.normalize( |
| self.vertex_feat_head(hs).view(B, -1, 2, self.vertex_feat_dim), dim=-1) |
| for aux, h in zip(outputs["aux_outputs"], intermediate_hs[:-1]): |
| aux["pred_vertex_feats"] = nn.functional.normalize( |
| self.vertex_feat_head(h).view(B, -1, 2, self.vertex_feat_dim), dim=-1) |
| return outputs |
|
|
| def forward(self, xyz, features, **kwargs): |
| B, N, _ = xyz.shape |
| src = self.input_proj(features) |
| pos_src = self.pos_encoder(xyz) |
| memory = self._encode(src + pos_src) |
|
|
| query_anchor_points = self.query_anchor_points.view(1, -1, 2, 3).repeat(B, 1, 1, 1) |
| query_pos = (self.pos_encoder(query_anchor_points[:, :, 0, :]) |
| + self.pos_encoder(query_anchor_points[:, :, 1, :])) |
| tgt = self.query_content.weight.unsqueeze(0).repeat(B, 1, 1) |
| hs, intermediate_hs = self.decoder(tgt, memory, tgt_pos=query_pos, memory_pos=pos_src) |
| return self._decode_outputs(hs, intermediate_hs, B) |
|
|
| @torch.no_grad() |
| def inference(self, xyz, features, threshold=0.5): |
| if xyz.ndim != 2: |
| raise NotImplementedError("Only implemented for single example.") |
| outputs = self(xyz.unsqueeze(0), features.unsqueeze(0)) |
| probs = outputs["pred_logits"].softmax(-1)[0] |
| scores, labels = probs[:, :-1].max(-1) |
| keep = scores > threshold |
| pred_coords = outputs["pred_coords"][0][keep].view(-1, 2, 3) |
| pred_vertices = pred_coords.reshape(-1, 3) |
| pred_edges = torch.arange(pred_vertices.shape[0], device=pred_vertices.device, dtype=torch.long).view(-1, 2) |
| results = { |
| "vertices": pred_vertices.cpu().numpy(), |
| "edges": pred_edges.cpu().numpy(), |
| "labels": labels[keep].cpu().numpy(), |
| "scores": scores[keep].cpu().numpy(), |
| } |
| if "pred_vertex_feats" in outputs: |
| results["vertex_feats"] = ( |
| outputs["pred_vertex_feats"][0][keep].view(-1, self.vertex_feat_dim).cpu().numpy() |
| ) |
| return results |
|
|
|
|
| class WireframeDETR(WireframeBase): |
| """WireframeBase with multi-scale encoder and CDN training support. |
| |
| Multi-scale: learned weighted aggregation of last K encoder layer outputs. |
| CDN: contrastive denoising queries injected at training time. |
| """ |
|
|
| def __init__( |
| self, |
| input_dim, |
| embed_dim=256, |
| num_queries=100, |
| num_classes=1, |
| num_encoder_layers=4, |
| num_decoder_layers=6, |
| **kwargs, |
| ): |
| super().__init__(input_dim, embed_dim, num_queries, num_classes, |
| num_encoder_layers, num_decoder_layers, **kwargs) |
| del self.encoder |
| self.encoder_layers = nn.ModuleList([ |
| nn.TransformerEncoderLayer(d_model=embed_dim, nhead=8, batch_first=True) |
| for _ in range(num_encoder_layers) |
| ]) |
| self.num_scale_layers = min(3, num_encoder_layers) |
| self.scale_weights = nn.Parameter(torch.zeros(self.num_scale_layers)) |
|
|
| def _encode(self, x: torch.Tensor) -> torch.Tensor: |
| enc_outputs = [] |
| for layer in self.encoder_layers: |
| x = layer(x) |
| enc_outputs.append(x) |
| scale_w = self.scale_weights.softmax(0) |
| return sum(scale_w[i] * enc_outputs[-(self.num_scale_layers - i)] |
| for i in range(self.num_scale_layers)) |
|
|
| def forward(self, xyz, features, **kwargs): |
| cdn_data = kwargs.get('cdn_data', None) |
| B, N, _ = xyz.shape |
|
|
| src = self.input_proj(features) |
| pos_src = self.pos_encoder(xyz) |
| memory = self._encode(src + pos_src) |
|
|
| query_anchor_points = self.query_anchor_points.view(1, -1, 2, 3).repeat(B, 1, 1, 1) |
| query_pos = (self.pos_encoder(query_anchor_points[:, :, 0, :]) |
| + self.pos_encoder(query_anchor_points[:, :, 1, :])) |
| tgt = self.query_content.weight.unsqueeze(0).repeat(B, 1, 1) |
|
|
| attn_mask = None |
| num_dn = 0 |
| if cdn_data is not None: |
| num_dn = cdn_data['num_dn'] |
| dn_pts = cdn_data['dn_coords'].view(B, num_dn, 2, 3) |
| dn_pos = (self.pos_encoder(dn_pts[:, :, 0, :]) + self.pos_encoder(dn_pts[:, :, 1, :])) |
| tgt = torch.cat([tgt, cdn_data['dn_content']], dim=1) |
| query_pos = torch.cat([query_pos, dn_pos], dim=1) |
| attn_mask = build_cdn_attn_mask(self.query_content.num_embeddings, num_dn, tgt.device) |
|
|
| hs, intermediate_hs = self.decoder(tgt, memory, tgt_pos=query_pos, |
| memory_pos=pos_src, attn_mask=attn_mask) |
|
|
| num_q = self.query_content.num_embeddings |
| hs_learn = hs[:, :num_q] |
| intermediate_learn = [h[:, :num_q] for h in intermediate_hs] |
| outputs = self._decode_outputs(hs_learn, intermediate_learn, B) |
|
|
| if num_dn > 0: |
| hs_dn = hs[:, num_q:] |
| outputs["dn_outputs"] = { |
| "pred_logits": self.class_head(hs_dn), |
| "pred_coords": self.coord_head(hs_dn).sigmoid(), |
| } |
| return outputs |
|
|
|
|
| |
| |
| |
|
|
| def _checkpoint_has_vertex_feat_head(checkpoint): |
| return any(k.startswith("vertex_feat_head.") for k in checkpoint.keys()) |
|
|
|
|
| def _checkpoint_vertex_feat_dim(checkpoint, default=32): |
| weight = checkpoint.get("vertex_feat_head.2.weight") |
| if weight is None: |
| return default |
| return weight.shape[0] // 2 |
|
|
|
|
| def get_model(checkpoint, num_classes=1, enable_vertex_feat_head=None): |
| embed_dim, feature_dim = checkpoint["input_proj.weight"].shape |
| num_queries = checkpoint["query_anchor_points"].shape[0] |
|
|
| encoder_layers = set() |
| decoder_layers = set() |
| for k in checkpoint.keys(): |
| if k.startswith("encoder_layers."): |
| encoder_layers.add(int(k.split(".")[1])) |
| elif k.startswith("encoder.layers."): |
| encoder_layers.add(int(k.split(".")[2])) |
| elif k.startswith("decoder.layers"): |
| decoder_layers.add(int(k.split(".")[2])) |
|
|
| if enable_vertex_feat_head is None: |
| enable_vertex_feat_head = _checkpoint_has_vertex_feat_head(checkpoint) |
|
|
| return WireframeDETR( |
| input_dim=feature_dim, |
| embed_dim=embed_dim, |
| num_queries=num_queries, |
| num_classes=num_classes, |
| num_encoder_layers=len(encoder_layers), |
| num_decoder_layers=len(decoder_layers), |
| vertex_feat_dim=_checkpoint_vertex_feat_dim(checkpoint), |
| enable_vertex_feat_head=enable_vertex_feat_head, |
| ) |
|
|
|
|
| def load_checkpoint_compat(model: WireframeBase, checkpoint: dict, |
| enable_vertex_feat_head=None, verbose: bool = True): |
| if enable_vertex_feat_head is None: |
| enable_vertex_feat_head = _checkpoint_has_vertex_feat_head(checkpoint) |
| if enable_vertex_feat_head: |
| model.enable_vertex_features() |
| else: |
| model.enable_vertex_feat_head = False |
|
|
| |
| if any(k.startswith("encoder.layers.") for k in checkpoint): |
| checkpoint = { |
| ("encoder_layers." + k[len("encoder.layers."):] if k.startswith("encoder.layers.") else k): v |
| for k, v in checkpoint.items() |
| } |
|
|
| missing_keys, unexpected_keys = model.load_state_dict(checkpoint, strict=False) |
| if verbose: |
| non_vertex_missing = [k for k in missing_keys if not k.startswith("vertex_feat_head.")] |
| if non_vertex_missing: |
| print(f"Missing checkpoint keys: {non_vertex_missing}") |
| if unexpected_keys: |
| print(f"Unexpected checkpoint keys: {unexpected_keys}") |
| if missing_keys and not non_vertex_missing: |
| print("Checkpoint has no vertex feature head; leaving the new head randomly initialized.") |
| return missing_keys, unexpected_keys |
|
|