scatteringnet / src /occupancy_encoder.py
scatteringnet-space
Slim Gradio Space: infer + demo only
bc4c433
Raw History Blame Contribute Delete
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))