import copy import torch import torch.nn as nn from .cdn import build_cdn_attn_mask, CDNModule, cdn_loss # noqa: F401 # --------------------------------------------------------------------------- # Positional Encoding # --------------------------------------------------------------------------- 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) # --------------------------------------------------------------------------- # Decoder # --------------------------------------------------------------------------- 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 # --------------------------------------------------------------------------- # Models # --------------------------------------------------------------------------- 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 # --------------------------------------------------------------------------- # Checkpoint helpers # --------------------------------------------------------------------------- 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 # Remap encoder.layers.X → encoder_layers.X for pre-refactor checkpoints 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