File size: 5,534 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
"""Surface envelope samples from a triangle mesh.

Area-weighted face darts (larger triangles get more samples). Each
sample is ``(x, y, z, nx, ny, nz)``: position plus the unit normal of
the triangle it sits on. World clouds are cached per
``(mesh_key, n_surface, seed)``. AABB is applied to XYZ only.
"""

from __future__ import annotations

import numpy as np
import trimesh
from numpy.typing import NDArray
from trimesh.sample import sample_surface

from scatteringnet.normalize import apply_normalization

PointsArray = NDArray[np.float32]
ENVELOPE_XYZ_DIM = 3
ENVELOPE_FEAT_DIM = 6

# Same OBJ + count + seed → same world samples (always 6-D).
_ENVELOPE_CACHE: dict[tuple[str, int, int], PointsArray] = {}


def clear_envelope_cache() -> None:
    """Drop cached world-space envelopes (tests / long-lived notebooks)."""
    _ENVELOPE_CACHE.clear()


def _triangle_normals(verts: np.ndarray, faces: np.ndarray) -> NDArray[np.float64]:
    v0 = verts[faces[:, 0]]
    v1 = verts[faces[:, 1]]
    v2 = verts[faces[:, 2]]
    cross = np.cross(v1 - v0, v2 - v0)
    length = np.linalg.norm(cross, axis=1, keepdims=True)
    ok = length[:, 0] > 1e-12
    normals = np.zeros_like(cross)
    normals[ok] = cross[ok] / length[ok]
    return normals


def _pack_xyz_normal(xyz: np.ndarray, normals: np.ndarray) -> PointsArray:
    """Concatenate XYZ with unit face normals → ``(N, 6)``."""
    pos = np.asarray(xyz, dtype=np.float32)
    nrm = np.asarray(normals, dtype=np.float32)
    if pos.shape != nrm.shape or pos.ndim != 2 or pos.shape[1] != 3:
        raise ValueError(
            f"xyz/normals must be (N, 3), got {tuple(pos.shape)} / {tuple(nrm.shape)}"
        )
    length = np.linalg.norm(nrm, axis=1, keepdims=True)
    ok = length[:, 0] > 1e-12
    unit = np.zeros_like(nrm)
    unit[ok] = nrm[ok] / length[ok]
    return np.concatenate([pos, unit], axis=1)


def apply_envelope_aabb(
    env: np.ndarray, center: np.ndarray, scale: float
) -> PointsArray:
    """AABB-normalize XYZ; leave normals as unit directions."""
    arr = np.asarray(env, dtype=np.float32)
    if arr.ndim != 2 or arr.shape[1] not in (ENVELOPE_XYZ_DIM, ENVELOPE_FEAT_DIM):
        raise ValueError(
            f"envelope must be (N, 3) or (N, 6), got {tuple(arr.shape)}"
        )
    out = arr.copy()
    out[:, :3] = apply_normalization(out[:, :3], center, scale)
    return out


def undo_envelope_aabb(
    env: np.ndarray, center: np.ndarray, scale: float
) -> PointsArray:
    """Undo AABB on XYZ only."""
    arr = np.asarray(env, dtype=np.float32)
    if arr.ndim != 2 or arr.shape[1] not in (ENVELOPE_XYZ_DIM, ENVELOPE_FEAT_DIM):
        raise ValueError(
            f"envelope must be (N, 3) or (N, 6), got {tuple(arr.shape)}"
        )
    out = arr.copy()
    c = np.asarray(center, dtype=np.float32).reshape(3)
    out[:, :3] = out[:, :3] * np.float32(scale) + c
    return out


def project_envelope_dim(env: np.ndarray, dim: int) -> PointsArray:
    """Keep XYZ+normal or drop to XYZ for an older checkpoint."""
    arr = np.asarray(env, dtype=np.float32)
    want = int(dim)
    if want == ENVELOPE_FEAT_DIM:
        if arr.ndim != 2 or arr.shape[1] != ENVELOPE_FEAT_DIM:
            raise ValueError(f"expected (N, 6) envelope, got {tuple(arr.shape)}")
        return arr
    if want == ENVELOPE_XYZ_DIM:
        if arr.ndim != 2 or arr.shape[1] not in (ENVELOPE_XYZ_DIM, ENVELOPE_FEAT_DIM):
            raise ValueError(f"expected (N, 3|6) envelope, got {tuple(arr.shape)}")
        return arr[:, :3]
    raise ValueError(f"envelope dim must be 3 or 6, got {want}")


def _sample_area(
    verts: np.ndarray, tris: np.ndarray, count: int, seed: int
) -> PointsArray:
    mesh = trimesh.Trimesh(vertices=verts, faces=tris, process=False)
    points, face_idx = sample_surface(mesh, count, seed=int(seed))
    nrm = _triangle_normals(verts, tris)[np.asarray(face_idx, dtype=np.int64)]
    out = _pack_xyz_normal(points, nrm)
    if out.shape != (count, ENVELOPE_FEAT_DIM):
        raise ValueError(
            f"expected envelope shape {(count, ENVELOPE_FEAT_DIM)}, got {tuple(out.shape)}"
        )
    return out


def sample_surface_points(
    vertices: np.ndarray,
    faces: np.ndarray,
    n_surface: int,
    *,
    seed: int = 1,
    cache_key: str | None = None,
) -> PointsArray:
    """
    Sample ``n_surface`` envelope points as ``(N, 6)`` XYZ + unit normal.

    Face-area weighted: larger triangles receive more darts. No crease
    or fold path — unused ``envelope_mix`` on old checkpoints is ignored.
    """
    count = int(n_surface)
    if count < 1:
        raise ValueError(f"n_surface must be >= 1, got {count}")
    verts = np.asarray(vertices, dtype=np.float64)
    tris = np.asarray(faces, dtype=np.int64)
    if verts.ndim != 2 or verts.shape[1] != 3:
        raise ValueError(f"vertices must have shape (V, 3), got {tuple(verts.shape)}")
    if tris.ndim != 2 or tris.shape[1] != 3:
        raise ValueError(f"faces must have shape (T, 3), got {tuple(tris.shape)}")

    key: tuple[str, int, int] | None = None
    if cache_key is not None:
        key = (str(cache_key), count, int(seed))
        cached = _ENVELOPE_CACHE.get(key)
        if cached is not None:
            return cached

    out = _sample_area(verts, tris, count, int(seed))
    if out.shape != (count, ENVELOPE_FEAT_DIM):
        raise ValueError(
            f"expected envelope shape {(count, ENVELOPE_FEAT_DIM)}, got {tuple(out.shape)}"
        )
    if key is not None:
        _ENVELOPE_CACHE[key] = out
    return out