Spaces:
Running on Zero
Running on Zero
File size: 9,125 Bytes
bc4c433 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 | """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))
|