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