Spaces:
Running on Zero
Running on Zero
| """Faithful, dependency-light port of the official PatchAlign3D stage-2 inference path. | |
| Source of truth: | |
| https://github.com/souhail-hadgi/PatchAlign3D | |
| src/models/point_transformer.py (encoder + patch grouping) | |
| src/inference/infer.py (single-shape inference) | |
| src/inference/eval.py (ShapeNetPart / FAUST evaluation) | |
| src/datasets/shapenet.py (pc_normalize, 2048-point sampling) | |
| Deviations from upstream, all behaviour-preserving: | |
| * `pointnet2_ops.furthest_point_sample` -> pure-torch FPS with the same | |
| deterministic seeding (start from index 0, squared distances, argmax). | |
| * `knn_cuda.KNN(..., transpose_mode=True)` -> pure-torch cdist + topk | |
| (ascending distance order, identical semantics). | |
| * open_clip `ViT-bigG-14 / laion2b_s39b_b160k` text tower -> the *same* | |
| weights served as a HF `CLIPTextModelWithProjection` (verified numerically | |
| identical up to fp16 storage rounding), so only the ~1.4 GB text tower is | |
| downloaded instead of the full 10 GB two-tower checkpoint. | |
| """ | |
| from __future__ import annotations | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| # -------------------------------------------------------------------------------------- | |
| # Config constants taken verbatim from the reference scripts | |
| # -------------------------------------------------------------------------------------- | |
| TRANS_DIM = 384 | |
| DEPTH = 12 | |
| NUM_HEADS = 6 | |
| ENCODER_DIMS = 256 | |
| DROP_PATH_RATE = 0.1 | |
| DEFAULT_NUM_GROUP = 128 | |
| DEFAULT_GROUP_SIZE = 32 | |
| DEFAULT_NPOINTS = 2048 | |
| DEFAULT_TAU = 0.07 | |
| CLIP_TEXT_REPO = "stabilityai/stable-diffusion-xl-base-1.0" | |
| CLIP_TEXT_SUBFOLDER = "text_encoder_2" | |
| CLIP_TOKENIZER_SUBFOLDER = "tokenizer_2" | |
| CLIP_TEXT_DIM = 1280 # ViT-bigG-14 joint embedding dim | |
| PART_ONLY_TEMPLATES = ["{}", "a {}", "{} part"] | |
| PART_PLUS_CAT_TEMPLATES = [ | |
| "a {} of a {}", | |
| "the {} of a {}", | |
| "{} of {}", | |
| "a {} part of a {}", | |
| ] | |
| def clean_text(s: str) -> str: | |
| """Upstream `_clean_text`: lowercase, underscores -> spaces, strip punctuation.""" | |
| s = s.strip().lower().replace("_", " ") | |
| out = [] | |
| for ch in s: | |
| out.append(ch if (ch.isalnum() or ch.isspace()) else " ") | |
| return " ".join("".join(out).split()) | |
| # -------------------------------------------------------------------------------------- | |
| # Pure-torch replacements for pointnet2_ops / knn_cuda | |
| # -------------------------------------------------------------------------------------- | |
| def furthest_point_sample(xyz: torch.Tensor, npoint: int) -> torch.Tensor: | |
| """Iterative FPS matching `pointnet2_ops.furthest_point_sample`. | |
| Starts from point index 0 and greedily picks the point with the largest | |
| squared distance to the already-selected set (exactly what the CUDA kernel | |
| does). Returns (B, npoint) long indices. | |
| """ | |
| B, N, _ = xyz.shape | |
| device = xyz.device | |
| idx = torch.zeros(B, npoint, dtype=torch.long, device=device) | |
| dist = torch.full((B, N), 1e10, device=device, dtype=xyz.dtype) | |
| farthest = torch.zeros(B, dtype=torch.long, device=device) | |
| ar = torch.arange(B, device=device) | |
| for i in range(npoint): | |
| idx[:, i] = farthest | |
| centroid = xyz[ar, farthest, :].view(B, 1, 3) | |
| d = ((xyz - centroid) ** 2).sum(-1) | |
| dist = torch.minimum(dist, d) | |
| farthest = dist.argmax(-1) | |
| return idx | |
| def fps(data: torch.Tensor, number: int) -> torch.Tensor: | |
| """(B, N, 3) -> (B, number, 3) furthest-point-sampled coordinates.""" | |
| idx = furthest_point_sample(data, number) | |
| return torch.gather(data, 1, idx.unsqueeze(-1).expand(-1, -1, data.shape[-1])) | |
| def knn_indices(ref: torch.Tensor, query: torch.Tensor, k: int) -> torch.Tensor: | |
| """`knn_cuda.KNN(k, transpose_mode=True)(ref, query)[1]`. | |
| ref: (B, Nr, 3), query: (B, Nq, 3) -> (B, Nq, k) indices into Nr, | |
| ordered by ascending distance. | |
| """ | |
| d = torch.cdist(query, ref) # (B, Nq, Nr) | |
| return d.topk(k, dim=-1, largest=False).indices | |
| # -------------------------------------------------------------------------------------- | |
| # Point-Transformer encoder (verbatim port of src/models/point_transformer.py) | |
| # -------------------------------------------------------------------------------------- | |
| class DropPath(nn.Module): | |
| """Stochastic depth. Identity at inference time (which is all we do here).""" | |
| def __init__(self, drop_prob: float = 0.0): | |
| super().__init__() | |
| self.drop_prob = drop_prob | |
| def forward(self, x): | |
| if self.drop_prob == 0.0 or not self.training: | |
| return x | |
| keep = 1.0 - self.drop_prob | |
| shape = (x.shape[0],) + (1,) * (x.ndim - 1) | |
| mask = x.new_empty(shape).bernoulli_(keep).div_(keep) | |
| return x * mask | |
| class PatchedGroup(nn.Module): | |
| """Same as upstream `PatchedGroup`, with FPS/KNN swapped for the torch versions.""" | |
| def __init__(self, num_group: int, group_size: int): | |
| super().__init__() | |
| self.num_group = num_group | |
| self.group_size = group_size | |
| def forward(self, xyz: torch.Tensor): | |
| batch_size, num_points, C = xyz.shape | |
| if C > 3: | |
| xyz_only = xyz[:, :, :3].contiguous() | |
| extra = xyz[:, :, 3:].contiguous() | |
| else: | |
| xyz_only = xyz.contiguous() | |
| extra = None | |
| center = fps(xyz_only, self.num_group) # (B, G, 3) | |
| idx = knn_indices(xyz_only, center, self.group_size) # (B, G, M) | |
| idx_rel = idx.clone() | |
| idx_base = torch.arange(0, batch_size, device=xyz.device).view(-1, 1, 1) * num_points | |
| idx_flat = (idx + idx_base).view(-1) | |
| neigh_xyz = xyz_only.reshape(batch_size * num_points, -1)[idx_flat, :].view( | |
| batch_size, self.num_group, self.group_size, 3 | |
| ) | |
| if extra is not None: | |
| neigh_extra = extra.reshape(batch_size * num_points, -1)[idx_flat, :].view( | |
| batch_size, self.num_group, self.group_size, -1 | |
| ) | |
| neighborhood = torch.cat((neigh_xyz - center.unsqueeze(2), neigh_extra), dim=-1) | |
| else: | |
| neighborhood = neigh_xyz - center.unsqueeze(2) | |
| return neighborhood.contiguous(), center.contiguous(), idx_rel | |
| class Encoder(nn.Module): | |
| def __init__(self, encoder_channel: int, color: bool = False): | |
| super().__init__() | |
| self.encoder_channel = encoder_channel | |
| self.first_conv = nn.Sequential( | |
| nn.Conv1d(6 if color else 3, 128, 1), | |
| nn.BatchNorm1d(128), | |
| nn.ReLU(inplace=True), | |
| nn.Conv1d(128, 256, 1), | |
| ) | |
| self.second_conv = nn.Sequential( | |
| nn.Conv1d(512, 512, 1), | |
| nn.BatchNorm1d(512), | |
| nn.ReLU(inplace=True), | |
| nn.Conv1d(512, self.encoder_channel, 1), | |
| ) | |
| def forward(self, point_groups): | |
| bs, g, n, c = point_groups.shape | |
| point_groups = point_groups.reshape(bs * g, n, c).permute(0, 2, 1) | |
| feature = self.first_conv(point_groups) | |
| feature_global = torch.max(feature, 2, keepdim=True)[0] | |
| feature_global = feature_global.repeat(1, 1, n) | |
| feature = torch.cat([feature_global, feature], 1) | |
| feature = self.second_conv(feature) | |
| feature = feature.max(dim=2)[0] | |
| return feature.reshape(bs, g, self.encoder_channel).contiguous() | |
| class MLP(nn.Module): | |
| def __init__(self, in_features, hidden_features=None, out_features=None, act_layer=nn.GELU, drop=0.0): | |
| super().__init__() | |
| out_features = out_features or in_features | |
| hidden_features = hidden_features or in_features | |
| self.fc1 = nn.Linear(in_features, hidden_features) | |
| self.act = act_layer() | |
| self.fc2 = nn.Linear(hidden_features, out_features) | |
| self.drop = nn.Dropout(drop) | |
| def forward(self, x): | |
| x = self.fc1(x) | |
| x = self.act(x) | |
| x = self.drop(x) | |
| x = self.fc2(x) | |
| x = self.drop(x) | |
| return x | |
| class Attention(nn.Module): | |
| def __init__(self, dim, num_heads=8, qkv_bias=False, qk_scale=None, attn_drop=0.0, proj_drop=0.0): | |
| super().__init__() | |
| self.num_heads = num_heads | |
| head_dim = dim // num_heads | |
| self.scale = qk_scale or head_dim ** -0.5 | |
| self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias) | |
| self.attn_drop = nn.Dropout(attn_drop) | |
| self.proj = nn.Linear(dim, dim) | |
| self.proj_drop = nn.Dropout(proj_drop) | |
| def forward(self, x): | |
| B, N, C = x.shape | |
| qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, C // self.num_heads).permute(2, 0, 3, 1, 4) | |
| q, k, v = qkv[0], qkv[1], qkv[2] | |
| attn = (q @ k.transpose(-2, -1)) * self.scale | |
| attn = attn.softmax(dim=-1) | |
| attn = self.attn_drop(attn) | |
| x = (attn @ v).transpose(1, 2).reshape(B, N, C) | |
| x = self.proj(x) | |
| return self.proj_drop(x) | |
| class Block(nn.Module): | |
| def __init__(self, dim, num_heads, mlp_ratio=4.0, qkv_bias=False, qk_scale=None, | |
| drop=0.0, attn_drop=0.0, drop_path=0.0, act_layer=nn.GELU): | |
| super().__init__() | |
| self.norm1 = nn.LayerNorm(dim) | |
| self.attn = Attention(dim, num_heads=num_heads, qkv_bias=qkv_bias, qk_scale=qk_scale, | |
| attn_drop=attn_drop, proj_drop=drop) | |
| self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity() | |
| self.norm2 = nn.LayerNorm(dim) | |
| self.mlp = MLP(in_features=dim, hidden_features=int(dim * mlp_ratio), act_layer=act_layer, drop=drop) | |
| def forward(self, x): | |
| x = x + self.drop_path(self.attn(self.norm1(x))) | |
| x = x + self.drop_path(self.mlp(self.norm2(x))) | |
| return x | |
| class TransformerEncoder(nn.Module): | |
| def __init__(self, embed_dim=768, depth=4, num_heads=12, mlp_ratio=4.0, qkv_bias=False, | |
| qk_scale=None, drop_rate=0.0, attn_drop_rate=0.0, drop_path_rate=0.0): | |
| super().__init__() | |
| def _drop_for_block(i): | |
| if isinstance(drop_path_rate, (list, tuple)): | |
| return drop_path_rate[i] | |
| return drop_path_rate | |
| self.blocks = nn.ModuleList([ | |
| Block(dim=embed_dim, num_heads=num_heads, mlp_ratio=mlp_ratio, qkv_bias=qkv_bias, | |
| qk_scale=qk_scale, drop=drop_rate, attn_drop=attn_drop_rate, | |
| drop_path=_drop_for_block(i)) | |
| for i in range(depth) | |
| ]) | |
| def forward(self, x, pos): | |
| for blk in self.blocks: | |
| x = blk(x + pos) | |
| return x | |
| class PointTransformer(nn.Module): | |
| """Upstream `point_transformer.get_model`.""" | |
| def __init__(self, num_group=DEFAULT_NUM_GROUP, group_size=DEFAULT_GROUP_SIZE, color=False): | |
| super().__init__() | |
| self.trans_dim = TRANS_DIM | |
| self.depth = DEPTH | |
| self.num_heads = NUM_HEADS | |
| self.encoder_dims = ENCODER_DIMS | |
| self.color = color | |
| self.group_size = group_size | |
| self.num_group = num_group | |
| self.group_divider = PatchedGroup(num_group=num_group, group_size=group_size) | |
| self.encoder = Encoder(encoder_channel=self.encoder_dims, color=color) | |
| self.reduce_dim = nn.Linear(self.encoder_dims, self.trans_dim) | |
| self.cls_token = nn.Parameter(torch.zeros(1, 1, self.trans_dim)) | |
| self.cls_pos = nn.Parameter(torch.randn(1, 1, self.trans_dim)) | |
| self.pos_embed = nn.Sequential(nn.Linear(3, 128), nn.GELU(), nn.Linear(128, self.trans_dim)) | |
| dpr = [x.item() for x in torch.linspace(0, DROP_PATH_RATE, self.depth)] | |
| self.blocks = TransformerEncoder(embed_dim=self.trans_dim, depth=self.depth, | |
| drop_path_rate=dpr, num_heads=self.num_heads) | |
| self.norm = nn.LayerNorm(self.trans_dim) | |
| def set_grouping(self, num_group: int, group_size: int) -> None: | |
| self.group_divider.num_group = int(num_group) | |
| self.group_divider.group_size = int(group_size) | |
| def forward_patches(self, pts: torch.Tensor): | |
| """pts: (B, C, N) with C >= 3. Returns patch_emb (B, D, G), centers (B, 3, G), idx (B, G, M).""" | |
| pts_bn = pts.transpose(-1, -2).contiguous() | |
| neighborhood, center, patch_idx = self.group_divider(pts_bn) | |
| group_tokens = self.encoder(neighborhood) | |
| group_tokens = self.reduce_dim(group_tokens) | |
| cls_tokens = self.cls_token.expand(group_tokens.size(0), -1, -1) | |
| cls_pos = self.cls_pos.expand(group_tokens.size(0), -1, -1) | |
| pos = self.pos_embed(center) | |
| x = torch.cat((cls_tokens, group_tokens), dim=1) | |
| pos = torch.cat((cls_pos, pos), dim=1) | |
| feature = self.blocks(x, pos) | |
| patch_emb = self.norm(feature)[:, 1:, :].transpose(-1, -2).contiguous() | |
| patch_centers = center.transpose(-1, -2).contiguous() | |
| return patch_emb, patch_centers, patch_idx | |
| class PatchToTextProj(nn.Module): | |
| def __init__(self, in_dim: int, out_dim: int): | |
| super().__init__() | |
| self.proj = nn.Linear(in_dim, out_dim) | |
| def forward(self, patch_emb): | |
| x = patch_emb.transpose(1, 2) | |
| x = self.proj(x) | |
| return F.normalize(x, dim=-1) | |
| # -------------------------------------------------------------------------------------- | |
| # Geometry helpers | |
| # -------------------------------------------------------------------------------------- | |
| def pc_normalize(pc: np.ndarray) -> np.ndarray: | |
| """Upstream `pc_normalize`: centre, then scale to the unit sphere.""" | |
| centroid = pc.mean(axis=0) | |
| pc = pc - centroid | |
| m = np.max(np.sqrt((pc ** 2).sum(axis=1))) | |
| if m <= 0: | |
| m = 1.0 | |
| return pc / m | |
| def prepare_points(points: torch.Tensor) -> torch.Tensor: | |
| """Upstream `prepare_points`: (B,N,C) -> (B,C,N) with the Y/Z axes swapped.""" | |
| if points.ndim != 3: | |
| raise ValueError(f"Expected (B,N,C), got {tuple(points.shape)}") | |
| pts = points.transpose(2, 1).contiguous() | |
| pts[:, [1, 2], :] = pts[:, [2, 1], :] | |
| return pts | |
| def assign_points_from_patches(points_xyz, patch_centers, patch_logits, patch_idx, mode="nearest"): | |
| """Upstream `assign_points_from_patches` (knn_cuda replaced by cdist/argmin).""" | |
| B, _, N = points_xyz.shape | |
| K = patch_logits.shape[-1] | |
| if mode == "membership": | |
| point_logits = torch.zeros(B, N, K, device=points_xyz.device, dtype=patch_logits.dtype) | |
| counts = torch.zeros(B, N, 1, device=points_xyz.device, dtype=patch_logits.dtype) | |
| for b in range(B): | |
| idx = patch_idx[b].reshape(-1) | |
| src = patch_logits[b].unsqueeze(1).expand_as(patch_idx[b].unsqueeze(-1).expand(-1, -1, K)).reshape(-1, K) | |
| point_logits[b].index_add_(0, idx, src) | |
| ones = torch.ones(idx.shape[0], 1, device=points_xyz.device, dtype=patch_logits.dtype) | |
| counts[b].index_add_(0, idx, ones) | |
| return point_logits / counts.clamp_min(1.0) | |
| nearest = knn_indices(patch_centers.transpose(1, 2).contiguous(), | |
| points_xyz.transpose(1, 2).contiguous(), 1).squeeze(-1) | |
| return patch_logits.gather(1, nearest.unsqueeze(-1).expand(-1, -1, K)) | |
| # -------------------------------------------------------------------------------------- | |
| # Text side | |
| # -------------------------------------------------------------------------------------- | |
| def build_prompts(name: str, category: str, setting: str) -> list[str]: | |
| """Prompt ensemble for one part label, mirroring `eval.py:encode_texts`.""" | |
| nm = clean_text(name) | |
| cname = clean_text(category or "") | |
| texts: list[str] = [] | |
| if setting in ("part_plus_cat", "ensemble") and cname: | |
| for tpl in PART_PLUS_CAT_TEMPLATES: | |
| slots = tpl.count("{}") | |
| if slots == 2: | |
| texts.append(tpl.format(nm, cname)) | |
| elif slots == 1: | |
| texts.append(tpl.format(f"{cname} {nm}")) | |
| else: | |
| texts.append(f"{cname} {nm}") | |
| if (setting in ("part_only", "ensemble")) or not cname: | |
| for tpl in PART_ONLY_TEMPLATES: | |
| texts.append(tpl.format(nm) if tpl.count("{}") == 1 else nm) | |
| return texts or [nm] | |
| def encode_labels(names, category, setting, text_model, tokenizer, device) -> torch.Tensor: | |
| """One L2-normalised CLIP text embedding per label -> (K, 1280).""" | |
| per_label = [] | |
| for nm in names: | |
| prompts = build_prompts(nm, category, setting) | |
| toks = tokenizer(prompts, padding="max_length", max_length=tokenizer.model_max_length, | |
| truncation=True, return_tensors="pt").to(device) | |
| feat = text_model(**toks).text_embeds.float() | |
| feat = F.normalize(feat, dim=-1) | |
| per_label.append(F.normalize(feat.mean(dim=0, keepdim=True), dim=-1)) | |
| return torch.cat(per_label, dim=0) | |
| # -------------------------------------------------------------------------------------- | |
| # Checkpoint | |
| # -------------------------------------------------------------------------------------- | |
| def load_patchalign3d(ckpt_path: str, num_group=DEFAULT_NUM_GROUP, group_size=DEFAULT_GROUP_SIZE): | |
| model = PointTransformer(num_group=num_group, group_size=group_size, color=False) | |
| proj = PatchToTextProj(in_dim=TRANS_DIM, out_dim=CLIP_TEXT_DIM) | |
| ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=False) | |
| if "model" in ckpt: | |
| res = model.load_state_dict(ckpt["model"], strict=False) | |
| print(f"[ckpt] encoder: missing={len(res.missing_keys)} unexpected={len(res.unexpected_keys)}") | |
| if res.missing_keys: | |
| print(" missing:", res.missing_keys) | |
| if res.unexpected_keys: | |
| print(" unexpected:", res.unexpected_keys) | |
| else: | |
| raise RuntimeError("checkpoint has no 'model' entry") | |
| if "proj" in ckpt: | |
| res = proj.load_state_dict(ckpt["proj"], strict=False) | |
| print(f"[ckpt] proj: missing={len(res.missing_keys)} unexpected={len(res.unexpected_keys)}") | |
| else: | |
| raise RuntimeError("checkpoint has no 'proj' entry") | |
| return model.eval(), proj.eval() | |
| # -------------------------------------------------------------------------------------- | |
| # End-to-end segmentation | |
| # -------------------------------------------------------------------------------------- | |
| def segment_point_cloud(points_np, label_names, model, proj, text_model, tokenizer, device, | |
| category="", text_setting="part_only", assign="nearest", | |
| tau=DEFAULT_TAU, num_group=DEFAULT_NUM_GROUP, group_size=DEFAULT_GROUP_SIZE): | |
| """points_np: (N,3) float array in original coordinates. Returns (pred, probs).""" | |
| model.set_grouping(num_group, group_size) | |
| pts = torch.as_tensor(np.ascontiguousarray(points_np[:, :3]), dtype=torch.float32).unsqueeze(0) | |
| pts = prepare_points(pts).to(device) | |
| patch_emb, patch_centers, patch_idx = model.forward_patches(pts) | |
| patch_feat = proj(patch_emb) | |
| text_feats = encode_labels(label_names, category, text_setting, text_model, tokenizer, device) | |
| logits = (patch_feat @ text_feats.t()) / max(float(tau), 1e-6) | |
| point_logits = assign_points_from_patches(pts[:, :3, :], patch_centers, logits, patch_idx, mode=assign) | |
| probs = point_logits.softmax(dim=-1).squeeze(0) | |
| pred = point_logits.argmax(dim=-1).squeeze(0) | |
| return pred.cpu().numpy().astype(np.int64), probs.cpu().numpy().astype(np.float32) | |