StarAtNyte1's picture
Upload s23dr_2026/model.py with huggingface_hub
4c75b21 verified
Raw
History Blame Contribute Delete
14.2 kB
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