Spaces:
Running on Zero
Running on Zero
Download src/occupancy_encoder.py from guyPerry/scatteringnet: direct link, hf CLI and curl.
- Browser
- Download file 9.13 kB
-
https://huggingface.co/spaces/guyPerry/scatteringnet/resolve/main/src/occupancy_encoder.py
- Command line
-
hf download hf://spaces/guyPerry/scatteringnet/src/occupancy_encoder.py
-
curl -L -o occupancy_encoder.py https://huggingface.co/spaces/guyPerry/scatteringnet/resolve/main/src/occupancy_encoder.py
9.13 kB
| """Geometry-conditioned occupancy: query XYZ plus a shape latent. | |
| ``OccupancyMLP`` stays xyz-only. ``shape_encoder: surface`` is the | |
| envelope PointNet over ``(N, 6)`` XYZ + face normal (or ``(N, 3)`` | |
| for older checkpoints). ``knn_k > 0`` adds per-query nearest-neighbor | |
| features (XYZ offset, plus the neighbor normal when the cloud is 6-D). | |
| The face-token head (``shape_encoder: mesh``) was removed. | |
| """ | |
| from __future__ import annotations | |
| import torch | |
| import torch.nn as nn | |
| from torch import Tensor | |
| from scatteringnet.occupancy_mlp import build_mlp | |
| # Distinct from OccupancyMLP so infer can tell the checkpoint apart. | |
| CHECKPOINT_KIND = "occupancy_encoder" | |
| class SurfaceEncoder(nn.Module): | |
| """ | |
| Per-point MLP + max-pool (PointNet) over an envelope cloud. | |
| Shapes | |
| ------ | |
| envelope: ``(U, N, C)`` unique meshes; ``C`` is 3 (XYZ) or 6 (XYZ+n) | |
| output: ``(U, D)`` one latent per mesh | |
| """ | |
| def __init__( | |
| self, latent_dim: int = 64, hidden: int = 64, *, in_dim: int = 3 | |
| ) -> None: | |
| super().__init__() | |
| if latent_dim < 1: | |
| raise ValueError(f"latent_dim must be >= 1, got {latent_dim}") | |
| if hidden < 1: | |
| raise ValueError(f"hidden must be >= 1, got {hidden}") | |
| dim = int(in_dim) | |
| if dim not in (3, 6): | |
| raise ValueError(f"in_dim must be 3 or 6, got {dim}") | |
| self.latent_dim = latent_dim | |
| self.hidden = hidden | |
| self.in_dim = dim | |
| self.point_mlp = nn.Sequential( | |
| nn.Linear(dim, hidden), | |
| nn.ReLU(inplace=True), | |
| nn.Linear(hidden, latent_dim), | |
| ) | |
| def forward(self, envelope: Tensor) -> Tensor: | |
| """Max-pool per-point features → one vector per unique mesh.""" | |
| if envelope.ndim != 3 or envelope.shape[-1] != self.in_dim: | |
| raise ValueError( | |
| f"envelope must have shape (U, N, {self.in_dim}), " | |
| f"got {tuple(envelope.shape)}" | |
| ) | |
| features = self.point_mlp(envelope) | |
| return features.max(dim=1).values | |
| def knn_offsets(xyz: Tensor, envelope: Tensor, k: int) -> Tensor: | |
| """ | |
| Offsets from each query to its ``k`` nearest envelope points. | |
| Distances use XYZ only (AABB Euclidean). ``k`` is clamped to N. | |
| If the cloud is ``(B, N, 6)``, each neighbor is | |
| ``(dx, dy, dz, nx, ny, nz)`` — relative position plus that | |
| neighbor's stored normal. Query points have no normal. | |
| Shapes: ``xyz (B, 3)``, ``envelope (B, N, 3|6)`` → ``(B, k, 3|6)``. | |
| """ | |
| if xyz.ndim != 2 or xyz.shape[-1] != 3: | |
| raise ValueError(f"xyz must have shape (B, 3), got {tuple(xyz.shape)}") | |
| if envelope.ndim != 3 or envelope.shape[-1] not in (3, 6): | |
| raise ValueError( | |
| f"envelope must have shape (B, N, 3 or 6), got {tuple(envelope.shape)}" | |
| ) | |
| if int(xyz.shape[0]) != int(envelope.shape[0]): | |
| raise ValueError( | |
| f"xyz/envelope batch mismatch: {tuple(xyz.shape)} vs {tuple(envelope.shape)}" | |
| ) | |
| n_env = int(envelope.shape[1]) | |
| if n_env < 1: | |
| raise ValueError("envelope length N must be >= 1") | |
| take = min(int(k), n_env) | |
| if take < 1: | |
| raise ValueError(f"k must be >= 1, got {k}") | |
| feat = int(envelope.shape[-1]) | |
| # k-NN is position-only; extras (normals) ride along after the gather. | |
| pos = envelope[..., :3] | |
| dist = torch.linalg.norm(pos - xyz.unsqueeze(1), dim=-1) | |
| idx = dist.topk(take, dim=-1, largest=False).indices | |
| nbrs = torch.gather(envelope, 1, idx.unsqueeze(-1).expand(-1, -1, feat)) | |
| rel_xyz = nbrs[..., :3] - xyz.unsqueeze(1) | |
| if feat == 3: | |
| return rel_xyz | |
| return torch.cat([rel_xyz, nbrs[..., 3:]], dim=-1) | |
| # Catalog trains used YAML seed 1. Old best.pt files omit ``seed``. | |
| DEFAULT_ENVELOPE_SEED = 1 | |
| def envelope_seed_from_ckpt(ckpt: dict) -> int: | |
| """ | |
| Envelope RNG seed this checkpoint was trained with. | |
| New trains store ``seed`` on ``best.pt``. Older files omit it; do | |
| not fall back to live YAML (that knob may have changed since train). | |
| """ | |
| raw = ckpt.get("seed") | |
| if raw is None: | |
| return DEFAULT_ENVELOPE_SEED | |
| return int(raw) | |
| def envelope_dim_from_ckpt(ckpt: dict) -> int: | |
| """ | |
| Envelope channel count this checkpoint was trained with. | |
| New trains store ``envelope_dim``. Older XYZ-only ``best.pt`` files | |
| omit it; the first SurfaceEncoder Linear in-features is then 3. | |
| """ | |
| raw = ckpt.get("envelope_dim") | |
| if raw is not None: | |
| dim = int(raw) | |
| if dim not in (3, 6): | |
| raise ValueError(f"envelope_dim must be 3 or 6, got {dim}") | |
| return dim | |
| weight = (ckpt.get("state_dict") or {}).get("surface.point_mlp.0.weight") | |
| if weight is not None: | |
| dim = int(weight.shape[1]) | |
| if dim in (3, 6): | |
| return dim | |
| return 3 | |
| class OccupancyEncoder(nn.Module): | |
| """ | |
| Occupancy logits from query XYZ and an envelope code. | |
| Unique ``shape_id`` values are encoded **once** per batch, then | |
| broadcast. ``knn_k > 0`` concatenates a local envelope code. | |
| Shapes | |
| ------ | |
| xyz: ``(B, 3)`` | |
| geom: ``(B, N, C)`` envelope; ``C`` is ``envelope_dim`` (3 or 6) | |
| shape_id: ``(B,)`` long | |
| output: ``(B, 1)`` logits | |
| """ | |
| def __init__( | |
| self, | |
| hidden: int = 64, | |
| depth: int = 4, | |
| latent_dim: int = 64, | |
| *, | |
| shape_encoder: str = "surface", | |
| knn_k: int = 0, | |
| knn_local_dim: int | None = None, | |
| envelope_dim: int = 6, | |
| ) -> None: | |
| super().__init__() | |
| if hidden < 1: | |
| raise ValueError(f"hidden must be >= 1, got {hidden}") | |
| if depth < 1: | |
| raise ValueError(f"depth must be >= 1, got {depth}") | |
| if latent_dim < 1: | |
| raise ValueError(f"latent_dim must be >= 1, got {latent_dim}") | |
| kind = str(shape_encoder).strip().lower() | |
| if kind != "surface": | |
| raise ValueError( | |
| "OccupancyEncoder only supports shape_encoder='surface' " | |
| f"(face-token 'mesh' was removed), got {shape_encoder!r}" | |
| ) | |
| k = int(knn_k) | |
| if k < 0: | |
| raise ValueError(f"knn_k must be >= 0, got {k}") | |
| local_dim = int(knn_local_dim) if knn_local_dim is not None else int(latent_dim) | |
| if k > 0 and local_dim < 1: | |
| raise ValueError(f"knn_local_dim must be >= 1, got {local_dim}") | |
| ed = int(envelope_dim) | |
| if ed not in (3, 6): | |
| raise ValueError(f"envelope_dim must be 3 or 6, got {ed}") | |
| self.hidden = hidden | |
| self.depth = depth | |
| self.latent_dim = latent_dim | |
| self.shape_encoder = kind | |
| self.knn_k = k | |
| self.knn_local_dim = local_dim if k > 0 else 0 | |
| self.envelope_dim = ed | |
| # Name ``surface`` is load-stable for existing envelope checkpoints. | |
| self.surface = SurfaceEncoder(latent_dim=latent_dim, hidden=hidden, in_dim=ed) | |
| self.token_dim = ed | |
| head_in = 3 + latent_dim | |
| if k > 0: | |
| # Same PointNet block as the global envelope, over k neighbor features. | |
| self.local = SurfaceEncoder(latent_dim=local_dim, hidden=hidden, in_dim=ed) | |
| head_in += local_dim | |
| self.head = build_mlp(head_in, hidden, depth) | |
| def encode_unique(self, geom: Tensor, shape_id: Tensor) -> Tensor: | |
| """ | |
| Encode each distinct ``shape_id`` once and scatter back to ``(B, D)``. | |
| Parameters | |
| ---------- | |
| geom, shape_id: | |
| Batched envelope clouds and integer mesh ids (same length B). | |
| """ | |
| if shape_id.ndim != 1 or int(shape_id.shape[0]) != int(geom.shape[0]): | |
| raise ValueError( | |
| f"shape_id must be (B,), got {tuple(shape_id.shape)} " | |
| f"for geom {tuple(geom.shape)}" | |
| ) | |
| unique_ids, inverse = torch.unique(shape_id, sorted=True, return_inverse=True) | |
| hits = shape_id.unsqueeze(0) == unique_ids.unsqueeze(1) | |
| first = hits.to(dtype=torch.int64).argmax(dim=1) | |
| z_unique = self.surface(geom[first]) | |
| return z_unique[inverse] | |
| def forward( | |
| self, | |
| xyz: Tensor, | |
| geom: Tensor, | |
| shape_id: Tensor, | |
| ) -> Tensor: | |
| """``cat(xyz, z_global[, z_local])`` → occupancy logit.""" | |
| if xyz.ndim != 2 or xyz.shape[-1] != 3: | |
| raise ValueError(f"xyz must have shape (B, 3), got {tuple(xyz.shape)}") | |
| if geom.ndim != 3 or geom.shape[-1] != self.token_dim: | |
| raise ValueError( | |
| f"geom must have shape (B, K, {self.token_dim}), got {tuple(geom.shape)}" | |
| ) | |
| if int(xyz.shape[0]) != int(geom.shape[0]): | |
| raise ValueError( | |
| f"xyz/geom batch mismatch: {tuple(xyz.shape)} vs {tuple(geom.shape)}" | |
| ) | |
| ids = shape_id.reshape(-1) | |
| z_shape = self.encode_unique(geom, ids) | |
| pieces = [xyz, z_shape] | |
| if self.knn_k > 0: | |
| rel = knn_offsets(xyz, geom, self.knn_k) | |
| pieces.append(self.local(rel)) | |
| return self.head(torch.cat(pieces, dim=-1)) | |