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))