scatteringnet-space commited on
Commit
bc4c433
·
0 Parent(s):

Slim Gradio Space: infer + demo only

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitignore +1 -0
  2. README.md +27 -0
  3. config.yaml +73 -0
  4. docs/inspect_checkpoint.yaml +5 -0
  5. models/.gitkeep +0 -0
  6. pyproject.toml +22 -0
  7. requirements.txt +4 -0
  8. scatteringnet/__init__.py +1 -0
  9. scatteringnet/config.py +611 -0
  10. scatteringnet/data_npz.py +460 -0
  11. scatteringnet/geometry/__init__.py +10 -0
  12. scatteringnet/geometry/mesh_io.py +80 -0
  13. scatteringnet/geometry/surface.py +154 -0
  14. scatteringnet/geometry/trimesh_util.py +19 -0
  15. scatteringnet/infer_multi_npz.py +202 -0
  16. scatteringnet/metrics.py +156 -0
  17. scatteringnet/normalize.py +104 -0
  18. scatteringnet/occupancy_encoder.py +248 -0
  19. scatteringnet/occupancy_mlp.py +87 -0
  20. scatteringnet/viewer/__init__.py +1 -0
  21. scatteringnet/viewer/infer_job.py +536 -0
  22. scatteringnet/viewer/mesh_access.py +52 -0
  23. scatteringnet/viewer/model_access.py +183 -0
  24. scatteringnet/viewer/obj_fill.py +138 -0
  25. src/__init__.py +1 -0
  26. src/config.py +611 -0
  27. src/data_npz.py +460 -0
  28. src/geometry/__init__.py +10 -0
  29. src/geometry/mesh_io.py +80 -0
  30. src/geometry/surface.py +154 -0
  31. src/geometry/trimesh_util.py +19 -0
  32. src/gradio/app.py +503 -0
  33. src/gradio/examples/Helix_bend.obj +0 -0
  34. src/gradio/examples/Obese.obj +0 -0
  35. src/gradio/examples/Player.obj +0 -0
  36. src/gradio/examples/TorusX3_box.obj +0 -0
  37. src/gradio/examples/dog.obj +0 -0
  38. src/gradio/examples/horse.obj +0 -0
  39. src/gradio/figure.py +255 -0
  40. src/gradio/orbit.js +564 -0
  41. src/gradio/pipeline.py +133 -0
  42. src/infer_multi_npz.py +202 -0
  43. src/metrics.py +156 -0
  44. src/normalize.py +104 -0
  45. src/occupancy_encoder.py +248 -0
  46. src/occupancy_mlp.py +87 -0
  47. src/viewer/__init__.py +1 -0
  48. src/viewer/infer_job.py +536 -0
  49. src/viewer/mesh_access.py +52 -0
  50. src/viewer/model_access.py +183 -0
.gitignore ADDED
@@ -0,0 +1 @@
 
 
1
+ *.pt
README.md ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ title: scatteringNet
3
+ emoji: 🧊
4
+ colorFrom: gray
5
+ colorTo: yellow
6
+ sdk: gradio
7
+ app_file: src/gradio/app.py
8
+ pinned: false
9
+ license: mit
10
+ short_description: "Occupancy fill: label inside points on a 3D mesh."
11
+ ---
12
+
13
+ # scatteringNet — occupancy fill
14
+
15
+ Public Gradio demo. Upload an OBJ (or pick a sample) and **Run model** to label query points inside the solid.
16
+
17
+ [Code](https://github.com/PerryGu/scatteringNet)
18
+
19
+ [Video](https://youtu.be/vU45O0Mu0o4)
20
+
21
+ These sample meshes were not in the training catalog.
22
+
23
+ **Weights.** After this slim tree is on the Space, upload:
24
+
25
+ `models/2026-09-14_07-43-34_prim_extruded_nr45_knn24_n2048_n6/best.pt`
26
+
27
+ This folder is infer + Gradio only. It is not the training repo.
config.yaml ADDED
@@ -0,0 +1,73 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Experiment settings for the occupancy MLP MVP.
2
+ # `device` is still detected at runtime in src/config.py (CUDA vs CPU).
3
+
4
+ # --- Model Architecture ---
5
+ hidden: 64 # Number of hidden units (channels) per layer in the OccupancyMLP
6
+ depth: 4 # Number of hidden linear layers in the MLP network
7
+ seed: 1 # Same seed as the previous extrude_nr1 train
8
+
9
+ # --- Path Configuration ---
10
+ data_dir: "E:/Work_stuff/scatteringNet/data" # Base root directory where all NPZ dataset files are stored
11
+
12
+ # --- Training Hyperparameters ---
13
+ epochs: 20 # Same length as the knn24-n2048 G5 clone (compare IoU to that best.pt)
14
+ lr: 0.001 # Initial learning rate
15
+ batch_size: 1024 # Points per optimizer step (GPU mini-batch)
16
+ optimizer: adam # adam | adamw | sgd
17
+ # BCE inside-class weight. auto = n_outside / n_inside on the train split.
18
+ # Omit or null = unweighted BCE (legacy / 18-09-31 clone). Head is unchanged.
19
+ pos_weight: auto
20
+
21
+ # --- Split ---
22
+ # Whole meshes: 80% train / 20% val of unique OBJs. All NPZs of one OBJ stay on one side.
23
+ # This split selects best.pt (not a locked holdout). Phase 2 Step 11 adds that.
24
+ val_fraction: 0.20
25
+
26
+ # --- Geometry ---
27
+ # none = xyz-only OccupancyMLP. surface = envelope (XYZ + face normal) + OccupancyEncoder.
28
+ # Occupancy labels stay in the NPZ. New trains write envelope_dim=6 on best.pt.
29
+ shape_encoder: surface # none / surface
30
+ n_surface: 2048 # Envelope samples on the joined OBJ (area-weighted face darts)
31
+ knn_k: 24 # Neighbors per query (not envelope count). Same head as knn24-n2048 inspect / G5 clone. Fresh train; do not resume that mix-75 best.pt.
32
+ # knn_local_dim: 64 # z_local width; omit to use latent_dim / hidden
33
+ # latent_dim: 64 # OccupancyEncoder z width; omit to use hidden
34
+
35
+ # --- Catalog ---
36
+ # npz_catalog wins over npz_glob. Each row is all meshes for that glob unless
37
+ # max_shapes is set (unique OBJs, sampled with seed; still max_files_per_shape NPZs).
38
+ # Varied/organic and combo* are omitted on purpose.
39
+ # Fallback glob is unused while npz_catalog is set.
40
+ # Smooth extruded_* is back (same glob as knn24-n2048 inspect). High-round
41
+ # extrude_* stay s0.08. Primitives and nr1 stay s0.15. nr3 is omitted.
42
+ # max_files_per_shape keeps lattice + one jitter.
43
+ npz_glob: "exports/dataset/*.npz"
44
+ npz_catalog:
45
+ - glob: "exports/dataset/cone_*.npz"
46
+ - glob: "exports/dataset/cube_*.npz"
47
+ - glob: "exports/dataset/cyl_*.npz"
48
+ - glob: "exports/dataset/gear_*.npz"
49
+ - glob: "exports/dataset/helix_*.npz"
50
+ - glob: "exports/dataset/pipe_*.npz"
51
+ - glob: "exports/dataset/platonic_*.npz"
52
+ - glob: "exports/dataset/prism_*.npz"
53
+ - glob: "exports/dataset/sphere_*.npz"
54
+ - glob: "exports/dataset/torus_*.npz"
55
+ - glob: "exports/dataset/extruded_*occupancy_s0.08*.npz"
56
+ - glob: "exports/dataset/extrude_*_nr1_*.npz"
57
+ max_shapes: 40
58
+ - glob: "exports/dataset/extrude_*_nr4_*occupancy_s0.08*.npz"
59
+ max_shapes: 90
60
+ - glob: "exports/dataset/extrude_*_nr5_*occupancy_s0.08*.npz"
61
+ max_files_per_shape: 2
62
+
63
+ # --- Run logs ---
64
+ # Suffix for runs/<YYYY-MM-DD_HH-MM-SS>_<name>/. Device is detected at runtime, not stored here.
65
+ run_name: prim_extruded_nr45_knn24_n2048_n6_pw
66
+
67
+ # --- Checkpoints ---
68
+ # Scalar used to decide models/<run_id>/best.pt (strict improve). No last.pt.
69
+ # The run snapshot (runs/<id>/config.yaml) records:
70
+ # total — planned epoch count (copied from epochs)
71
+ # checkpoint — epoch index stored in best.pt (runtime; not a project knob)
72
+ # val_iou = keep the epoch with the best inside overlap (not overall accuracy).
73
+ checkpoint_metric: val_iou
docs/inspect_checkpoint.yaml ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ # INSPECT checkpoint alias (single source of truth).
2
+ # Both the Three.js inspect viewer and Gradio default to this run when
3
+ # models/<id>/best.pt exists. Do not call this alias "best".
4
+ # 2026-09-16 area-only and the pos_weight run are documented A/B logs only.
5
+ run_id: "2026-09-14_07-43-34_prim_extruded_nr45_knn24_n2048_n6"
models/.gitkeep ADDED
File without changes
pyproject.toml ADDED
@@ -0,0 +1,22 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ [build-system]
2
+ requires = ["setuptools>=61"]
3
+ build-backend = "setuptools.build_meta"
4
+
5
+ [project]
6
+ name = "scatteringnet"
7
+ version = "0.1.0"
8
+ description = "Occupancy fill demo (Gradio Space)"
9
+ requires-python = ">=3.10"
10
+ dependencies = [
11
+ "numpy",
12
+ "pyyaml",
13
+ "trimesh>=4.0",
14
+ ]
15
+
16
+ [tool.setuptools]
17
+ package-dir = { "scatteringnet" = "src" }
18
+ packages = [
19
+ "scatteringnet",
20
+ "scatteringnet.geometry",
21
+ "scatteringnet.viewer",
22
+ ]
requirements.txt ADDED
@@ -0,0 +1,4 @@
 
 
 
 
 
1
+ numpy
2
+ pyyaml
3
+ trimesh>=4.0
4
+ torch
scatteringnet/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ """Occupancy fill: catalog training, encoder, and inspect helper."""
scatteringnet/config.py ADDED
@@ -0,0 +1,611 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Hybrid config loader for catalog occupancy training.
2
+
3
+ Static experiment knobs live in ``config.yaml``. ``device`` is resolved
4
+ here from CUDA availability. Training knobs (epochs, lr, batch_size,
5
+ optimizer, catalog, val split) are YAML-owned so ``train_multi_npz``
6
+ does not hardcode them.
7
+
8
+ This module does not read ``.env`` and does not open NPZ files.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import os
14
+ from dataclasses import dataclass
15
+ from pathlib import Path
16
+ from typing import Any, Mapping, TypedDict
17
+
18
+ import torch
19
+ import yaml
20
+
21
+ # Repo root: src/config.py → parents[1].
22
+ _REPO_ROOT = Path(__file__).resolve().parents[1]
23
+ _DEFAULT_YAML = _REPO_ROOT / "config.yaml"
24
+
25
+ _REQUIRED_YAML_KEYS = (
26
+ "hidden",
27
+ "depth",
28
+ "seed",
29
+ "data_dir",
30
+ "epochs",
31
+ "lr",
32
+ )
33
+
34
+
35
+ class YamlKnobs(TypedDict):
36
+ """Subset of OccupancyConfig that is stored in YAML (paths as strings)."""
37
+
38
+ hidden: int
39
+ depth: int
40
+ seed: int
41
+ data_dir: str
42
+ epochs: int
43
+ lr: float
44
+ val_fraction: float
45
+ latent_dim: int | None
46
+ npz_glob: str
47
+ npz_paths: tuple[str, ...]
48
+ npz_catalog: tuple[tuple[str, int | None], ...]
49
+ max_files_per_shape: int | None
50
+ run_name: str
51
+ checkpoint_metric: str
52
+ batch_size: int
53
+ optimizer: str
54
+ n_surface: int
55
+ knn_k: int
56
+ knn_local_dim: int | None
57
+ shape_encoder: str
58
+ # Explicit BCE pos_weight; None when omitted or when auto is set.
59
+ pos_weight: float | None
60
+ pos_weight_auto: bool
61
+
62
+
63
+ def get_device() -> torch.device:
64
+ """
65
+ CUDA when a GPU is visible; otherwise CPU.
66
+
67
+ Hugging Face CPU Spaces have no CUDA: infer still runs (slower).
68
+ ``SCATTERINGNET_DEVICE=cpu`` forces CPU even if a GPU exists.
69
+ ``SCATTERINGNET_DEVICE=cuda`` uses CUDA only when ``is_available()``;
70
+ otherwise it falls back to CPU (no crash).
71
+
72
+ Training scripts should still warn on a long catalog run on CPU.
73
+ """
74
+ forced = os.environ.get("SCATTERINGNET_DEVICE", "").strip().lower()
75
+ if forced == "cpu":
76
+ return torch.device("cpu")
77
+ if forced == "cuda":
78
+ if torch.cuda.is_available():
79
+ return torch.device("cuda")
80
+ return torch.device("cpu")
81
+ if torch.cuda.is_available():
82
+ return torch.device("cuda")
83
+ return torch.device("cpu")
84
+
85
+
86
+ def repo_root() -> Path:
87
+ """Git / project root (folder that contains ``src/`` and ``config.yaml``)."""
88
+ return _REPO_ROOT
89
+
90
+
91
+ def gpu_name(device: torch.device | None = None) -> str | None:
92
+ """
93
+ Human GPU name for the run snapshot (``None`` on CPU).
94
+
95
+ Uses ``cfg.device`` when given so a forced-CPU train does not stamp a
96
+ card that was not used.
97
+ """
98
+ dev = device if device is not None else get_device()
99
+ if dev.type != "cuda" or not torch.cuda.is_available():
100
+ return None
101
+ index = 0 if dev.index is None else int(dev.index)
102
+ if index < 0 or index >= torch.cuda.device_count():
103
+ return None
104
+ name = str(torch.cuda.get_device_name(index)).strip()
105
+ return name or None
106
+
107
+
108
+ def _as_positive_int(name: str, value: Any) -> int:
109
+ """YAML may yield int or (rarely) str; occupancy dims must be int >= 1."""
110
+ try:
111
+ parsed = int(value)
112
+ except (TypeError, ValueError) as exc:
113
+ raise ValueError(f"{name} must be an integer, got {value!r}") from exc
114
+ if parsed < 1:
115
+ raise ValueError(f"{name} must be >= 1, got {parsed}")
116
+ return parsed
117
+
118
+
119
+ def _as_int_in_range(name: str, value: Any, lo: int, hi: int) -> int:
120
+ """Inclusive integer range (YAML ``knn_k`` is 0–4096)."""
121
+ try:
122
+ parsed = int(value)
123
+ except (TypeError, ValueError) as exc:
124
+ raise ValueError(f"{name} must be an integer, got {value!r}") from exc
125
+ if parsed < lo or parsed > hi:
126
+ raise ValueError(f"{name} must be in [{lo}, {hi}], got {parsed}")
127
+ return parsed
128
+
129
+
130
+ def _as_positive_float(name: str, value: Any) -> float:
131
+ """Learning-rate style knobs must be a finite float > 0."""
132
+ try:
133
+ parsed = float(value)
134
+ except (TypeError, ValueError) as exc:
135
+ raise ValueError(f"{name} must be a float, got {value!r}") from exc
136
+ if parsed <= 0.0 or parsed != parsed:
137
+ raise ValueError(f"{name} must be > 0, got {parsed}")
138
+ return parsed
139
+
140
+
141
+ def _as_pos_weight_pair(raw: Mapping[str, Any]) -> tuple[float | None, bool]:
142
+ """
143
+ YAML ``pos_weight``: omit / null → unweighted BCE.
144
+
145
+ ``auto`` → compute n_outside / n_inside on the train split at train time.
146
+ A finite float > 0 is used as-is (1.0 is unweighted).
147
+ """
148
+ if "pos_weight" not in raw:
149
+ return None, False
150
+ value = raw["pos_weight"]
151
+ if value is None or value is False:
152
+ return None, False
153
+ if isinstance(value, str):
154
+ text = value.strip().lower()
155
+ if text in ("", "none", "off", "false"):
156
+ return None, False
157
+ if text == "auto":
158
+ return None, True
159
+ parsed = _as_positive_float("pos_weight", value)
160
+ return parsed, False
161
+
162
+
163
+ def _as_open_unit_interval(name: str, value: Any) -> float:
164
+ """Hold-out fractions must be in (0, 1) so both splits are non-empty."""
165
+ try:
166
+ parsed = float(value)
167
+ except (TypeError, ValueError) as exc:
168
+ raise ValueError(f"{name} must be a float, got {value!r}") from exc
169
+ if parsed != parsed or parsed <= 0.0 or parsed >= 1.0:
170
+ raise ValueError(f"{name} must be in (0, 1), got {parsed}")
171
+ return parsed
172
+
173
+
174
+ def _as_nonempty_path_string(name: str, value: Any) -> str:
175
+ if value is None or (isinstance(value, str) and not value.strip()):
176
+ raise ValueError(f"{name} must be a non-empty path string in config.yaml")
177
+ return str(value).strip()
178
+
179
+
180
+ def _as_data_dir_string(value: Any) -> str:
181
+ return _as_nonempty_path_string("data_dir", value)
182
+
183
+
184
+ def _as_optional_positive_int(name: str, value: Any) -> int | None:
185
+ """YAML null → unlimited catalog cap; otherwise int >= 1."""
186
+ if value is None:
187
+ return None
188
+ return _as_positive_int(name, value)
189
+
190
+
191
+ def _as_run_name(value: Any) -> str:
192
+ """Optional YAML suffix for ``runs/<timestamp>_<name>/``; empty → ``run``."""
193
+ if value is None:
194
+ return "run"
195
+ text = str(value).strip()
196
+ return text if text else "run"
197
+
198
+
199
+ _CHECKPOINT_METRIC_ALIASES = {
200
+ "test_acc": "val_acc",
201
+ "test_iou": "val_iou",
202
+ }
203
+
204
+
205
+ def _as_checkpoint_metric(value: Any) -> str:
206
+ """Name of the scalar used to decide ``best.pt`` (strict improve)."""
207
+ if value is None or (isinstance(value, str) and not value.strip()):
208
+ raise ValueError("checkpoint_metric must be a non-empty string")
209
+ name = str(value).strip()
210
+ return _CHECKPOINT_METRIC_ALIASES.get(name, name)
211
+
212
+
213
+ def _as_val_fraction(raw: Mapping[str, Any]) -> float:
214
+ """Prefer ``val_fraction``; accept legacy ``test_fraction``."""
215
+ if "val_fraction" in raw:
216
+ return _as_open_unit_interval("val_fraction", raw["val_fraction"])
217
+ if "test_fraction" in raw:
218
+ return _as_open_unit_interval("test_fraction", raw["test_fraction"])
219
+ raise ValueError("Config YAML missing keys: val_fraction")
220
+
221
+
222
+ def _as_optional_latent_dim(raw: Mapping[str, Any]) -> int | None:
223
+ """YAML omit / null → use ``hidden`` at train time."""
224
+ if "latent_dim" not in raw or raw["latent_dim"] is None:
225
+ return None
226
+ return _as_positive_int("latent_dim", raw["latent_dim"])
227
+
228
+
229
+ # Names accepted in config.yaml ``optimizer``. Used by train_multi_npz.
230
+ _ALLOWED_OPTIMIZERS = ("adam", "adamw", "sgd")
231
+ # ``none`` = OccupancyMLP (xyz only). ``surface`` = envelope encoder.
232
+ _ALLOWED_SHAPE_ENCODERS = ("none", "surface")
233
+
234
+
235
+ def _as_optimizer(value: Any) -> str:
236
+ """Optimizer family for multi-NPZ train; default Adam."""
237
+ if value is None or (isinstance(value, str) and not str(value).strip()):
238
+ return "adam"
239
+ name = str(value).strip().lower()
240
+ if name not in _ALLOWED_OPTIMIZERS:
241
+ allowed = ", ".join(_ALLOWED_OPTIMIZERS)
242
+ raise ValueError(f"optimizer must be one of {allowed}, got {value!r}")
243
+ return name
244
+
245
+
246
+ def _as_shape_encoder(value: Any) -> str:
247
+ """Occupancy head family; default xyz-only so older YAML still loads."""
248
+ if value is None or (isinstance(value, str) and not str(value).strip()):
249
+ return "none"
250
+ name = str(value).strip().lower()
251
+ if name not in _ALLOWED_SHAPE_ENCODERS:
252
+ allowed = ", ".join(_ALLOWED_SHAPE_ENCODERS)
253
+ raise ValueError(f"shape_encoder must be one of {allowed}, got {value!r}")
254
+ return name
255
+
256
+
257
+ def _as_npz_paths(value: Any) -> tuple[str, ...]:
258
+ """Explicit NPZ list relative to data_dir (empty → use glob)."""
259
+ if value is None:
260
+ return ()
261
+ if isinstance(value, str):
262
+ item = value.strip()
263
+ return (item,) if item else ()
264
+ if not isinstance(value, list):
265
+ raise ValueError(f"npz_paths must be a list of strings or null, got {type(value).__name__}")
266
+ out: list[str] = []
267
+ for i, raw in enumerate(value):
268
+ text = _as_nonempty_path_string(f"npz_paths[{i}]", raw)
269
+ out.append(text)
270
+ return tuple(out)
271
+
272
+
273
+ def _as_npz_catalog(value: Any) -> tuple[tuple[str, int | None], ...]:
274
+ """Union of globs; optional ``max_shapes`` is unique meshes per glob."""
275
+ if value is None:
276
+ return ()
277
+ if not isinstance(value, list):
278
+ raise ValueError(
279
+ f"npz_catalog must be a list or null, got {type(value).__name__}"
280
+ )
281
+ out: list[tuple[str, int | None]] = []
282
+ for i, raw in enumerate(value):
283
+ if isinstance(raw, str):
284
+ glob_s = _as_nonempty_path_string(f"npz_catalog[{i}]", raw)
285
+ out.append((glob_s, None))
286
+ continue
287
+ if not isinstance(raw, Mapping):
288
+ raise ValueError(
289
+ f"npz_catalog[{i}] must be a glob string or mapping, "
290
+ f"got {type(raw).__name__}"
291
+ )
292
+ if "glob" not in raw:
293
+ raise ValueError(f"npz_catalog[{i}] missing glob")
294
+ glob_s = _as_nonempty_path_string(f"npz_catalog[{i}].glob", raw["glob"])
295
+ max_shapes: int | None = None
296
+ if "max_shapes" in raw and raw["max_shapes"] is not None:
297
+ max_shapes = _as_positive_int(
298
+ f"npz_catalog[{i}].max_shapes", raw["max_shapes"]
299
+ )
300
+ out.append((glob_s, max_shapes))
301
+ return tuple(out)
302
+
303
+
304
+ def as_repo_relative(path: Path | str, *, root: Path | None = None) -> str:
305
+ """
306
+ POSIX string relative to the git repo when ``path`` is inside it.
307
+
308
+ Already-relative inputs are returned as POSIX. Absolute paths on another
309
+ drive (the dataset disk) cannot be repo-relative and stay absolute POSIX.
310
+ """
311
+ text = str(path).strip()
312
+ if not text:
313
+ return text
314
+ parsed = Path(text)
315
+ if not parsed.is_absolute():
316
+ return parsed.as_posix()
317
+ base = (root or _REPO_ROOT).resolve()
318
+ try:
319
+ return parsed.resolve().relative_to(base).as_posix()
320
+ except ValueError:
321
+ return parsed.resolve().as_posix()
322
+
323
+
324
+ def as_data_relative(path: Path | str, data_dir: Path | str) -> str:
325
+ """
326
+ POSIX string relative to ``data_dir`` (``exports/...``, not ``E:/...``).
327
+
328
+ Already-relative inputs are returned as POSIX. Paths outside ``data_dir``
329
+ (unit-test temp trees) fall back to absolute POSIX.
330
+ """
331
+ text = str(path).strip()
332
+ if not text:
333
+ return text
334
+ parsed = Path(text)
335
+ if not parsed.is_absolute():
336
+ return parsed.as_posix()
337
+ root = Path(data_dir).expanduser().resolve()
338
+ resolved = parsed.expanduser().resolve()
339
+ try:
340
+ return resolved.relative_to(root).as_posix()
341
+ except ValueError:
342
+ return resolved.as_posix()
343
+
344
+
345
+ def load_yaml_knobs(path: Path) -> YamlKnobs:
346
+ """
347
+ Read YAML settings. Does not check that data_dir exists on disk.
348
+
349
+ Parameters
350
+ ----------
351
+ path:
352
+ Path to ``config.yaml``.
353
+
354
+ Returns
355
+ -------
356
+ YamlKnobs
357
+ Typed dict of experiment knobs (paths still strings).
358
+ """
359
+ if not path.is_file():
360
+ raise FileNotFoundError(f"Config YAML not found: {path}")
361
+ raw = yaml.safe_load(path.read_text(encoding="utf-8"))
362
+ if not isinstance(raw, Mapping):
363
+ raise ValueError(f"Config YAML must be a mapping, got {type(raw).__name__}")
364
+ missing = [k for k in _REQUIRED_YAML_KEYS if k not in raw]
365
+ if missing:
366
+ raise ValueError(f"Config YAML missing keys: {', '.join(missing)}")
367
+ pos_weight, pos_weight_auto = _as_pos_weight_pair(raw)
368
+ return {
369
+ "hidden": _as_positive_int("hidden", raw["hidden"]),
370
+ "depth": _as_positive_int("depth", raw["depth"]),
371
+ "seed": _as_positive_int("seed", raw["seed"]),
372
+ "data_dir": _as_data_dir_string(raw["data_dir"]),
373
+ "epochs": _as_positive_int("epochs", raw["epochs"]),
374
+ "lr": _as_positive_float("lr", raw["lr"]),
375
+ "val_fraction": _as_val_fraction(raw),
376
+ "latent_dim": _as_optional_latent_dim(raw),
377
+ "npz_glob": (
378
+ _as_nonempty_path_string("npz_glob", raw["npz_glob"])
379
+ if "npz_glob" in raw
380
+ else "exports/dataset/*.npz"
381
+ ),
382
+ "npz_paths": _as_npz_paths(raw.get("npz_paths")),
383
+ "npz_catalog": _as_npz_catalog(raw.get("npz_catalog")),
384
+ "max_files_per_shape": (
385
+ _as_optional_positive_int("max_files_per_shape", raw["max_files_per_shape"])
386
+ if "max_files_per_shape" in raw
387
+ else 2
388
+ ),
389
+ "run_name": (
390
+ _as_run_name(raw["run_name"]) if "run_name" in raw else "run"
391
+ ),
392
+ "checkpoint_metric": (
393
+ _as_checkpoint_metric(raw["checkpoint_metric"])
394
+ if "checkpoint_metric" in raw
395
+ else "val_acc"
396
+ ),
397
+ "batch_size": (
398
+ _as_positive_int("batch_size", raw["batch_size"])
399
+ if "batch_size" in raw
400
+ else 1024
401
+ ),
402
+ "optimizer": (
403
+ _as_optimizer(raw["optimizer"]) if "optimizer" in raw else "adam"
404
+ ),
405
+ "n_surface": (
406
+ _as_positive_int("n_surface", raw["n_surface"])
407
+ if "n_surface" in raw
408
+ else 1024
409
+ ),
410
+ "knn_k": (
411
+ _as_int_in_range("knn_k", raw["knn_k"], 0, 4096)
412
+ if "knn_k" in raw
413
+ else 0
414
+ ),
415
+ "knn_local_dim": (
416
+ _as_positive_int("knn_local_dim", raw["knn_local_dim"])
417
+ if "knn_local_dim" in raw
418
+ else None
419
+ ),
420
+ "shape_encoder": (
421
+ _as_shape_encoder(raw["shape_encoder"])
422
+ if "shape_encoder" in raw
423
+ else "none"
424
+ ),
425
+ "pos_weight": pos_weight,
426
+ "pos_weight_auto": pos_weight_auto,
427
+ }
428
+
429
+
430
+ def _warn_missing_data_dir(data_dir: Path, yaml_path: Path) -> None:
431
+ """Print a terminal hint so the user can fix config.yaml (no .env involved)."""
432
+ print(
433
+ "\n"
434
+ "Dataset folder not found.\n"
435
+ f" Looked for: {data_dir}\n"
436
+ "\n"
437
+ "Update `data_dir` in config.yaml to the folder that contains "
438
+ "`exports/` and `meshes/`.\n"
439
+ f" Config file: {yaml_path}\n"
440
+ )
441
+
442
+
443
+ def require_data_dir(data_dir: Path, *, yaml_path: Path) -> None:
444
+ """Validate the dataset root before training or NPZ loading."""
445
+ if data_dir.is_dir():
446
+ return
447
+ _warn_missing_data_dir(data_dir, yaml_path)
448
+ raise FileNotFoundError(
449
+ f"Dataset directory does not exist: {data_dir}. "
450
+ f"Set data_dir in {yaml_path}."
451
+ )
452
+
453
+
454
+ @dataclass(frozen=True)
455
+ class OccupancyConfig:
456
+ """Resolved experiment settings from YAML plus detected device."""
457
+
458
+ data_dir: Path
459
+ device: torch.device
460
+ hidden: int
461
+ depth: int
462
+ seed: int
463
+ epochs: int
464
+ lr: float
465
+ # Fraction of catalog **meshes** held out as val (selection split, not a locked test).
466
+ val_fraction: float
467
+ # Catalog knobs (optional in YAML; omitted keys keep these defaults).
468
+ npz_glob: str = "exports/dataset/*.npz"
469
+ npz_paths: tuple[str, ...] = ()
470
+ # ``(glob, max_shapes)`` rows. Empty → use ``npz_glob``. ``max_shapes``
471
+ # None keeps every mesh that glob hits (after ``max_files_per_shape``).
472
+ npz_catalog: tuple[tuple[str, int | None], ...] = ()
473
+ max_files_per_shape: int | None = 2
474
+ # Suffix for runs/<timestamp>_<name>/ (device stays runtime-only).
475
+ run_name: str = "run"
476
+ # Which logged scalar selects best.pt (strict improve).
477
+ checkpoint_metric: str = "val_acc"
478
+ # Encoder latent width; None → use ``hidden`` at train / infer time.
479
+ latent_dim: int | None = None
480
+ # Mini-batch size and optimizer family (train_multi_npz).
481
+ batch_size: int = 1024
482
+ optimizer: str = "adam"
483
+ # Envelope sample count (YAML). Used when ``shape_encoder`` is ``surface``.
484
+ n_surface: int = 1024
485
+ # 0 = global envelope z only. >0 = that many nearest envelope dots per query.
486
+ knn_k: int = 0
487
+ # Width of z_local; None → same as occupancy latent_dim / hidden.
488
+ knn_local_dim: int | None = None
489
+ # ``none`` keeps OccupancyMLP; ``surface`` uses the envelope PointNet.
490
+ shape_encoder: str = "none"
491
+ # BCE inside-class weight. None + auto=False = unweighted (legacy).
492
+ pos_weight: float | None = None
493
+ pos_weight_auto: bool = False
494
+
495
+
496
+ def load_config(
497
+ yaml_path: Path | None = None,
498
+ *,
499
+ require_existing_data_dir: bool = True,
500
+ ) -> OccupancyConfig:
501
+ """
502
+ Compose OccupancyConfig from ``config.yaml`` (not from ``.env``).
503
+
504
+ When ``require_existing_data_dir`` is True (default), a missing folder
505
+ prints a short instruction and then raises FileNotFoundError — used for
506
+ training and data loading.
507
+
508
+ Parameters
509
+ ----------
510
+ yaml_path:
511
+ Config file; default is repo-root ``config.yaml``.
512
+ require_existing_data_dir:
513
+ If True, refuse to return a config whose ``data_dir`` is missing.
514
+
515
+ Returns
516
+ -------
517
+ OccupancyConfig
518
+ YAML knobs plus detected ``device``.
519
+ """
520
+ cfg_path = yaml_path or _DEFAULT_YAML
521
+ knobs = load_yaml_knobs(cfg_path)
522
+ data_dir = Path(knobs["data_dir"])
523
+ if require_existing_data_dir:
524
+ require_data_dir(data_dir, yaml_path=cfg_path)
525
+ return OccupancyConfig(
526
+ data_dir=data_dir,
527
+ device=get_device(),
528
+ hidden=knobs["hidden"],
529
+ depth=knobs["depth"],
530
+ seed=knobs["seed"],
531
+ epochs=knobs["epochs"],
532
+ lr=knobs["lr"],
533
+ val_fraction=knobs["val_fraction"],
534
+ latent_dim=knobs["latent_dim"],
535
+ npz_glob=knobs["npz_glob"],
536
+ npz_paths=knobs["npz_paths"],
537
+ npz_catalog=knobs["npz_catalog"],
538
+ max_files_per_shape=knobs["max_files_per_shape"],
539
+ run_name=knobs["run_name"],
540
+ checkpoint_metric=knobs["checkpoint_metric"],
541
+ batch_size=knobs["batch_size"],
542
+ optimizer=knobs["optimizer"],
543
+ n_surface=knobs["n_surface"],
544
+ knn_k=knobs["knn_k"],
545
+ knn_local_dim=knobs["knn_local_dim"],
546
+ shape_encoder=knobs["shape_encoder"],
547
+ pos_weight=knobs["pos_weight"],
548
+ pos_weight_auto=knobs["pos_weight_auto"],
549
+ )
550
+
551
+
552
+ def encoder_latent_dim(cfg: OccupancyConfig) -> int:
553
+ """OccupancyEncoder ``z`` width: YAML ``latent_dim`` or ``hidden``."""
554
+ if cfg.latent_dim is None:
555
+ return int(cfg.hidden)
556
+ return int(cfg.latent_dim)
557
+
558
+
559
+ def encoder_knn_local_dim(cfg: OccupancyConfig) -> int:
560
+ """Local envelope code width: YAML ``knn_local_dim`` or global latent."""
561
+ if cfg.knn_local_dim is None:
562
+ return encoder_latent_dim(cfg)
563
+ return int(cfg.knn_local_dim)
564
+
565
+
566
+ def format_config(cfg: OccupancyConfig) -> str:
567
+ """
568
+ Pretty-print for CLI smoke checks.
569
+
570
+ Parameters
571
+ ----------
572
+ cfg:
573
+ Resolved config.
574
+
575
+ Returns
576
+ -------
577
+ str
578
+ Multi-line ``OccupancyConfig(...)`` dump.
579
+ """
580
+ return (
581
+ f"OccupancyConfig(\n"
582
+ f" data_dir={cfg.data_dir}\n"
583
+ f" device={cfg.device}\n"
584
+ f" gpu={gpu_name(cfg.device)}\n"
585
+ f" hidden={cfg.hidden}\n"
586
+ f" depth={cfg.depth}\n"
587
+ f" seed={cfg.seed}\n"
588
+ f" epochs={cfg.epochs}\n"
589
+ f" lr={cfg.lr}\n"
590
+ f" val_fraction={cfg.val_fraction}\n"
591
+ f" latent_dim={cfg.latent_dim}\n"
592
+ f" npz_glob={cfg.npz_glob}\n"
593
+ f" npz_paths={list(cfg.npz_paths)}\n"
594
+ f" npz_catalog={list(cfg.npz_catalog)}\n"
595
+ f" max_files_per_shape={cfg.max_files_per_shape}\n"
596
+ f" run_name={cfg.run_name}\n"
597
+ f" checkpoint_metric={cfg.checkpoint_metric}\n"
598
+ f" batch_size={cfg.batch_size}\n"
599
+ f" optimizer={cfg.optimizer}\n"
600
+ f" n_surface={cfg.n_surface}\n"
601
+ f" knn_k={cfg.knn_k}\n"
602
+ f" knn_local_dim={cfg.knn_local_dim}\n"
603
+ f" shape_encoder={cfg.shape_encoder}\n"
604
+ f" pos_weight={cfg.pos_weight}\n"
605
+ f" pos_weight_auto={cfg.pos_weight_auto}\n"
606
+ f")"
607
+ )
608
+
609
+
610
+ if __name__ == "__main__":
611
+ print(format_config(load_config()))
scatteringnet/data_npz.py ADDED
@@ -0,0 +1,460 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Load occupancy query points and labels from scatter NPZ files.
2
+
3
+ :func:`load_points_labels` reads **one** file. A catalog resolver lists
4
+ many NPZs (glob or explicit paths) without training.
5
+
6
+ ``load_points_labels`` still returns only ``points`` and ``labels``.
7
+ ``load_points_labels_mesh`` also resolves ``mesh_path`` against ``data_dir``.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ import glob as globlib
13
+ import random
14
+ from pathlib import Path
15
+ from typing import Sequence
16
+
17
+ import numpy as np
18
+ from numpy.typing import NDArray
19
+
20
+ # Labels are converted here (not deferred to the Dataset) so every caller gets
21
+ # the same dtypes: float32 XYZ and float32 {0, 1} occupancy.
22
+ PointsArray = NDArray[np.float32]
23
+ LabelsArray = NDArray[np.float32]
24
+
25
+
26
+ def load_points_labels(path: Path) -> tuple[PointsArray, LabelsArray]:
27
+ """
28
+ Read query coordinates and inside/outside labels from one NPZ file.
29
+
30
+ Parameters
31
+ ----------
32
+ path:
33
+ Path to a ``.npz`` with arrays ``points`` ``(N, 3)`` and
34
+ ``labels`` ``(N,)`` (typically uint8 0/1).
35
+
36
+ Returns
37
+ -------
38
+ points:
39
+ ``float32`` array of shape ``(N, 3)``.
40
+ labels:
41
+ ``float32`` array of shape ``(N,)`` with values in ``{0.0, 1.0}``
42
+ (0 = outside, 1 = inside).
43
+ """
44
+ npz_path = Path(path)
45
+ if not npz_path.is_file():
46
+ raise FileNotFoundError(f"NPZ not found: {npz_path}")
47
+
48
+ # allow_pickle=False: we only need numeric arrays, not object payloads.
49
+ with np.load(npz_path, allow_pickle=False) as raw:
50
+ files = set(raw.files)
51
+ if "points" not in files or "labels" not in files:
52
+ raise KeyError(
53
+ f"NPZ must contain 'points' and 'labels', got {sorted(files)} "
54
+ f"in {npz_path}"
55
+ )
56
+ points = np.asarray(raw["points"])
57
+ labels = np.asarray(raw["labels"])
58
+
59
+ if points.ndim != 2 or points.shape[1] != 3:
60
+ raise ValueError(
61
+ f"points must have shape (N, 3), got {tuple(points.shape)} in {npz_path}"
62
+ )
63
+ n = int(points.shape[0])
64
+ if labels.shape != (n,):
65
+ raise ValueError(
66
+ f"labels must have shape (N,) with N={n}, got {tuple(labels.shape)} "
67
+ f"in {npz_path}"
68
+ )
69
+
70
+ points_f32 = np.asarray(points, dtype=np.float32)
71
+ labels_f32 = np.asarray(labels, dtype=np.float32)
72
+ unique = np.unique(labels_f32)
73
+ if not np.all((unique == 0.0) | (unique == 1.0)):
74
+ raise ValueError(
75
+ f"labels must be in {{0, 1}}, got unique={unique.tolist()} in {npz_path}"
76
+ )
77
+ return points_f32, labels_f32
78
+
79
+
80
+ def count_npz_points(path: Path | str) -> int:
81
+ """
82
+ Query count in one NPZ without keeping the arrays.
83
+
84
+ Catalog construct uses this so ``len(part)`` / ``n_points`` do not
85
+ require loading every lattice into RAM.
86
+ """
87
+ npz_path = Path(path)
88
+ if not npz_path.is_file():
89
+ raise FileNotFoundError(f"NPZ not found: {npz_path}")
90
+ with np.load(npz_path, allow_pickle=False) as raw:
91
+ files = set(raw.files)
92
+ if "points" not in files or "labels" not in files:
93
+ raise KeyError(
94
+ f"NPZ must contain 'points' and 'labels', got {sorted(files)} "
95
+ f"in {npz_path}"
96
+ )
97
+ n = int(np.asarray(raw["points"]).shape[0])
98
+ n_y = int(np.asarray(raw["labels"]).shape[0])
99
+ if n_y != n:
100
+ raise ValueError(
101
+ f"labels must have shape (N,) with N={n}, got N={n_y} in {npz_path}"
102
+ )
103
+ return n
104
+
105
+
106
+ def count_npz_labels(path: Path | str) -> tuple[int, int]:
107
+ """
108
+ Outside / inside counts in one NPZ without keeping the point cloud.
109
+
110
+ Used for ``pos_weight: auto`` (n_outside / n_inside on the train split).
111
+ """
112
+ npz_path = Path(path)
113
+ if not npz_path.is_file():
114
+ raise FileNotFoundError(f"NPZ not found: {npz_path}")
115
+ with np.load(npz_path, allow_pickle=False) as raw:
116
+ if "labels" not in set(raw.files):
117
+ raise KeyError(f"NPZ must contain 'labels', got {sorted(raw.files)} in {npz_path}")
118
+ labels = np.asarray(raw["labels"]).reshape(-1)
119
+ # Labels are 0/1 occupancy; nonzero is inside.
120
+ n_in = int(np.count_nonzero(labels))
121
+ n_out = int(labels.size) - n_in
122
+ return n_out, n_in
123
+
124
+
125
+ def read_npz_mesh_path(path: Path | str) -> str:
126
+ """
127
+ Read the stored ``mesh_path`` string from one occupancy NPZ.
128
+
129
+ Step 2 writes a ``data_dir``-relative POSIX path (for example
130
+ ``meshes/Primitives/Sphere/sphere_r0p5_sa16_sh16.obj``).
131
+
132
+ Parameters
133
+ ----------
134
+ path:
135
+ Occupancy ``.npz`` that contains ``mesh_path``.
136
+
137
+ Returns
138
+ -------
139
+ str
140
+ Stored path string (relative or absolute). Not resolved here.
141
+ """
142
+ npz_path = Path(path)
143
+ if not npz_path.is_file():
144
+ raise FileNotFoundError(f"NPZ not found: {npz_path}")
145
+ # allow_pickle=True: some exports store a 0-d string / object array.
146
+ with np.load(npz_path, allow_pickle=True) as raw:
147
+ if "mesh_path" not in raw.files:
148
+ raise KeyError(f"NPZ has no 'mesh_path' in {npz_path}")
149
+ stored = str(np.asarray(raw["mesh_path"]).item()).strip()
150
+ if not stored:
151
+ raise ValueError(f"mesh_path is empty in {npz_path}")
152
+ return stored
153
+
154
+
155
+ def resolve_mesh_path(stored: str, data_dir: Path | str) -> Path:
156
+ """
157
+ Resolve a stored ``mesh_path`` against ``data_dir``.
158
+
159
+ Relative entries are joined to ``data_dir``. Absolute entries are
160
+ used as-is. Missing files raise ``FileNotFoundError``.
161
+
162
+ Parameters
163
+ ----------
164
+ stored:
165
+ Value from :func:`read_npz_mesh_path`.
166
+ data_dir:
167
+ Dataset root (``config.yaml`` ``data_dir``).
168
+
169
+ Returns
170
+ -------
171
+ Path
172
+ Existing resolved mesh file.
173
+ """
174
+ text = str(stored).strip()
175
+ if not text:
176
+ raise ValueError("mesh_path is empty")
177
+ item = Path(text)
178
+ root = Path(data_dir)
179
+ resolved = item if item.is_absolute() else (root / item)
180
+ resolved = resolved.resolve()
181
+ if not resolved.is_file():
182
+ raise FileNotFoundError(f"mesh not found: {resolved} (stored={text!r})")
183
+ return resolved
184
+
185
+
186
+ def load_points_labels_mesh(
187
+ path: Path | str,
188
+ data_dir: Path | str,
189
+ ) -> tuple[PointsArray, LabelsArray, Path]:
190
+ """
191
+ Read occupancy arrays and resolve the source OBJ.
192
+
193
+ Keeps :func:`load_points_labels` unchanged (xyz + labels only).
194
+
195
+ Parameters
196
+ ----------
197
+ path:
198
+ Occupancy ``.npz`` with ``points``, ``labels``, and ``mesh_path``.
199
+ data_dir:
200
+ Root used to resolve a relative ``mesh_path``.
201
+
202
+ Returns
203
+ -------
204
+ points, labels, mesh_path:
205
+ Same arrays as :func:`load_points_labels`, plus the existing OBJ.
206
+ """
207
+ npz_path = Path(path)
208
+ points, labels = load_points_labels(npz_path)
209
+ stored = read_npz_mesh_path(npz_path)
210
+ mesh_path = resolve_mesh_path(stored, data_dir)
211
+ return points, labels, mesh_path
212
+
213
+
214
+ def _is_combo_npz(path: Path) -> bool:
215
+ """True when the filename looks like a combo dump (excluded from the catalog)."""
216
+ return "combo" in path.name.lower()
217
+
218
+
219
+ def shape_key(path: Path) -> str:
220
+ """
221
+ Group NPZs that belong to the same mesh.
222
+
223
+ Uses the stem before ``__`` (dataset_builder tag), else the full stem.
224
+
225
+ Parameters
226
+ ----------
227
+ path:
228
+ NPZ path.
229
+
230
+ Returns
231
+ -------
232
+ str
233
+ Stable key for ``max_files_per_shape``.
234
+ """
235
+ stem = Path(path).stem
236
+ if "__" in stem:
237
+ return stem.split("__", 1)[0]
238
+ return stem
239
+
240
+
241
+ def _cap_per_shape(
242
+ paths: Sequence[Path],
243
+ max_files_per_shape: int | None,
244
+ ) -> list[Path]:
245
+ """Keep at most ``max_files_per_shape`` files per :func:`shape_key` (sorted order)."""
246
+ if max_files_per_shape is None:
247
+ return list(paths)
248
+ if max_files_per_shape < 1:
249
+ raise ValueError(f"max_files_per_shape must be >= 1 or None, got {max_files_per_shape}")
250
+ counts: dict[str, int] = {}
251
+ out: list[Path] = []
252
+ for path in paths:
253
+ key = shape_key(path)
254
+ taken = counts.get(key, 0)
255
+ if taken >= max_files_per_shape:
256
+ continue
257
+ counts[key] = taken + 1
258
+ out.append(path)
259
+ return out
260
+
261
+
262
+ def _glob_npz(root: Path, pattern: str) -> list[Path]:
263
+ """Match ``pattern`` under ``root`` (``*`` / ``**``)."""
264
+ full = str(root / pattern)
265
+ recursive = "**" in pattern.replace("\\", "/")
266
+ found = globlib.glob(full, recursive=recursive)
267
+ return [Path(p).resolve() for p in found if Path(p).is_file()]
268
+
269
+
270
+ def _is_parameterized_stem(path: Path) -> bool:
271
+ """
272
+ Maya catalog names are ``family_param_...``. Varied one-off stems
273
+ (``Cone.obj`` → ``Cone__occupancy.npz``) have no ``_`` in the shape key
274
+ and must not ride along when Windows glob is case-insensitive.
275
+ """
276
+ return "_" in shape_key(path)
277
+
278
+
279
+ def _subsample_shapes(
280
+ paths: Sequence[Path],
281
+ max_shapes: int,
282
+ seed: int,
283
+ ) -> list[Path]:
284
+ """Keep NPZs for at most ``max_shapes`` unique :func:`shape_key` values."""
285
+ if max_shapes < 1:
286
+ raise ValueError(f"max_shapes must be >= 1, got {max_shapes}")
287
+ keys: list[str] = []
288
+ seen: set[str] = set()
289
+ for path in paths:
290
+ key = shape_key(path)
291
+ if key in seen:
292
+ continue
293
+ seen.add(key)
294
+ keys.append(key)
295
+ if max_shapes >= len(keys):
296
+ return list(paths)
297
+ # Sort then sample so the same seed always picks the same meshes.
298
+ chosen_keys = set(random.Random(int(seed)).sample(sorted(keys), max_shapes))
299
+ return [path for path in paths if shape_key(path) in chosen_keys]
300
+
301
+
302
+ def resolve_npz_catalog(
303
+ data_dir: Path | str,
304
+ *,
305
+ npz_glob: str = "exports/dataset/*.npz",
306
+ npz_paths: Sequence[str | Path] | None = None,
307
+ npz_catalog: Sequence[tuple[str, int | None]] | None = None,
308
+ max_files_per_shape: int | None = 2,
309
+ exclude_combo: bool = True,
310
+ seed: int = 1,
311
+ ) -> list[Path]:
312
+ """
313
+ Resolve occupancy NPZ paths under ``data_dir`` (no point loading).
314
+
315
+ Priority: explicit ``npz_paths``, else ``npz_catalog`` (union of globs),
316
+ else ``npz_glob``. Relative entries are joined to ``data_dir``.
317
+ Missing files in ``npz_paths`` raise ``FileNotFoundError``.
318
+
319
+ Parameters
320
+ ----------
321
+ data_dir:
322
+ Dataset root (``config.yaml`` ``data_dir``).
323
+ npz_glob:
324
+ Single glob relative to ``data_dir`` (``*`` and ``**`` allowed).
325
+ npz_paths:
326
+ Explicit relative or absolute NPZ paths. Empty / None → use glob(s).
327
+ npz_catalog:
328
+ ``(glob, max_shapes)`` rows. ``max_shapes`` is unique meshes after
329
+ the per-shape file cap; ``None`` keeps every mesh the glob hits.
330
+ max_files_per_shape:
331
+ Cap per :func:`shape_key` after sort. ``None`` = no cap.
332
+ exclude_combo:
333
+ Drop filenames containing ``combo``.
334
+ seed:
335
+ RNG for ``max_shapes`` subsampling (YAML ``seed``).
336
+
337
+ Returns
338
+ -------
339
+ list[Path]
340
+ Sorted existing ``.npz`` files.
341
+ """
342
+ root = Path(data_dir)
343
+ chosen: list[Path]
344
+ if npz_paths:
345
+ chosen = []
346
+ for raw in npz_paths:
347
+ item = Path(raw)
348
+ resolved = item if item.is_absolute() else (root / item)
349
+ if not resolved.is_file():
350
+ raise FileNotFoundError(f"NPZ not found: {resolved}")
351
+ chosen.append(resolved.resolve())
352
+ elif npz_catalog:
353
+ # Union in YAML order. Same file from two globs is kept once.
354
+ seen: set[Path] = set()
355
+ chosen = []
356
+ for pattern, max_shapes in npz_catalog:
357
+ hit = _glob_npz(root, str(pattern))
358
+ if exclude_combo:
359
+ hit = [p for p in hit if not _is_combo_npz(p)]
360
+ hit = [
361
+ p
362
+ for p in hit
363
+ if p.suffix.lower() == ".npz" and _is_parameterized_stem(p)
364
+ ]
365
+ hit = sorted(hit)
366
+ hit = _cap_per_shape(hit, max_files_per_shape)
367
+ if max_shapes is not None:
368
+ hit = _subsample_shapes(hit, int(max_shapes), int(seed))
369
+ for path in hit:
370
+ if path in seen:
371
+ continue
372
+ seen.add(path)
373
+ chosen.append(path)
374
+ else:
375
+ chosen = _glob_npz(root, npz_glob)
376
+
377
+ npz_only = [p for p in chosen if p.suffix.lower() == ".npz"]
378
+ if exclude_combo:
379
+ npz_only = [p for p in npz_only if not _is_combo_npz(p)]
380
+ if npz_catalog and not npz_paths:
381
+ # Already capped per glob; keep YAML union order (not a global sort).
382
+ capped = npz_only
383
+ else:
384
+ npz_only = sorted(npz_only)
385
+ capped = _cap_per_shape(npz_only, max_files_per_shape)
386
+ if not capped:
387
+ raise FileNotFoundError(
388
+ f"No occupancy NPZ files matched under {root} "
389
+ f"(glob={npz_glob!r}, catalog={bool(npz_catalog)}, "
390
+ f"explicit={bool(npz_paths)})"
391
+ )
392
+ return capped
393
+
394
+
395
+ def _summarize(points: PointsArray, labels: LabelsArray) -> str:
396
+ n = int(points.shape[0])
397
+ inside = float(labels.mean()) if n else float("nan")
398
+ xyz_min = points.min(axis=0) if n else np.full(3, np.nan, dtype=np.float32)
399
+ xyz_max = points.max(axis=0) if n else np.full(3, np.nan, dtype=np.float32)
400
+ return (
401
+ f"N={n}\n"
402
+ f"inside_fraction={inside:.6f}\n"
403
+ f"xyz_min={xyz_min.tolist()}\n"
404
+ f"xyz_max={xyz_max.tolist()}"
405
+ )
406
+
407
+
408
+ # Default smoke-check file from the v2 plan (dataset_test sphere).
409
+ _SAMPLE_RELATIVE = Path("exports") / "dataset_test" / "sphere__raycast_z_raut_s0.15_inout.npz"
410
+
411
+
412
+ if __name__ == "__main__":
413
+ import sys
414
+
415
+ from scatteringnet.config import load_config
416
+
417
+ cfg = load_config()
418
+ if "--catalog" in sys.argv:
419
+ paths = resolve_npz_catalog(
420
+ cfg.data_dir,
421
+ npz_glob=cfg.npz_glob,
422
+ npz_paths=cfg.npz_paths or None,
423
+ npz_catalog=cfg.npz_catalog or None,
424
+ max_files_per_shape=cfg.max_files_per_shape,
425
+ seed=cfg.seed,
426
+ )
427
+ print(f"data_dir={cfg.data_dir}")
428
+ print(f"npz_glob={cfg.npz_glob}")
429
+ print(f"npz_catalog={list(cfg.npz_catalog)}")
430
+ print(f"max_files_per_shape={cfg.max_files_per_shape}")
431
+ print(f"files={len(paths)}")
432
+ # Load a few files only — the full catalog can be thousands of NPZs.
433
+ preview = paths[:3]
434
+ n_all = 0
435
+ n_in = 0
436
+ for path in preview:
437
+ pts, labs = load_points_labels(path)
438
+ n = int(pts.shape[0])
439
+ inside = int((labs == 1.0).sum())
440
+ n_all += n
441
+ n_in += inside
442
+ print(
443
+ f" {path.name} N={n} inside={inside} outside={n - inside}"
444
+ )
445
+ print(
446
+ f"preview_files={len(preview)} preview_N={n_all} "
447
+ f"preview_inside={n_in} preview_outside={n_all - n_in}"
448
+ )
449
+ from scatteringnet.dataset import OccupancyMultiNpzDataset, make_dataloader
450
+
451
+ ds = OccupancyMultiNpzDataset(paths)
452
+ loader = make_dataloader(ds.parts[0], batch_size=8, shuffle=False)
453
+ xyz, y = next(iter(loader))
454
+ print(f"files={len(ds)} n_points={ds.n_points} parts={len(ds.parts)}")
455
+ print(f"batch xyz={tuple(xyz.shape)} y={tuple(y.shape)}")
456
+ else:
457
+ sample = cfg.data_dir / _SAMPLE_RELATIVE
458
+ pts, labs = load_points_labels(sample)
459
+ print(f"file={sample}")
460
+ print(_summarize(pts, labs))
scatteringnet/geometry/__init__.py ADDED
@@ -0,0 +1,10 @@
 
 
 
 
 
 
 
 
 
 
 
1
+ """Training-side geometry: OBJ triangles and envelope samples."""
2
+
3
+ from scatteringnet.geometry.mesh_io import clear_triangle_cache
4
+ from scatteringnet.geometry.surface import clear_envelope_cache
5
+
6
+
7
+ def clear_geometry_caches() -> None:
8
+ """Drop OBJ triangle and envelope process caches."""
9
+ clear_triangle_cache()
10
+ clear_envelope_cache()
scatteringnet/geometry/mesh_io.py ADDED
@@ -0,0 +1,80 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Load OBJ triangle meshes for the NPZ ↔ mesh join.
2
+
3
+ This module only returns ``vertices (V, 3)`` and ``faces (T, 3)``.
4
+ It does **not** sample the envelope.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ from pathlib import Path
10
+ import numpy as np
11
+ import trimesh
12
+ from numpy.typing import NDArray
13
+
14
+ from scatteringnet.geometry.trimesh_util import as_trimesh
15
+
16
+ VerticesArray = NDArray[np.float32]
17
+ FacesArray = NDArray[np.int32]
18
+
19
+ # Same resolved OBJ can back several NPZs (``max_files_per_shape``). Cache
20
+ # the triangle arrays so catalog load does not re-parse the file.
21
+ _TRIANGLE_CACHE: dict[str, tuple[VerticesArray, FacesArray]] = {}
22
+
23
+
24
+ def clear_triangle_cache() -> None:
25
+ """Drop cached OBJ arrays (tests / long-lived notebooks)."""
26
+ _TRIANGLE_CACHE.clear()
27
+
28
+
29
+ def load_obj_triangles(
30
+ path: Path | str,
31
+ *,
32
+ cache: bool = True,
33
+ ) -> tuple[VerticesArray, FacesArray]:
34
+ """
35
+ Read one OBJ as triangle vertices and face indices.
36
+
37
+ Parameters
38
+ ----------
39
+ path:
40
+ Existing ``.obj`` file.
41
+ cache:
42
+ Reuse arrays for the same resolved path (catalog load).
43
+
44
+ Returns
45
+ -------
46
+ vertices:
47
+ ``float32`` array of shape ``(V, 3)``.
48
+ faces:
49
+ ``int32`` array of shape ``(T, 3)`` (0-based vertex indices).
50
+ """
51
+ obj_path = Path(path)
52
+ if not obj_path.is_file():
53
+ raise FileNotFoundError(f"OBJ not found: {obj_path}")
54
+ if obj_path.suffix.lower() != ".obj":
55
+ raise ValueError(f"expected .obj, got {obj_path.suffix!r} ({obj_path})")
56
+
57
+ cache_key = str(obj_path.resolve())
58
+ if cache and cache_key in _TRIANGLE_CACHE:
59
+ return _TRIANGLE_CACHE[cache_key]
60
+
61
+ # process=False keeps the authored vertices; we only need the join.
62
+ loaded = trimesh.load(obj_path, force=None, process=False)
63
+ mesh = as_trimesh(loaded)
64
+ vertices = np.asarray(mesh.vertices, dtype=np.float32)
65
+ faces = np.asarray(mesh.faces, dtype=np.int32)
66
+ if vertices.ndim != 2 or vertices.shape[1] != 3:
67
+ raise ValueError(
68
+ f"vertices must have shape (V, 3), got {tuple(vertices.shape)} in {obj_path}"
69
+ )
70
+ if faces.ndim != 2 or faces.shape[1] != 3:
71
+ raise ValueError(
72
+ f"faces must have shape (T, 3), got {tuple(faces.shape)} in {obj_path}"
73
+ )
74
+ if int(faces.shape[0]) < 1:
75
+ raise ValueError(f"OBJ has no triangles: {obj_path}")
76
+
77
+ arrays = (vertices, faces)
78
+ if cache:
79
+ _TRIANGLE_CACHE[cache_key] = arrays
80
+ return arrays
scatteringnet/geometry/surface.py ADDED
@@ -0,0 +1,154 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Surface envelope samples from a triangle mesh.
2
+
3
+ Area-weighted face darts (larger triangles get more samples). Each
4
+ sample is ``(x, y, z, nx, ny, nz)``: position plus the unit normal of
5
+ the triangle it sits on. World clouds are cached per
6
+ ``(mesh_key, n_surface, seed)``. AABB is applied to XYZ only.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import numpy as np
12
+ import trimesh
13
+ from numpy.typing import NDArray
14
+ from trimesh.sample import sample_surface
15
+
16
+ from scatteringnet.normalize import apply_normalization
17
+
18
+ PointsArray = NDArray[np.float32]
19
+ ENVELOPE_XYZ_DIM = 3
20
+ ENVELOPE_FEAT_DIM = 6
21
+
22
+ # Same OBJ + count + seed → same world samples (always 6-D).
23
+ _ENVELOPE_CACHE: dict[tuple[str, int, int], PointsArray] = {}
24
+
25
+
26
+ def clear_envelope_cache() -> None:
27
+ """Drop cached world-space envelopes (tests / long-lived notebooks)."""
28
+ _ENVELOPE_CACHE.clear()
29
+
30
+
31
+ def _triangle_normals(verts: np.ndarray, faces: np.ndarray) -> NDArray[np.float64]:
32
+ v0 = verts[faces[:, 0]]
33
+ v1 = verts[faces[:, 1]]
34
+ v2 = verts[faces[:, 2]]
35
+ cross = np.cross(v1 - v0, v2 - v0)
36
+ length = np.linalg.norm(cross, axis=1, keepdims=True)
37
+ ok = length[:, 0] > 1e-12
38
+ normals = np.zeros_like(cross)
39
+ normals[ok] = cross[ok] / length[ok]
40
+ return normals
41
+
42
+
43
+ def _pack_xyz_normal(xyz: np.ndarray, normals: np.ndarray) -> PointsArray:
44
+ """Concatenate XYZ with unit face normals → ``(N, 6)``."""
45
+ pos = np.asarray(xyz, dtype=np.float32)
46
+ nrm = np.asarray(normals, dtype=np.float32)
47
+ if pos.shape != nrm.shape or pos.ndim != 2 or pos.shape[1] != 3:
48
+ raise ValueError(
49
+ f"xyz/normals must be (N, 3), got {tuple(pos.shape)} / {tuple(nrm.shape)}"
50
+ )
51
+ length = np.linalg.norm(nrm, axis=1, keepdims=True)
52
+ ok = length[:, 0] > 1e-12
53
+ unit = np.zeros_like(nrm)
54
+ unit[ok] = nrm[ok] / length[ok]
55
+ return np.concatenate([pos, unit], axis=1)
56
+
57
+
58
+ def apply_envelope_aabb(
59
+ env: np.ndarray, center: np.ndarray, scale: float
60
+ ) -> PointsArray:
61
+ """AABB-normalize XYZ; leave normals as unit directions."""
62
+ arr = np.asarray(env, dtype=np.float32)
63
+ if arr.ndim != 2 or arr.shape[1] not in (ENVELOPE_XYZ_DIM, ENVELOPE_FEAT_DIM):
64
+ raise ValueError(
65
+ f"envelope must be (N, 3) or (N, 6), got {tuple(arr.shape)}"
66
+ )
67
+ out = arr.copy()
68
+ out[:, :3] = apply_normalization(out[:, :3], center, scale)
69
+ return out
70
+
71
+
72
+ def undo_envelope_aabb(
73
+ env: np.ndarray, center: np.ndarray, scale: float
74
+ ) -> PointsArray:
75
+ """Undo AABB on XYZ only."""
76
+ arr = np.asarray(env, dtype=np.float32)
77
+ if arr.ndim != 2 or arr.shape[1] not in (ENVELOPE_XYZ_DIM, ENVELOPE_FEAT_DIM):
78
+ raise ValueError(
79
+ f"envelope must be (N, 3) or (N, 6), got {tuple(arr.shape)}"
80
+ )
81
+ out = arr.copy()
82
+ c = np.asarray(center, dtype=np.float32).reshape(3)
83
+ out[:, :3] = out[:, :3] * np.float32(scale) + c
84
+ return out
85
+
86
+
87
+ def project_envelope_dim(env: np.ndarray, dim: int) -> PointsArray:
88
+ """Keep XYZ+normal or drop to XYZ for an older checkpoint."""
89
+ arr = np.asarray(env, dtype=np.float32)
90
+ want = int(dim)
91
+ if want == ENVELOPE_FEAT_DIM:
92
+ if arr.ndim != 2 or arr.shape[1] != ENVELOPE_FEAT_DIM:
93
+ raise ValueError(f"expected (N, 6) envelope, got {tuple(arr.shape)}")
94
+ return arr
95
+ if want == ENVELOPE_XYZ_DIM:
96
+ if arr.ndim != 2 or arr.shape[1] not in (ENVELOPE_XYZ_DIM, ENVELOPE_FEAT_DIM):
97
+ raise ValueError(f"expected (N, 3|6) envelope, got {tuple(arr.shape)}")
98
+ return arr[:, :3]
99
+ raise ValueError(f"envelope dim must be 3 or 6, got {want}")
100
+
101
+
102
+ def _sample_area(
103
+ verts: np.ndarray, tris: np.ndarray, count: int, seed: int
104
+ ) -> PointsArray:
105
+ mesh = trimesh.Trimesh(vertices=verts, faces=tris, process=False)
106
+ points, face_idx = sample_surface(mesh, count, seed=int(seed))
107
+ nrm = _triangle_normals(verts, tris)[np.asarray(face_idx, dtype=np.int64)]
108
+ out = _pack_xyz_normal(points, nrm)
109
+ if out.shape != (count, ENVELOPE_FEAT_DIM):
110
+ raise ValueError(
111
+ f"expected envelope shape {(count, ENVELOPE_FEAT_DIM)}, got {tuple(out.shape)}"
112
+ )
113
+ return out
114
+
115
+
116
+ def sample_surface_points(
117
+ vertices: np.ndarray,
118
+ faces: np.ndarray,
119
+ n_surface: int,
120
+ *,
121
+ seed: int = 1,
122
+ cache_key: str | None = None,
123
+ ) -> PointsArray:
124
+ """
125
+ Sample ``n_surface`` envelope points as ``(N, 6)`` XYZ + unit normal.
126
+
127
+ Face-area weighted: larger triangles receive more darts. No crease
128
+ or fold path — unused ``envelope_mix`` on old checkpoints is ignored.
129
+ """
130
+ count = int(n_surface)
131
+ if count < 1:
132
+ raise ValueError(f"n_surface must be >= 1, got {count}")
133
+ verts = np.asarray(vertices, dtype=np.float64)
134
+ tris = np.asarray(faces, dtype=np.int64)
135
+ if verts.ndim != 2 or verts.shape[1] != 3:
136
+ raise ValueError(f"vertices must have shape (V, 3), got {tuple(verts.shape)}")
137
+ if tris.ndim != 2 or tris.shape[1] != 3:
138
+ raise ValueError(f"faces must have shape (T, 3), got {tuple(tris.shape)}")
139
+
140
+ key: tuple[str, int, int] | None = None
141
+ if cache_key is not None:
142
+ key = (str(cache_key), count, int(seed))
143
+ cached = _ENVELOPE_CACHE.get(key)
144
+ if cached is not None:
145
+ return cached
146
+
147
+ out = _sample_area(verts, tris, count, int(seed))
148
+ if out.shape != (count, ENVELOPE_FEAT_DIM):
149
+ raise ValueError(
150
+ f"expected envelope shape {(count, ENVELOPE_FEAT_DIM)}, got {tuple(out.shape)}"
151
+ )
152
+ if key is not None:
153
+ _ENVELOPE_CACHE[key] = out
154
+ return out
scatteringnet/geometry/trimesh_util.py ADDED
@@ -0,0 +1,19 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Shared trimesh flatten. No Open3D — safe for train-side mesh_io."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from typing import Any
6
+
7
+ import trimesh
8
+
9
+
10
+ def as_trimesh(mesh: Any) -> trimesh.Trimesh:
11
+ """Flatten a Trimesh or a Scene of triangle meshes to one Trimesh."""
12
+ if isinstance(mesh, trimesh.Scene):
13
+ geoms = [g for g in mesh.geometry.values() if isinstance(g, trimesh.Trimesh)]
14
+ if not geoms:
15
+ raise ValueError("Scene contains no triangle meshes")
16
+ mesh = trimesh.util.concatenate(geoms)
17
+ if not isinstance(mesh, trimesh.Trimesh):
18
+ raise TypeError(f"Unsupported mesh type: {type(mesh)!r}")
19
+ return mesh
scatteringnet/infer_multi_npz.py ADDED
@@ -0,0 +1,202 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Classify occupancy NPZs from ``models/<run_id>/best.pt``.
2
+
3
+ Rebuilds ``OccupancyMLP`` or ``OccupancyEncoder`` from the checkpoint
4
+ ``kind``. Query XYZ (and envelope, when conditioned) use the stored
5
+ per-mesh AABB, not a fresh map. Face-token checkpoints are rejected.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ from pathlib import Path
11
+ from typing import Any
12
+
13
+ import numpy as np
14
+ import torch
15
+
16
+ from scatteringnet.config import OccupancyConfig, as_data_relative, load_config, repo_root
17
+ from scatteringnet.data_npz import load_points_labels, load_points_labels_mesh
18
+ from scatteringnet.geometry.mesh_io import load_obj_triangles
19
+ from scatteringnet.geometry.surface import apply_envelope_aabb, project_envelope_dim, sample_surface_points
20
+ from scatteringnet.metrics import occupancy_metrics
21
+ from scatteringnet.normalize import apply_normalization
22
+ from scatteringnet.occupancy_encoder import (
23
+ OccupancyEncoder,
24
+ envelope_dim_from_ckpt,
25
+ envelope_seed_from_ckpt,
26
+ )
27
+ from scatteringnet.occupancy_encoder import CHECKPOINT_KIND as ENCODER_KIND
28
+ from scatteringnet.occupancy_mlp import OccupancyMLP
29
+ from scatteringnet.occupancy_mlp import CHECKPOINT_KIND as MLP_KIND
30
+
31
+
32
+ def resolve_best_pt(
33
+ *,
34
+ checkpoint: Path | str | None = None,
35
+ run_id: str | None = None,
36
+ root: Path | None = None,
37
+ ) -> Path:
38
+ """Resolve ``best.pt`` from an explicit path or ``runs/<id>/checkpoint_dir.txt``."""
39
+ base = (root or repo_root()).resolve()
40
+ if checkpoint is not None:
41
+ path = Path(checkpoint)
42
+ if not path.is_absolute():
43
+ path = base / path
44
+ if not path.is_file():
45
+ raise FileNotFoundError(f"checkpoint not found: {path}")
46
+ return path
47
+ if not run_id:
48
+ raise ValueError("pass checkpoint= or run_id=")
49
+ pointer = base / "runs" / str(run_id) / "checkpoint_dir.txt"
50
+ if not pointer.is_file():
51
+ raise FileNotFoundError(f"run pointer not found: {pointer}")
52
+ rel = pointer.read_text(encoding="utf-8").strip()
53
+ path = Path(rel)
54
+ if not path.is_absolute():
55
+ path = base / path
56
+ best = path / "best.pt" if path.is_dir() else path
57
+ if not best.is_file():
58
+ raise FileNotFoundError(f"best.pt not found: {best}")
59
+ return best
60
+
61
+
62
+ def load_occupancy_model(
63
+ ckpt: dict[str, Any],
64
+ device: torch.device,
65
+ ) -> OccupancyMLP | OccupancyEncoder:
66
+ """Rebuild the head recorded in ``ckpt['kind']`` and load weights."""
67
+ hidden = int(ckpt["hidden"])
68
+ depth = int(ckpt["depth"])
69
+ kind = str(ckpt["kind"])
70
+ if kind == MLP_KIND:
71
+ model: OccupancyMLP | OccupancyEncoder = OccupancyMLP(hidden=hidden, depth=depth)
72
+ elif kind == ENCODER_KIND:
73
+ latent = int(ckpt["latent_dim"]) if ckpt.get("latent_dim") is not None else hidden
74
+ enc = str(ckpt.get("shape_encoder") or "surface").strip().lower()
75
+ if enc == "mesh":
76
+ raise ValueError(
77
+ "face-token occupancy checkpoints (shape_encoder='mesh') "
78
+ "are no longer supported"
79
+ )
80
+ knn_k = int(ckpt["knn_k"]) if ckpt.get("knn_k") is not None else 0
81
+ knn_local = None
82
+ if knn_k > 0 and ckpt.get("knn_local_dim") is not None:
83
+ knn_local = int(ckpt["knn_local_dim"])
84
+ model = OccupancyEncoder(
85
+ hidden=hidden,
86
+ depth=depth,
87
+ latent_dim=latent,
88
+ shape_encoder=enc,
89
+ knn_k=knn_k,
90
+ knn_local_dim=knn_local,
91
+ envelope_dim=envelope_dim_from_ckpt(ckpt),
92
+ )
93
+ else:
94
+ raise ValueError(f"unsupported checkpoint kind {kind!r}")
95
+ model.load_state_dict(ckpt["state_dict"])
96
+ return model.to(device).eval()
97
+
98
+
99
+ def _part_for_npz(ckpt: dict[str, Any], npz_path: Path, data_dir: Path) -> dict[str, Any]:
100
+ rel = as_data_relative(npz_path, data_dir)
101
+ name = npz_path.name
102
+ for part in ckpt.get("parts") or []:
103
+ stored = str(part.get("npz", ""))
104
+ if stored == rel or Path(stored).name == name:
105
+ return part
106
+ raise KeyError(f"no AABB part for {npz_path.name} in checkpoint")
107
+
108
+
109
+ def infer_npz(
110
+ cfg: OccupancyConfig,
111
+ *,
112
+ checkpoint: Path | str | None = None,
113
+ run_id: str | None = None,
114
+ npz_path: Path | str | None = None,
115
+ root: Path | None = None,
116
+ ) -> dict[str, float]:
117
+ """
118
+ Score one NPZ with ``best.pt``. Returns accuracy / IoU / F1.
119
+ """
120
+ ckpt_path = resolve_best_pt(checkpoint=checkpoint, run_id=run_id, root=root)
121
+ ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=False)
122
+ model = load_occupancy_model(ckpt, cfg.device)
123
+ if npz_path is None:
124
+ rels = ckpt.get("npz_paths") or []
125
+ if not rels:
126
+ raise ValueError("checkpoint has no npz_paths; pass npz_path=")
127
+ npz = cfg.data_dir / str(rels[0])
128
+ else:
129
+ npz = Path(npz_path)
130
+ if not npz.is_absolute():
131
+ npz = cfg.data_dir / npz
132
+ part = _part_for_npz(ckpt, npz, cfg.data_dir)
133
+ center = np.asarray(part["center"], dtype=np.float32)
134
+ scale = float(part["scale"])
135
+ geom = None
136
+ shape_id = None
137
+ enc = str(ckpt.get("shape_encoder", "none")).strip().lower()
138
+ if enc == "mesh":
139
+ raise ValueError(
140
+ "face-token occupancy checkpoints (shape_encoder='mesh') "
141
+ "are no longer supported"
142
+ )
143
+ if enc == "surface":
144
+ points, labels, mesh_path = load_points_labels_mesh(npz, cfg.data_dir)
145
+ vertices, faces = load_obj_triangles(mesh_path)
146
+ cache_key = str(mesh_path.resolve())
147
+ n_surface = int(ckpt.get("n_surface") or cfg.n_surface)
148
+ world = sample_surface_points(
149
+ vertices,
150
+ faces,
151
+ n_surface,
152
+ seed=envelope_seed_from_ckpt(ckpt),
153
+ cache_key=cache_key,
154
+ )
155
+ arr = apply_envelope_aabb(world, center, scale)
156
+ arr = project_envelope_dim(arr, envelope_dim_from_ckpt(ckpt))
157
+ geom = torch.from_numpy(arr).unsqueeze(0)
158
+ shape_id = torch.zeros((), dtype=torch.long)
159
+ else:
160
+ points, labels = load_points_labels(npz)
161
+ xyz = torch.from_numpy(apply_normalization(points, center, scale))
162
+ y = torch.from_numpy(np.asarray(labels, dtype=np.float32)).unsqueeze(1)
163
+ logits_rows: list[torch.Tensor] = []
164
+ with torch.no_grad():
165
+ for start in range(0, int(xyz.shape[0]), int(cfg.batch_size)):
166
+ sl = slice(start, start + int(cfg.batch_size))
167
+ batch_xyz = xyz[sl].to(cfg.device)
168
+ if geom is None:
169
+ logits_rows.append(model(batch_xyz).cpu())
170
+ else:
171
+ b = int(batch_xyz.shape[0])
172
+ geom_b = geom.expand(b, -1, -1).to(cfg.device)
173
+ sid = shape_id.expand(b).to(cfg.device)
174
+ logits_rows.append(model(batch_xyz, geom_b, sid).cpu())
175
+ logits = torch.cat(logits_rows, dim=0)
176
+ scores = occupancy_metrics(logits, y)
177
+ print(
178
+ f"checkpoint={ckpt_path} npz={npz.name} "
179
+ f"acc={scores.accuracy:.4f} iou={scores.inside_iou:.4f} "
180
+ f"f1={scores.inside_f1:.4f}"
181
+ )
182
+ return {
183
+ "accuracy": scores.accuracy,
184
+ "inside_iou": scores.inside_iou,
185
+ "inside_f1": scores.inside_f1,
186
+ }
187
+
188
+
189
+ if __name__ == "__main__":
190
+ import argparse
191
+
192
+ parser = argparse.ArgumentParser(description="Infer occupancy from models/<id>/best.pt")
193
+ parser.add_argument("--checkpoint", default=None, help="Path to best.pt")
194
+ parser.add_argument("--run-id", default=None, help="runs/<id> folder name")
195
+ parser.add_argument("--npz", default=None, help="NPZ path (data_dir-relative or absolute)")
196
+ args = parser.parse_args()
197
+ infer_npz(
198
+ load_config(),
199
+ checkpoint=args.checkpoint,
200
+ run_id=args.run_id,
201
+ npz_path=args.npz,
202
+ )
scatteringnet/metrics.py ADDED
@@ -0,0 +1,156 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Occupancy classification metrics from logits vs labels.
2
+
3
+ The model emits raw logits (no sigmoid in ``forward``). Metrics apply
4
+ sigmoid only at decision time so they stay consistent with
5
+ ``BCEWithLogitsLoss`` and with inference (threshold ``0.5``).
6
+
7
+ Inside (label ``1``) is the product-relevant class: points that belong
8
+ in the interior. Precision / recall are therefore reported for that
9
+ class only. No extra packages (sklearn, etc.).
10
+ """
11
+
12
+ from __future__ import annotations
13
+
14
+ from dataclasses import dataclass
15
+
16
+ import torch
17
+ from torch import Tensor
18
+
19
+
20
+ @dataclass(frozen=True)
21
+ class OccupancyMetrics:
22
+ """Pointwise occupancy scores for one batch or a full set of queries."""
23
+
24
+ accuracy: float
25
+ inside_precision: float
26
+ inside_recall: float
27
+ inside_iou: float
28
+ inside_f1: float
29
+
30
+
31
+ def _safe_div(numerator: float, denominator: float) -> float:
32
+ """Return ``0.0`` when the count in the denominator is zero (no sklearn)."""
33
+ if denominator <= 0.0:
34
+ return 0.0
35
+ return numerator / denominator
36
+
37
+
38
+ def _flatten_pair(logits: Tensor, labels: Tensor) -> tuple[Tensor, Tensor]:
39
+ """
40
+ Collapse ``(B, 1)`` or ``(B,)`` logits / labels to a shared 1-D view.
41
+
42
+ Last dim of logits is 1 when it comes from ``OccupancyMLP``; labels from
43
+ the Dataset match that. A 1-D label vector is accepted so callers do not
44
+ have to unsqueeze.
45
+ """
46
+ if logits.numel() != labels.numel():
47
+ raise ValueError(
48
+ f"logits and labels must have the same number of elements, "
49
+ f"got logits={tuple(logits.shape)} labels={tuple(labels.shape)}"
50
+ )
51
+ return logits.reshape(-1), labels.reshape(-1)
52
+
53
+
54
+ def occupancy_metrics(
55
+ logits: Tensor,
56
+ labels: Tensor,
57
+ *,
58
+ threshold: float = 0.5,
59
+ ) -> OccupancyMetrics:
60
+ """
61
+ Accuracy plus inside precision / recall from occupancy logits.
62
+
63
+ Parameters
64
+ ----------
65
+ logits:
66
+ Unnormalized scores, shape ``(B, 1)`` or ``(B,)``. Positive → inside.
67
+ labels:
68
+ Float ``{0, 1}`` with the same number of elements as ``logits``.
69
+ threshold:
70
+ Decision cut on ``sigmoid(logit)``. Default ``0.5`` matches inference.
71
+
72
+ Returns
73
+ -------
74
+ OccupancyMetrics
75
+ Scalar floats on CPU (safe to print or average across batches).
76
+ """
77
+ # Sigmoid here only: training still uses BCE-with-logits on raw logits.
78
+ tp, fp, fn, correct, n = occupancy_counts(logits, labels, threshold=threshold)
79
+ return occupancy_metrics_from_counts(tp=tp, fp=fp, fn=fn, correct=correct, n=n)
80
+
81
+
82
+ def occupancy_counts(
83
+ logits: Tensor,
84
+ labels: Tensor,
85
+ *,
86
+ threshold: float = 0.5,
87
+ ) -> tuple[float, float, float, float, float]:
88
+ """Return ``(tp, fp, fn, correct, n)`` for a micro-average over points."""
89
+ logits_flat, labels_flat = _flatten_pair(logits, labels)
90
+ pred_inside = logits_flat.sigmoid() >= threshold
91
+ true_inside = labels_flat > 0.5
92
+ pred_f = pred_inside.to(dtype=torch.float32)
93
+ true_f = true_inside.to(dtype=torch.float32)
94
+ tp = float((pred_f * true_f).sum().item())
95
+ fp = float((pred_f * (1.0 - true_f)).sum().item())
96
+ fn = float(((1.0 - pred_f) * true_f).sum().item())
97
+ correct = float((pred_inside == true_inside).to(dtype=torch.float32).sum().item())
98
+ n = float(pred_inside.numel())
99
+ return tp, fp, fn, correct, n
100
+
101
+
102
+ def occupancy_metrics_from_counts(
103
+ *,
104
+ tp: float,
105
+ fp: float,
106
+ fn: float,
107
+ correct: float,
108
+ n: float,
109
+ ) -> OccupancyMetrics:
110
+ """Build metrics from accumulated confusion counts (point micro-average)."""
111
+ precision = _safe_div(tp, tp + fp)
112
+ recall = _safe_div(tp, tp + fn)
113
+ return OccupancyMetrics(
114
+ accuracy=_safe_div(correct, n),
115
+ inside_precision=precision,
116
+ inside_recall=recall,
117
+ inside_iou=_safe_div(tp, tp + fp + fn),
118
+ inside_f1=_safe_div(2.0 * precision * recall, precision + recall),
119
+ )
120
+
121
+
122
+ def accuracy_from_logits(
123
+ logits: Tensor,
124
+ labels: Tensor,
125
+ *,
126
+ threshold: float = 0.5,
127
+ ) -> float:
128
+ """Fraction of points whose thresholded sigmoid matches the 0/1 label.
129
+
130
+ Parameters
131
+ ----------
132
+ logits, labels, threshold:
133
+ Same meaning as :func:`occupancy_metrics`.
134
+
135
+ Returns
136
+ -------
137
+ float
138
+ Accuracy in ``[0, 1]``.
139
+ """
140
+ return occupancy_metrics(logits, labels, threshold=threshold).accuracy
141
+
142
+
143
+ if __name__ == "__main__":
144
+ # Deterministic 8-point batch: 4 true insides, 4 true outsides, all correct.
145
+ demo_logits = torch.tensor(
146
+ [[4.0], [3.0], [2.0], [1.0], [-1.0], [-2.0], [-3.0], [-4.0]]
147
+ )
148
+ demo_labels = torch.tensor(
149
+ [[1.0], [1.0], [1.0], [1.0], [0.0], [0.0], [0.0], [0.0]]
150
+ )
151
+ scores = occupancy_metrics(demo_logits, demo_labels)
152
+ print(f"accuracy={scores.accuracy:.4f}")
153
+ print(f"inside_precision={scores.inside_precision:.4f}")
154
+ print(f"inside_recall={scores.inside_recall:.4f}")
155
+ print(f"inside_iou={scores.inside_iou:.4f}")
156
+ print(f"inside_f1={scores.inside_f1:.4f}")
scatteringnet/normalize.py ADDED
@@ -0,0 +1,104 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """AABB normalization for occupancy XYZ.
2
+
3
+ Maps a point cloud into a roughly ``[-1, 1]^3`` cube so the MLP sees
4
+ comparable coordinates across differently sized meshes.
5
+
6
+ center = midpoint of the axis-aligned bounding box
7
+ scale = maximum half-extent (longest AABB side / 2)
8
+
9
+ Normalized point: ``(xyz - center) / scale``.
10
+ This module does not touch the occupancy model.
11
+ """
12
+
13
+ from __future__ import annotations
14
+
15
+ import numpy as np
16
+ from numpy.typing import NDArray
17
+
18
+ PointsArray = NDArray[np.float32]
19
+
20
+
21
+ def compute_center_scale(points: np.ndarray) -> tuple[PointsArray, float]:
22
+ """
23
+ AABB center and max half-extent for an ``(N, 3)`` point array.
24
+
25
+ Parameters
26
+ ----------
27
+ points:
28
+ Query XYZ, shape ``(N, 3)``, at least one row.
29
+
30
+ Returns
31
+ -------
32
+ center:
33
+ ``float32`` vector of shape ``(3,)``.
34
+ scale:
35
+ Positive float (max half-extent). Raises if the cloud has no extent.
36
+ """
37
+ pts = np.asarray(points, dtype=np.float32)
38
+ if pts.ndim != 2 or pts.shape[1] != 3:
39
+ raise ValueError(f"points must have shape (N, 3), got {tuple(pts.shape)}")
40
+ if pts.shape[0] == 0:
41
+ raise ValueError("points must contain at least one row")
42
+
43
+ xyz_min = pts.min(axis=0)
44
+ xyz_max = pts.max(axis=0)
45
+ center = 0.5 * (xyz_min + xyz_max)
46
+ half_extents = 0.5 * (xyz_max - xyz_min)
47
+ scale = float(np.max(half_extents))
48
+ if scale <= 0.0:
49
+ raise ValueError(
50
+ "scale must be > 0; all points appear to share the same location"
51
+ )
52
+ return center.astype(np.float32, copy=False), scale
53
+
54
+
55
+ def apply_normalization(
56
+ points: np.ndarray,
57
+ center: np.ndarray,
58
+ scale: float,
59
+ ) -> PointsArray:
60
+ """
61
+ Return ``(points - center) / scale`` as ``float32 (N, 3)``.
62
+
63
+ Parameters
64
+ ----------
65
+ points:
66
+ Query XYZ, shape ``(N, 3)``.
67
+ center:
68
+ AABB midpoint, shape ``(3,)``.
69
+ scale:
70
+ Positive max half-extent.
71
+
72
+ Returns
73
+ -------
74
+ ndarray
75
+ Normalized points, ``float32 (N, 3)``.
76
+ """
77
+ if scale <= 0.0:
78
+ raise ValueError(f"scale must be > 0, got {scale}")
79
+ pts = np.asarray(points, dtype=np.float32)
80
+ if pts.ndim != 2 or pts.shape[1] != 3:
81
+ raise ValueError(f"points must have shape (N, 3), got {tuple(pts.shape)}")
82
+ c = np.asarray(center, dtype=np.float32).reshape(3)
83
+ return (pts - c) / np.float32(scale)
84
+
85
+
86
+ if __name__ == "__main__":
87
+ from scatteringnet.config import load_config
88
+ from scatteringnet.data_npz import load_points_labels
89
+
90
+ sample = (
91
+ load_config().data_dir
92
+ / "exports"
93
+ / "dataset_test"
94
+ / "sphere__raycast_z_raut_s0.15_inout.npz"
95
+ )
96
+ points, _labels = load_points_labels(sample)
97
+ center, scale = compute_center_scale(points)
98
+ normed = apply_normalization(points, center, scale)
99
+ recovered = normed[:3] * np.float32(scale) + center
100
+ print(f"file={sample}")
101
+ print(f"center={center.tolist()} scale={scale:.6f}")
102
+ print(f"normed_min={normed.min(axis=0).tolist()}")
103
+ print(f"normed_max={normed.max(axis=0).tolist()}")
104
+ print(f"inverse_ok={np.allclose(recovered, points[:3], rtol=1e-5, atol=1e-5)}")
scatteringnet/occupancy_encoder.py ADDED
@@ -0,0 +1,248 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Geometry-conditioned occupancy: query XYZ plus a shape latent.
2
+
3
+ ``OccupancyMLP`` stays xyz-only. ``shape_encoder: surface`` is the
4
+ envelope PointNet over ``(N, 6)`` XYZ + face normal (or ``(N, 3)``
5
+ for older checkpoints). ``knn_k > 0`` adds per-query nearest-neighbor
6
+ features (XYZ offset, plus the neighbor normal when the cloud is 6-D).
7
+ The face-token head (``shape_encoder: mesh``) was removed.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ import torch
13
+ import torch.nn as nn
14
+ from torch import Tensor
15
+
16
+ from scatteringnet.occupancy_mlp import build_mlp
17
+
18
+ # Distinct from OccupancyMLP so infer can tell the checkpoint apart.
19
+ CHECKPOINT_KIND = "occupancy_encoder"
20
+
21
+
22
+ class SurfaceEncoder(nn.Module):
23
+ """
24
+ Per-point MLP + max-pool (PointNet) over an envelope cloud.
25
+
26
+ Shapes
27
+ ------
28
+ envelope: ``(U, N, C)`` unique meshes; ``C`` is 3 (XYZ) or 6 (XYZ+n)
29
+ output: ``(U, D)`` one latent per mesh
30
+ """
31
+
32
+ def __init__(
33
+ self, latent_dim: int = 64, hidden: int = 64, *, in_dim: int = 3
34
+ ) -> None:
35
+ super().__init__()
36
+ if latent_dim < 1:
37
+ raise ValueError(f"latent_dim must be >= 1, got {latent_dim}")
38
+ if hidden < 1:
39
+ raise ValueError(f"hidden must be >= 1, got {hidden}")
40
+ dim = int(in_dim)
41
+ if dim not in (3, 6):
42
+ raise ValueError(f"in_dim must be 3 or 6, got {dim}")
43
+ self.latent_dim = latent_dim
44
+ self.hidden = hidden
45
+ self.in_dim = dim
46
+ self.point_mlp = nn.Sequential(
47
+ nn.Linear(dim, hidden),
48
+ nn.ReLU(inplace=True),
49
+ nn.Linear(hidden, latent_dim),
50
+ )
51
+
52
+ def forward(self, envelope: Tensor) -> Tensor:
53
+ """Max-pool per-point features → one vector per unique mesh."""
54
+ if envelope.ndim != 3 or envelope.shape[-1] != self.in_dim:
55
+ raise ValueError(
56
+ f"envelope must have shape (U, N, {self.in_dim}), "
57
+ f"got {tuple(envelope.shape)}"
58
+ )
59
+ features = self.point_mlp(envelope)
60
+ return features.max(dim=1).values
61
+
62
+
63
+ def knn_offsets(xyz: Tensor, envelope: Tensor, k: int) -> Tensor:
64
+ """
65
+ Offsets from each query to its ``k`` nearest envelope points.
66
+
67
+ Distances use XYZ only (AABB Euclidean). ``k`` is clamped to N.
68
+ If the cloud is ``(B, N, 6)``, each neighbor is
69
+ ``(dx, dy, dz, nx, ny, nz)`` — relative position plus that
70
+ neighbor's stored normal. Query points have no normal.
71
+
72
+ Shapes: ``xyz (B, 3)``, ``envelope (B, N, 3|6)`` → ``(B, k, 3|6)``.
73
+ """
74
+ if xyz.ndim != 2 or xyz.shape[-1] != 3:
75
+ raise ValueError(f"xyz must have shape (B, 3), got {tuple(xyz.shape)}")
76
+ if envelope.ndim != 3 or envelope.shape[-1] not in (3, 6):
77
+ raise ValueError(
78
+ f"envelope must have shape (B, N, 3 or 6), got {tuple(envelope.shape)}"
79
+ )
80
+ if int(xyz.shape[0]) != int(envelope.shape[0]):
81
+ raise ValueError(
82
+ f"xyz/envelope batch mismatch: {tuple(xyz.shape)} vs {tuple(envelope.shape)}"
83
+ )
84
+ n_env = int(envelope.shape[1])
85
+ if n_env < 1:
86
+ raise ValueError("envelope length N must be >= 1")
87
+ take = min(int(k), n_env)
88
+ if take < 1:
89
+ raise ValueError(f"k must be >= 1, got {k}")
90
+ feat = int(envelope.shape[-1])
91
+ # k-NN is position-only; extras (normals) ride along after the gather.
92
+ pos = envelope[..., :3]
93
+ dist = torch.linalg.norm(pos - xyz.unsqueeze(1), dim=-1)
94
+ idx = dist.topk(take, dim=-1, largest=False).indices
95
+ nbrs = torch.gather(envelope, 1, idx.unsqueeze(-1).expand(-1, -1, feat))
96
+ rel_xyz = nbrs[..., :3] - xyz.unsqueeze(1)
97
+ if feat == 3:
98
+ return rel_xyz
99
+ return torch.cat([rel_xyz, nbrs[..., 3:]], dim=-1)
100
+
101
+
102
+ # Catalog trains used YAML seed 1. Old best.pt files omit ``seed``.
103
+ DEFAULT_ENVELOPE_SEED = 1
104
+
105
+
106
+ def envelope_seed_from_ckpt(ckpt: dict) -> int:
107
+ """
108
+ Envelope RNG seed this checkpoint was trained with.
109
+
110
+ New trains store ``seed`` on ``best.pt``. Older files omit it; do
111
+ not fall back to live YAML (that knob may have changed since train).
112
+ """
113
+ raw = ckpt.get("seed")
114
+ if raw is None:
115
+ return DEFAULT_ENVELOPE_SEED
116
+ return int(raw)
117
+
118
+
119
+ def envelope_dim_from_ckpt(ckpt: dict) -> int:
120
+ """
121
+ Envelope channel count this checkpoint was trained with.
122
+
123
+ New trains store ``envelope_dim``. Older XYZ-only ``best.pt`` files
124
+ omit it; the first SurfaceEncoder Linear in-features is then 3.
125
+ """
126
+ raw = ckpt.get("envelope_dim")
127
+ if raw is not None:
128
+ dim = int(raw)
129
+ if dim not in (3, 6):
130
+ raise ValueError(f"envelope_dim must be 3 or 6, got {dim}")
131
+ return dim
132
+ weight = (ckpt.get("state_dict") or {}).get("surface.point_mlp.0.weight")
133
+ if weight is not None:
134
+ dim = int(weight.shape[1])
135
+ if dim in (3, 6):
136
+ return dim
137
+ return 3
138
+
139
+
140
+ class OccupancyEncoder(nn.Module):
141
+ """
142
+ Occupancy logits from query XYZ and an envelope code.
143
+
144
+ Unique ``shape_id`` values are encoded **once** per batch, then
145
+ broadcast. ``knn_k > 0`` concatenates a local envelope code.
146
+
147
+ Shapes
148
+ ------
149
+ xyz: ``(B, 3)``
150
+ geom: ``(B, N, C)`` envelope; ``C`` is ``envelope_dim`` (3 or 6)
151
+ shape_id: ``(B,)`` long
152
+ output: ``(B, 1)`` logits
153
+ """
154
+
155
+ def __init__(
156
+ self,
157
+ hidden: int = 64,
158
+ depth: int = 4,
159
+ latent_dim: int = 64,
160
+ *,
161
+ shape_encoder: str = "surface",
162
+ knn_k: int = 0,
163
+ knn_local_dim: int | None = None,
164
+ envelope_dim: int = 6,
165
+ ) -> None:
166
+ super().__init__()
167
+ if hidden < 1:
168
+ raise ValueError(f"hidden must be >= 1, got {hidden}")
169
+ if depth < 1:
170
+ raise ValueError(f"depth must be >= 1, got {depth}")
171
+ if latent_dim < 1:
172
+ raise ValueError(f"latent_dim must be >= 1, got {latent_dim}")
173
+ kind = str(shape_encoder).strip().lower()
174
+ if kind != "surface":
175
+ raise ValueError(
176
+ "OccupancyEncoder only supports shape_encoder='surface' "
177
+ f"(face-token 'mesh' was removed), got {shape_encoder!r}"
178
+ )
179
+ k = int(knn_k)
180
+ if k < 0:
181
+ raise ValueError(f"knn_k must be >= 0, got {k}")
182
+ local_dim = int(knn_local_dim) if knn_local_dim is not None else int(latent_dim)
183
+ if k > 0 and local_dim < 1:
184
+ raise ValueError(f"knn_local_dim must be >= 1, got {local_dim}")
185
+ ed = int(envelope_dim)
186
+ if ed not in (3, 6):
187
+ raise ValueError(f"envelope_dim must be 3 or 6, got {ed}")
188
+ self.hidden = hidden
189
+ self.depth = depth
190
+ self.latent_dim = latent_dim
191
+ self.shape_encoder = kind
192
+ self.knn_k = k
193
+ self.knn_local_dim = local_dim if k > 0 else 0
194
+ self.envelope_dim = ed
195
+ # Name ``surface`` is load-stable for existing envelope checkpoints.
196
+ self.surface = SurfaceEncoder(latent_dim=latent_dim, hidden=hidden, in_dim=ed)
197
+ self.token_dim = ed
198
+ head_in = 3 + latent_dim
199
+ if k > 0:
200
+ # Same PointNet block as the global envelope, over k neighbor features.
201
+ self.local = SurfaceEncoder(latent_dim=local_dim, hidden=hidden, in_dim=ed)
202
+ head_in += local_dim
203
+ self.head = build_mlp(head_in, hidden, depth)
204
+
205
+ def encode_unique(self, geom: Tensor, shape_id: Tensor) -> Tensor:
206
+ """
207
+ Encode each distinct ``shape_id`` once and scatter back to ``(B, D)``.
208
+
209
+ Parameters
210
+ ----------
211
+ geom, shape_id:
212
+ Batched envelope clouds and integer mesh ids (same length B).
213
+ """
214
+ if shape_id.ndim != 1 or int(shape_id.shape[0]) != int(geom.shape[0]):
215
+ raise ValueError(
216
+ f"shape_id must be (B,), got {tuple(shape_id.shape)} "
217
+ f"for geom {tuple(geom.shape)}"
218
+ )
219
+ unique_ids, inverse = torch.unique(shape_id, sorted=True, return_inverse=True)
220
+ hits = shape_id.unsqueeze(0) == unique_ids.unsqueeze(1)
221
+ first = hits.to(dtype=torch.int64).argmax(dim=1)
222
+ z_unique = self.surface(geom[first])
223
+ return z_unique[inverse]
224
+
225
+ def forward(
226
+ self,
227
+ xyz: Tensor,
228
+ geom: Tensor,
229
+ shape_id: Tensor,
230
+ ) -> Tensor:
231
+ """``cat(xyz, z_global[, z_local])`` → occupancy logit."""
232
+ if xyz.ndim != 2 or xyz.shape[-1] != 3:
233
+ raise ValueError(f"xyz must have shape (B, 3), got {tuple(xyz.shape)}")
234
+ if geom.ndim != 3 or geom.shape[-1] != self.token_dim:
235
+ raise ValueError(
236
+ f"geom must have shape (B, K, {self.token_dim}), got {tuple(geom.shape)}"
237
+ )
238
+ if int(xyz.shape[0]) != int(geom.shape[0]):
239
+ raise ValueError(
240
+ f"xyz/geom batch mismatch: {tuple(xyz.shape)} vs {tuple(geom.shape)}"
241
+ )
242
+ ids = shape_id.reshape(-1)
243
+ z_shape = self.encode_unique(geom, ids)
244
+ pieces = [xyz, z_shape]
245
+ if self.knn_k > 0:
246
+ rel = knn_offsets(xyz, geom, self.knn_k)
247
+ pieces.append(self.local(rel))
248
+ return self.head(torch.cat(pieces, dim=-1))
scatteringnet/occupancy_mlp.py ADDED
@@ -0,0 +1,87 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Occupancy MLP: raw XYZ → inside/outside logit."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import torch
6
+ import torch.nn as nn
7
+ from torch import Tensor
8
+
9
+ # Checkpoint schema tag (not a YAML knob).
10
+ CHECKPOINT_KIND = "occupancy_mlp"
11
+
12
+
13
+ def build_mlp(in_dim: int, hidden: int, depth: int) -> nn.Sequential:
14
+ """Linear→ReLU × ``depth`` then a 1-logit head. Shared by OccupancyMLP."""
15
+ if in_dim < 1:
16
+ raise ValueError(f"in_dim must be >= 1, got {in_dim}")
17
+ if hidden < 1:
18
+ raise ValueError(f"hidden must be >= 1, got {hidden}")
19
+ if depth < 1:
20
+ raise ValueError(f"depth must be >= 1, got {depth}")
21
+ layers: list[nn.Module] = []
22
+ dim = in_dim
23
+ for _ in range(depth):
24
+ layers.append(nn.Linear(dim, hidden))
25
+ layers.append(nn.ReLU(inplace=True))
26
+ dim = hidden
27
+ layers.append(nn.Linear(dim, 1))
28
+ return nn.Sequential(*layers)
29
+
30
+
31
+ class OccupancyMLP(nn.Module):
32
+ """
33
+ Tiny fully-connected occupancy field.
34
+
35
+ Maps a batch of 3D query coordinates to a single unnormalized logit per
36
+ point. A later training step will apply ``binary_cross_entropy_with_logits``
37
+ (do not softmax / sigmoid inside ``forward``).
38
+
39
+ Shapes
40
+ ------
41
+ xyz: ``(B, 3)`` batch of query points (device follows the caller)
42
+ output: ``(B, 1)`` logits; positive → inside, negative → outside
43
+
44
+ Device
45
+ ------
46
+ Parameters live on whatever device the module was moved to
47
+ (``.to(device)`` / ``.cuda()``). ``xyz`` must already be on that same
48
+ device; this module does not copy tensors.
49
+ """
50
+
51
+ def __init__(self, hidden: int = 64, depth: int = 4) -> None:
52
+ """
53
+ Build Linear→ReLU blocks then a 1-logit head.
54
+
55
+ Parameters
56
+ ----------
57
+ hidden:
58
+ Channel width of each hidden Linear (must be ``>= 1``).
59
+ depth:
60
+ Number of hidden Linear+ReLU blocks (must be ``>= 1``).
61
+ """
62
+ super().__init__()
63
+ self.hidden = hidden
64
+ self.depth = depth
65
+ # First Linear is 3 → H; remaining blocks are H → H.
66
+ self.net = build_mlp(3, hidden, depth)
67
+
68
+ def forward(self, xyz: Tensor) -> Tensor:
69
+ """
70
+ Evaluate occupancy logits at query coordinates.
71
+
72
+ Parameters
73
+ ----------
74
+ xyz:
75
+ Float tensor of shape ``(B, 3)``. Last dim is Cartesian XYZ.
76
+
77
+ Returns
78
+ -------
79
+ Tensor
80
+ Float tensor of shape ``(B, 1)`` on the same device as ``xyz``.
81
+ """
82
+ if xyz.ndim != 2 or xyz.shape[-1] != 3:
83
+ raise ValueError(
84
+ f"xyz must have shape (B, 3), got {tuple(xyz.shape)}"
85
+ )
86
+ # Sequential Linear layers require matching dtype/device with parameters.
87
+ return self.net(xyz)
scatteringnet/viewer/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ """Three.js inspect helper (HTTP + Job A / Job B)."""
scatteringnet/viewer/infer_job.py ADDED
@@ -0,0 +1,536 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Job A / B: classify uploaded points with a ``models/*/best.pt``.
2
+
3
+ Reuses ``load_occupancy_model`` and AABB helpers. Envelope tokens follow
4
+ the **checkpoint**, not live ``config.yaml``. Face-token checkpoints
5
+ are rejected. Does not modify occupancy train/infer modules.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import base64
11
+ import threading
12
+ import time
13
+ from pathlib import Path
14
+ from typing import Any, TYPE_CHECKING
15
+
16
+ if TYPE_CHECKING:
17
+ from scatteringnet.config import OccupancyConfig
18
+
19
+ import numpy as np
20
+
21
+
22
+ from scatteringnet.infer_multi_npz import load_occupancy_model # noqa: E402
23
+ from scatteringnet.geometry.mesh_io import load_obj_triangles # noqa: E402
24
+ from scatteringnet.geometry.surface import ( # noqa: E402
25
+ apply_envelope_aabb,
26
+ project_envelope_dim,
27
+ sample_surface_points,
28
+ )
29
+ from scatteringnet.metrics import occupancy_metrics # noqa: E402
30
+ from scatteringnet.normalize import apply_normalization, compute_center_scale # noqa: E402
31
+ from scatteringnet.occupancy_encoder import CHECKPOINT_KIND as ENCODER_KIND # noqa: E402
32
+ from scatteringnet.occupancy_encoder import envelope_dim_from_ckpt # noqa: E402
33
+ from scatteringnet.occupancy_encoder import envelope_seed_from_ckpt # noqa: E402
34
+
35
+ from scatteringnet.viewer.mesh_access import resolve_viewer_mesh # noqa: E402
36
+ from scatteringnet.viewer.model_access import match_checkpoint_part, resolve_viewer_checkpoint # noqa: E402
37
+ from scatteringnet.viewer.obj_fill import triangles_from_obj_text # noqa: E402
38
+
39
+ MAX_POINTS = 2_000_000
40
+
41
+ # One occupancy head in this process. Same checkpoint + device → skip torch.load.
42
+ _CACHE_LOCK = threading.Lock()
43
+ _MODEL_CACHE: tuple[tuple[str, int, str], Any, dict[str, Any]] | None = None
44
+
45
+
46
+ def clear_model_cache() -> None:
47
+ """Drop the cached head (tests / swapped GPU)."""
48
+ global _MODEL_CACHE
49
+ with _CACHE_LOCK:
50
+ _MODEL_CACHE = None
51
+
52
+
53
+ def _model_cache_key(ckpt_path: Path, device: Any) -> tuple[str, int, str]:
54
+ stat = ckpt_path.stat()
55
+ return (str(ckpt_path.resolve()), int(stat.st_mtime_ns), str(device))
56
+
57
+
58
+ def load_cached_occupancy(
59
+ run_id: str,
60
+ models_root: Path,
61
+ device: Any,
62
+ ) -> tuple[Any, dict[str, Any], Path, bool]:
63
+ """
64
+ Load ``best.pt`` once per (path, mtime, device).
65
+
66
+ Returns ``(model, ckpt, path, cache_hit)``. Switching run_id replaces
67
+ the slot so GPU RAM does not keep every inspect checkpoint.
68
+ """
69
+ import torch
70
+
71
+ ckpt_path = resolve_viewer_checkpoint(run_id, models_root)
72
+ key = _model_cache_key(ckpt_path, device)
73
+ global _MODEL_CACHE
74
+ with _CACHE_LOCK:
75
+ slot = _MODEL_CACHE
76
+ if slot is not None and slot[0] == key:
77
+ return slot[1], slot[2], ckpt_path, True
78
+ ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=False)
79
+ model = load_occupancy_model(ckpt, device)
80
+ with _CACHE_LOCK:
81
+ _MODEL_CACHE = (key, model, ckpt)
82
+ return model, ckpt, ckpt_path, False
83
+
84
+
85
+ def _sync_if_cuda(device) -> None:
86
+ """Wait for GPU work so lap times are not just the CPU launch."""
87
+ import torch
88
+
89
+ dev = device if hasattr(device, "type") else torch.device(str(device))
90
+ if str(dev.type) == "cuda":
91
+ torch.cuda.synchronize()
92
+
93
+
94
+ def _round_s(t0: float) -> float:
95
+ return round(time.perf_counter() - t0, 3)
96
+
97
+
98
+ def pred_from_probs(probs: np.ndarray, threshold: float = 0.5) -> np.ndarray:
99
+ """Hard inside labels: 1 iff sigmoid probability is at least ``threshold``."""
100
+ t = float(threshold)
101
+ if not np.isfinite(t):
102
+ t = 0.5
103
+ t = min(1.0, max(0.0, t))
104
+ return (np.asarray(probs, dtype=np.float32).reshape(-1) >= t).astype(np.uint8)
105
+
106
+
107
+ def _sigmoid_probs_and_pred(logits: Any) -> tuple[np.ndarray, np.ndarray]:
108
+ """
109
+ Float32 sigmoid of a 1-D logit tensor, plus a 0.5-cut pred for tests/compat.
110
+
111
+ The inspect page re-cuts from ``prob_b64``; occupancy train metrics stay at 0.5.
112
+ """
113
+ import torch
114
+
115
+ probs = (
116
+ torch.sigmoid(logits.reshape(-1))
117
+ .detach()
118
+ .cpu()
119
+ .numpy()
120
+ .astype(np.float32, copy=False)
121
+ )
122
+ return probs, pred_from_probs(probs, 0.5)
123
+
124
+
125
+ def aabb_for_viewer(
126
+ ckpt: dict[str, Any],
127
+ *,
128
+ npz_name: str,
129
+ mesh_path: str,
130
+ points: np.ndarray,
131
+ data_dir: Path | None,
132
+ vertices: np.ndarray | None = None,
133
+ uploaded_obj: bool = False,
134
+ ) -> tuple[np.ndarray, float, str]:
135
+ """
136
+ Checkpoint AABB when this catalog NPZ/mesh was trained; else mesh vertices.
137
+
138
+ Uploaded OBJ Fill (Job B) always uses this mesh's vertices. Catalog
139
+ ``parts`` matched by filename would apply another cube's box.
140
+ """
141
+ if not uploaded_obj:
142
+ part = match_checkpoint_part(ckpt, npz_name, mesh_path)
143
+ if part is not None:
144
+ center = np.asarray(part["center"], dtype=np.float32).reshape(3)
145
+ scale = float(part["scale"])
146
+ return center, scale, "checkpoint"
147
+ if vertices is not None and int(np.asarray(vertices).shape[0]) > 0:
148
+ center, scale = compute_center_scale(np.asarray(vertices, dtype=np.float32))
149
+ return center, scale, "mesh"
150
+ if mesh_path and data_dir is not None:
151
+ try:
152
+ resolved = resolve_viewer_mesh(mesh_path, data_dir)
153
+ vertices, _faces = load_obj_triangles(resolved)
154
+ center, scale = compute_center_scale(vertices)
155
+ return center, scale, "mesh"
156
+ except (PermissionError, FileNotFoundError, ValueError):
157
+ pass
158
+ center, scale = compute_center_scale(points)
159
+ return center, scale, "points"
160
+
161
+
162
+ def _ckpt_shape_encoder(ckpt: dict[str, Any]) -> str:
163
+ """
164
+ Encoder stored on ``best.pt``, not live YAML.
165
+
166
+ Missing ``shape_encoder`` on an occupancy-encoder checkpoint is the
167
+ original envelope head.
168
+ """
169
+ raw = str(ckpt.get("shape_encoder") or "").strip().lower()
170
+ if raw in ("surface", "none"):
171
+ return raw
172
+ if raw == "mesh":
173
+ raise ValueError(
174
+ "face-token occupancy checkpoints (shape_encoder='mesh') "
175
+ "are no longer supported"
176
+ )
177
+ kind = str(ckpt.get("kind") or "")
178
+ if kind == ENCODER_KIND:
179
+ return "surface"
180
+ return "none"
181
+
182
+
183
+ def _runtime_cfg(cfg: OccupancyConfig | None):
184
+ """Use the caller cfg (tests) or load repo YAML for device / batch / seed."""
185
+ if cfg is not None:
186
+ return cfg
187
+ from scatteringnet.config import load_config
188
+
189
+ return load_config()
190
+
191
+
192
+ def _forward_logits(model, cfg, xyz, geom, shape_id):
193
+ """Batched occupancy logits. ``geom`` is envelope ``(1, N, 3)`` or ``None``."""
194
+ import torch
195
+
196
+ n = int(xyz.shape[0])
197
+ logits_rows: list[Any] = []
198
+ with torch.no_grad():
199
+ for start in range(0, n, int(cfg.batch_size)):
200
+ sl = slice(start, start + int(cfg.batch_size))
201
+ batch_xyz = xyz[sl].to(cfg.device)
202
+ if geom is None:
203
+ logits_rows.append(model(batch_xyz).cpu())
204
+ else:
205
+ b = int(batch_xyz.shape[0])
206
+ geom_b = geom.expand(b, -1, -1).to(cfg.device)
207
+ sid = shape_id.reshape(()).expand(b).to(cfg.device)
208
+ logits_rows.append(model(batch_xyz, geom_b, sid).cpu())
209
+ return torch.cat(logits_rows, dim=0)
210
+
211
+
212
+ def _geom_from_mesh(ckpt, cfg, vertices, faces, center, scale, cache_key: str):
213
+ """
214
+ Build the envelope this checkpoint was trained with.
215
+
216
+ ``surface`` → ``(1, n_surface, C)`` envelope (C from checkpoint).
217
+ ``none`` → no geometry (xyz-only MLP).
218
+ Count comes from the checkpoint. Sampling is always area-weighted
219
+ (old ``envelope_mix`` on ``best.pt`` is ignored).
220
+ """
221
+ import torch
222
+
223
+ enc = _ckpt_shape_encoder(ckpt)
224
+ seed = envelope_seed_from_ckpt(ckpt)
225
+ if enc == "surface":
226
+ n_surface = (
227
+ int(ckpt["n_surface"]) if ckpt.get("n_surface") is not None else 1024
228
+ )
229
+ world = sample_surface_points(
230
+ vertices,
231
+ faces,
232
+ n_surface,
233
+ seed=seed,
234
+ cache_key=cache_key,
235
+ )
236
+ env = apply_envelope_aabb(world, center, scale)
237
+ env = project_envelope_dim(env, envelope_dim_from_ckpt(ckpt))
238
+ geom = torch.from_numpy(env).unsqueeze(0)
239
+ shape_id = torch.zeros(1, dtype=torch.long)
240
+ return geom, shape_id
241
+ return None, None
242
+
243
+
244
+ def decode_points_b64(text: str) -> np.ndarray:
245
+ """Little-endian float32 XYZ from standard base64."""
246
+ raw = base64.b64decode(text)
247
+ pts = np.frombuffer(raw, dtype=np.float32)
248
+ if pts.size % 3 != 0:
249
+ raise ValueError("points buffer length is not a multiple of 3")
250
+ return np.ascontiguousarray(pts.reshape(-1, 3))
251
+
252
+
253
+ def decode_labels_b64(text: str, n: int) -> np.ndarray:
254
+ """uint8 {0,1} labels, length ``n``."""
255
+ raw = base64.b64decode(text)
256
+ labels = np.frombuffer(raw, dtype=np.uint8)
257
+ if int(labels.shape[0]) != int(n):
258
+ raise ValueError(
259
+ f"labels length {labels.shape[0]} does not match points {n}"
260
+ )
261
+ return np.ascontiguousarray(labels)
262
+
263
+
264
+ def infer_uploaded_npz(
265
+ *,
266
+ run_id: str,
267
+ models_root: Path,
268
+ data_dir: Path | None,
269
+ npz_name: str,
270
+ mesh_path: str,
271
+ points: np.ndarray,
272
+ labels: np.ndarray,
273
+ cfg: OccupancyConfig | None = None,
274
+ ) -> dict[str, Any]:
275
+ """
276
+ Classify ``points`` with ``best.pt``. Returns pred bytes (base64) and scores.
277
+
278
+ Envelope rebuild follows ``ckpt['shape_encoder']``.
279
+ """
280
+ import torch
281
+
282
+ n = int(points.shape[0])
283
+ if n < 1:
284
+ raise ValueError("no query points")
285
+ if n > MAX_POINTS:
286
+ raise ValueError(f"too many points ({n}; max {MAX_POINTS})")
287
+ if points.ndim != 2 or points.shape[1] != 3:
288
+ raise ValueError(f"points must be (N, 3), got {tuple(points.shape)}")
289
+ if int(labels.shape[0]) != n:
290
+ raise ValueError("labels length does not match points")
291
+
292
+ t_all = time.perf_counter()
293
+ t0 = time.perf_counter()
294
+ cfg = _runtime_cfg(cfg)
295
+ model, ckpt, ckpt_path, cache_hit = load_cached_occupancy(
296
+ run_id, models_root, cfg.device
297
+ )
298
+ _sync_if_cuda(cfg.device)
299
+ load_s = _round_s(t0)
300
+ center, scale, aabb_src = aabb_for_viewer(
301
+ ckpt,
302
+ npz_name=npz_name,
303
+ mesh_path=mesh_path,
304
+ points=points,
305
+ data_dir=data_dir,
306
+ )
307
+ enc = _ckpt_shape_encoder(ckpt)
308
+ geom = None
309
+ shape_id = None
310
+ parse_s = 0.0
311
+ envelope_s = 0.0
312
+ if enc in ("surface", "mesh"):
313
+ if not mesh_path:
314
+ raise ValueError("this checkpoint needs a mesh_path for the geometry encoder")
315
+ if data_dir is None:
316
+ raise ValueError("helper has no data_dir; cannot load the OBJ for the encoder")
317
+ t0 = time.perf_counter()
318
+ mesh_file = resolve_viewer_mesh(mesh_path, data_dir)
319
+ vertices, faces = load_obj_triangles(mesh_file)
320
+ parse_s = _round_s(t0)
321
+ t0 = time.perf_counter()
322
+ geom, shape_id = _geom_from_mesh(
323
+ ckpt, cfg, vertices, faces, center, scale, str(mesh_file.resolve())
324
+ )
325
+ envelope_s = _round_s(t0)
326
+ if geom is None:
327
+ raise ValueError(f"checkpoint shape_encoder={enc!r} built no geometry tokens")
328
+
329
+ t0 = time.perf_counter()
330
+ xyz = torch.from_numpy(apply_normalization(points, center, scale))
331
+ y = torch.from_numpy(labels.astype(np.float32)).unsqueeze(1)
332
+ logits = _forward_logits(model, cfg, xyz, geom, shape_id)
333
+ _sync_if_cuda(cfg.device)
334
+ forward_s = _round_s(t0)
335
+ scores = occupancy_metrics(logits, y)
336
+ probs, pred = _sigmoid_probs_and_pred(logits)
337
+ timings = {
338
+ "parse_obj": parse_s,
339
+ "load_model": load_s,
340
+ "envelope": envelope_s,
341
+ "forward": forward_s,
342
+ "server_total": _round_s(t_all),
343
+ "load_cached": cache_hit,
344
+ "device": str(cfg.device),
345
+ }
346
+ print(
347
+ "viewer infer-npz timings (s) n=%s parse=%.3f load=%.3f envelope=%.3f "
348
+ "forward=%.3f total=%.3f cached=%s device=%s"
349
+ % (
350
+ n,
351
+ parse_s,
352
+ load_s,
353
+ envelope_s,
354
+ forward_s,
355
+ timings["server_total"],
356
+ cache_hit,
357
+ cfg.device,
358
+ ),
359
+ flush=True,
360
+ )
361
+ gt = labels > 0
362
+ pred_bool = pred > 0
363
+ n_fn = int(np.count_nonzero(gt & ~pred_bool))
364
+ n_fp = int(np.count_nonzero(~gt & pred_bool))
365
+ return {
366
+ "pred_b64": base64.b64encode(np.ascontiguousarray(pred)).decode("ascii"),
367
+ "prob_b64": base64.b64encode(np.ascontiguousarray(probs)).decode("ascii"),
368
+ "n": n,
369
+ "accuracy": float(scores.accuracy),
370
+ "inside_iou": float(scores.inside_iou),
371
+ "inside_f1": float(scores.inside_f1),
372
+ "kind": str(ckpt.get("kind") or ""),
373
+ "shape_encoder": enc,
374
+ "aabb": aabb_src,
375
+ "n_error": n_fn + n_fp,
376
+ "n_fn": n_fn,
377
+ "n_fp": n_fp,
378
+ "run_id": Path(ckpt_path).parent.name,
379
+ "timings": timings,
380
+ }
381
+
382
+
383
+ def infer_uploaded_obj(
384
+ *,
385
+ run_id: str,
386
+ models_root: Path,
387
+ data_dir: Path | None,
388
+ obj_name: str,
389
+ obj_text: str,
390
+ points: np.ndarray,
391
+ cfg: OccupancyConfig | None = None,
392
+ ) -> dict[str, Any]:
393
+ """
394
+ Classify fill-lattice XYZ with ``best.pt``.
395
+
396
+ Envelope tokens are built from the uploaded OBJ when the
397
+ checkpoint is ``surface``. No file labels (Job B): metrics are omitted.
398
+ """
399
+ import torch
400
+
401
+ n = int(points.shape[0])
402
+ if n < 1:
403
+ raise ValueError("no query points; Fill points first")
404
+ if n > MAX_POINTS:
405
+ raise ValueError(f"too many points ({n}; max {MAX_POINTS})")
406
+ if points.ndim != 2 or points.shape[1] != 3:
407
+ raise ValueError(f"points must be (N, 3), got {tuple(points.shape)}")
408
+
409
+ t_all = time.perf_counter()
410
+ t0 = time.perf_counter()
411
+ vertices, faces = triangles_from_obj_text(obj_text)
412
+ parse_s = _round_s(t0)
413
+ t0 = time.perf_counter()
414
+ cfg = _runtime_cfg(cfg)
415
+ model, ckpt, ckpt_path, cache_hit = load_cached_occupancy(
416
+ run_id, models_root, cfg.device
417
+ )
418
+ _sync_if_cuda(cfg.device)
419
+ load_s = _round_s(t0)
420
+ center, scale, aabb_src = aabb_for_viewer(
421
+ ckpt,
422
+ npz_name="",
423
+ mesh_path=str(obj_name or ""),
424
+ points=points,
425
+ data_dir=data_dir,
426
+ vertices=vertices,
427
+ uploaded_obj=True,
428
+ )
429
+ enc = _ckpt_shape_encoder(ckpt)
430
+ t0 = time.perf_counter()
431
+ geom, shape_id = _geom_from_mesh(
432
+ ckpt, cfg, vertices, faces, center, scale, "upload:" + str(obj_name or "obj")
433
+ )
434
+ envelope_s = _round_s(t0)
435
+ if enc in ("surface", "mesh") and geom is None:
436
+ raise ValueError(f"checkpoint shape_encoder={enc!r} built no geometry tokens")
437
+ t0 = time.perf_counter()
438
+ xyz = torch.from_numpy(apply_normalization(points, center, scale))
439
+ logits = _forward_logits(model, cfg, xyz, geom, shape_id)
440
+ _sync_if_cuda(cfg.device)
441
+ forward_s = _round_s(t0)
442
+ probs, pred = _sigmoid_probs_and_pred(logits)
443
+ n_in = int(np.count_nonzero(pred > 0))
444
+ timings = {
445
+ "parse_obj": parse_s,
446
+ "load_model": load_s,
447
+ "envelope": envelope_s,
448
+ "forward": forward_s,
449
+ "server_total": _round_s(t_all),
450
+ "load_cached": cache_hit,
451
+ "device": str(cfg.device),
452
+ }
453
+ print(
454
+ "viewer infer-obj timings (s) n=%s parse=%.3f load=%.3f envelope=%.3f "
455
+ "forward=%.3f total=%.3f cached=%s device=%s"
456
+ % (
457
+ n,
458
+ parse_s,
459
+ load_s,
460
+ envelope_s,
461
+ forward_s,
462
+ timings["server_total"],
463
+ cache_hit,
464
+ cfg.device,
465
+ ),
466
+ flush=True,
467
+ )
468
+ return {
469
+ "pred_b64": base64.b64encode(np.ascontiguousarray(pred)).decode("ascii"),
470
+ "prob_b64": base64.b64encode(np.ascontiguousarray(probs)).decode("ascii"),
471
+ "n": n,
472
+ "n_inside": n_in,
473
+ "n_outside": n - n_in,
474
+ "kind": str(ckpt.get("kind") or ""),
475
+ "shape_encoder": enc,
476
+ "aabb": aabb_src,
477
+ "run_id": Path(ckpt_path).parent.name,
478
+ "timings": timings,
479
+ }
480
+
481
+
482
+ _WARM_OBJ = (
483
+ "v 0 0 0\nv 1 0 0\nv 0 1 0\nv 0 0 1\n"
484
+ "f 1 2 3\nf 1 2 4\nf 1 3 4\nf 2 3 4\n"
485
+ )
486
+
487
+
488
+ def warmup_viewer_helper(models_root: Path) -> None:
489
+ """
490
+ Pay first-click costs before the browser opens: torch, CUDA, trimesh
491
+ fill, load the INSPECT ``best.pt`` (else newest), tiny envelope + forward.
492
+ """
493
+ t0 = time.perf_counter()
494
+ print("warmup: torch / CUDA…", flush=True)
495
+ import torch
496
+
497
+ from scatteringnet.config import load_config
498
+ from scatteringnet.viewer.obj_fill import fill_aabb_lattice
499
+
500
+ cfg = load_config()
501
+ print("warmup: occupancy device=" + str(cfg.device), flush=True)
502
+ if str(cfg.device.type) == "cuda":
503
+ torch.zeros(1, device=cfg.device)
504
+ torch.cuda.synchronize()
505
+ print("warmup: CUDA " + str(cfg.device), flush=True)
506
+ else:
507
+ print("warmup: CPU only", flush=True)
508
+
509
+ print("warmup: tiny fill…", flush=True)
510
+ verts, faces = triangles_from_obj_text(_WARM_OBJ)
511
+ points, _used, _grid = fill_aabb_lattice(verts, 0.40)
512
+ if int(points.shape[0]) < 1:
513
+ points = np.array([[0.2, 0.2, 0.2], [0.8, 0.2, 0.2]], dtype=np.float32)
514
+
515
+ from scatteringnet.viewer.model_access import inspect_run_id, list_viewer_models
516
+
517
+ rows = list_viewer_models(models_root)
518
+ if not rows:
519
+ print("warmup: no models/; skip infer", flush=True)
520
+ print("warmup: done in %.1fs" % (time.perf_counter() - t0), flush=True)
521
+ return
522
+ ids = [str(row["id"]) for row in rows]
523
+ wanted = inspect_run_id()
524
+ # Warm the inspect alias when that folder exists; otherwise newest.
525
+ run_id = wanted if wanted in ids else ids[0]
526
+ print("warmup: infer " + run_id + "…", flush=True)
527
+ infer_uploaded_obj(
528
+ run_id=run_id,
529
+ models_root=models_root,
530
+ data_dir=None,
531
+ obj_name="warmup.obj",
532
+ obj_text=_WARM_OBJ,
533
+ points=points,
534
+ cfg=cfg,
535
+ )
536
+ print("warmup: done in %.1fs" % (time.perf_counter() - t0), flush=True)
scatteringnet/viewer/mesh_access.py ADDED
@@ -0,0 +1,52 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Resolve NPZ ``mesh_path`` for the occupancy viewer helper.
2
+
3
+ Uses :func:`data_npz.resolve_mesh_path` and then requires the file to sit
4
+ under ``data_dir`` (no arbitrary filesystem reads).
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ from pathlib import Path
10
+
11
+ from scatteringnet.data_npz import resolve_mesh_path
12
+
13
+
14
+ def resolve_viewer_mesh(stored: str, data_dir: Path | str) -> Path:
15
+ """
16
+ Return an existing ``.obj`` under ``data_dir`` for a stored mesh_path.
17
+
18
+ Parameters
19
+ ----------
20
+ stored:
21
+ Relative or absolute path from the NPZ ``mesh_path`` array.
22
+ data_dir:
23
+ Dataset root (``config.yaml`` ``data_dir``).
24
+
25
+ Returns
26
+ -------
27
+ Path
28
+ Resolved OBJ path.
29
+
30
+ Raises
31
+ ------
32
+ ValueError
33
+ Empty path or not an OBJ.
34
+ FileNotFoundError
35
+ Mesh does not exist (from :func:`resolve_mesh_path`).
36
+ PermissionError
37
+ Resolved path is outside ``data_dir``.
38
+ """
39
+ text = str(stored).strip()
40
+ if not text:
41
+ raise ValueError("mesh_path is empty")
42
+ root = Path(data_dir).resolve()
43
+ resolved = resolve_mesh_path(text, root)
44
+ try:
45
+ resolved.relative_to(root)
46
+ except ValueError as exc:
47
+ raise PermissionError(
48
+ f"mesh is outside data_dir: {resolved} (stored={text!r})"
49
+ ) from exc
50
+ if resolved.suffix.lower() != ".obj":
51
+ raise ValueError(f"mesh is not an OBJ: {resolved}")
52
+ return resolved
scatteringnet/viewer/model_access.py ADDED
@@ -0,0 +1,183 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """List and resolve ``models/<run_id>/best.pt`` for the occupancy viewer.
2
+
3
+ Read-only: paths must stay under the repo ``models/`` folder.
4
+ Does not import torch.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ from pathlib import Path
10
+
11
+ import yaml
12
+
13
+ # Repo root: this file lives at src/viewer/model_access.py.
14
+ _REPO_ROOT = Path(__file__).resolve().parents[2]
15
+ # Pointers (YAML). Occupancy train / infer math does not read these.
16
+ INSPECT_POINTER = _REPO_ROOT / "docs" / "inspect_checkpoint.yaml"
17
+ HOLDOUT_POINTER = _REPO_ROOT / "docs" / "locked_holdout_objs.yaml"
18
+ # Same id as the committed pointer; used if that file is missing.
19
+ _INSPECT_FALLBACK = "2026-09-14_07-43-34_prim_extruded_nr45_knn24_n2048_n6"
20
+
21
+
22
+ def inspect_run_id() -> str:
23
+ """INSPECT alias: ``run_id`` in ``docs/inspect_checkpoint.yaml``.
24
+
25
+ Viewers use this as the default ``models/<id>/best.pt`` when that file
26
+ exists. Occupancy train / infer math is unchanged.
27
+ """
28
+ try:
29
+ raw = yaml.safe_load(INSPECT_POINTER.read_text(encoding="utf-8"))
30
+ except OSError:
31
+ return _INSPECT_FALLBACK
32
+ if not isinstance(raw, dict):
33
+ return _INSPECT_FALLBACK
34
+ token = str(raw.get("run_id") or "").strip()
35
+ return token[:200] if token else _INSPECT_FALLBACK
36
+
37
+
38
+ def locked_holdout_objs() -> list[str]:
39
+ """OBJ basenames from ``docs/locked_holdout_objs.yaml``.
40
+
41
+ Train / catalog construction does not consult this list. It is the
42
+ locked inspect set for humans and for tests.
43
+ """
44
+ try:
45
+ raw = yaml.safe_load(HOLDOUT_POINTER.read_text(encoding="utf-8"))
46
+ except OSError:
47
+ return []
48
+ if not isinstance(raw, dict):
49
+ return []
50
+ objs = raw.get("objs") or []
51
+ if not isinstance(objs, list):
52
+ return []
53
+ names: list[str] = []
54
+ for item in objs:
55
+ name = str(item or "").strip()
56
+ if name:
57
+ names.append(name)
58
+ return names
59
+
60
+
61
+ def _shape_encoder_from_run(runs_root: Path | None, run_id: str) -> str:
62
+ """Read ``shape_encoder`` from ``runs/<id>/config.yaml`` (no torch)."""
63
+ if runs_root is None:
64
+ return ""
65
+ path = Path(runs_root) / run_id / "config.yaml"
66
+ try:
67
+ if not path.is_file():
68
+ return ""
69
+ for line in path.read_text(encoding="utf-8").splitlines():
70
+ stripped = line.strip()
71
+ if stripped.startswith("#") or not stripped.startswith("shape_encoder:"):
72
+ continue
73
+ raw = stripped.split(":", 1)[1].split("#", 1)[0].strip().strip("\"'")
74
+ kind = raw.lower()
75
+ if kind in ("surface", "mesh", "none"):
76
+ return kind
77
+ return ""
78
+ except OSError:
79
+ return ""
80
+ return ""
81
+
82
+
83
+ def list_viewer_models(
84
+ models_root: Path | str,
85
+ *,
86
+ runs_root: Path | str | None = None,
87
+ ) -> list[dict[str, str | int]]:
88
+ """
89
+ Return ``best.pt`` checkpoints one level under ``models_root``.
90
+
91
+ Each item: ``id`` (folder name), ``path`` (repo-relative), ``mtime``,
92
+ and ``shape_encoder`` when ``runs/<id>/config.yaml`` is present.
93
+ Newest first. Missing folder → empty list.
94
+ """
95
+ root = Path(models_root)
96
+ try:
97
+ root = root.resolve()
98
+ except OSError:
99
+ return []
100
+ if not root.is_dir():
101
+ return []
102
+ runs = Path(runs_root) if runs_root is not None else None
103
+ items: list[dict[str, str | int]] = []
104
+ for best in root.glob("*/best.pt"):
105
+ if not best.is_file():
106
+ continue
107
+ run_id = best.parent.name
108
+ try:
109
+ mtime = int(best.stat().st_mtime)
110
+ except OSError:
111
+ mtime = 0
112
+ enc = _shape_encoder_from_run(runs, run_id)
113
+ row: dict[str, str | int] = {
114
+ "id": run_id,
115
+ "path": "models/" + run_id + "/best.pt",
116
+ "mtime": mtime,
117
+ }
118
+ if enc:
119
+ row["shape_encoder"] = enc
120
+ items.append(row)
121
+ items.sort(key=lambda row: (-int(row["mtime"]), str(row["id"])))
122
+ return items
123
+
124
+
125
+ def resolve_viewer_checkpoint(run_id: str, models_root: Path | str) -> Path:
126
+ """
127
+ Return ``models_root / <run_id> / best.pt``.
128
+
129
+ ``run_id`` is a single folder name (no slashes). Also accepts
130
+ ``models/<run_id>/best.pt`` and strips it down to the folder name.
131
+
132
+ Raises
133
+ ------
134
+ ValueError
135
+ Empty or unsafe id.
136
+ FileNotFoundError
137
+ ``best.pt`` is missing.
138
+ PermissionError
139
+ Resolved path is outside ``models_root``.
140
+ """
141
+ raw = str(run_id or "").strip().replace("\\", "/")
142
+ if raw.endswith("/best.pt"):
143
+ raw = raw[: -len("/best.pt")]
144
+ if raw.startswith("models/"):
145
+ raw = raw[len("models/") :]
146
+ name = raw.strip("/")
147
+ if not name or "/" in name or name in (".", "..") or ".." in name:
148
+ raise ValueError("invalid model id")
149
+ root = Path(models_root).resolve()
150
+ best = (root / name / "best.pt").resolve()
151
+ try:
152
+ best.relative_to(root)
153
+ except ValueError as exc:
154
+ raise PermissionError(
155
+ f"checkpoint is outside models/: {best} (id={run_id!r})"
156
+ ) from exc
157
+ if not best.is_file():
158
+ raise FileNotFoundError(f"best.pt not found for {name!r}")
159
+ return best
160
+
161
+
162
+ def match_checkpoint_part(ckpt: dict, npz_name: str, mesh_path: str) -> dict | None:
163
+ """
164
+ Catalog AABB from ``ckpt['parts']``.
165
+
166
+ NPZ: full stored path or the occupancy filename (those stems are unique).
167
+ Mesh: exact stored path only. Basename matches (``cube.obj``) steal a
168
+ catalog box for an OOD upload of the same name — Job B must not do that.
169
+ """
170
+ parts = ckpt.get("parts") or []
171
+ name = Path(str(npz_name or "")).name
172
+ mesh = str(mesh_path or "").replace("\\", "/").strip()
173
+ if name:
174
+ for part in parts:
175
+ stored = str(part.get("npz") or "").replace("\\", "/")
176
+ if stored == name or Path(stored).name == name:
177
+ return part
178
+ if mesh:
179
+ for part in parts:
180
+ stored_m = str(part.get("mesh") or "").replace("\\", "/").strip()
181
+ if stored_m and stored_m == mesh:
182
+ return part
183
+ return None
scatteringnet/viewer/obj_fill.py ADDED
@@ -0,0 +1,138 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Job B: unlabeled AABB lattice inside an uploaded OBJ (no occupancy GT).
2
+
3
+ Lattice math matches ``scatter_generation.raycast_scatter`` occupancy grid
4
+ (``_padded_bounds`` / ``_uniform_grid_points``) without importing that module
5
+ (Open3D / package ``__init__``). Does not import torch and does not raycast-label.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import base64
11
+ import os
12
+ import tempfile
13
+ from pathlib import Path
14
+
15
+ import numpy as np
16
+
17
+
18
+ from scatteringnet.geometry.mesh_io import load_obj_triangles # noqa: E402
19
+
20
+ MAX_FILL_POINTS = 200_000
21
+ SPACING_MIN = 0.04
22
+ SPACING_MAX = 0.50
23
+ # Slider 0 = coarse (0.40), 100 = dense (0.05). Training occupancy often uses 0.15.
24
+ SPACING_COARSE = 0.40
25
+ SPACING_FINE = 0.05
26
+
27
+
28
+ def spacing_from_slider(value: float) -> float:
29
+ """Map UI density 0–100 to lattice spacing (higher = denser = smaller step)."""
30
+ t = min(1.0, max(0.0, float(value) / 100.0))
31
+ return float(SPACING_COARSE + (SPACING_FINE - SPACING_COARSE) * t)
32
+
33
+
34
+ def clamp_spacing(spacing: float) -> float:
35
+ s = float(spacing)
36
+ if s < SPACING_MIN or s > SPACING_MAX:
37
+ raise ValueError(
38
+ f"spacing must be between {SPACING_MIN} and {SPACING_MAX}, got {s}"
39
+ )
40
+ return s
41
+
42
+
43
+ def _padded_bounds(bounds: np.ndarray, pad: float) -> np.ndarray:
44
+ """Expand AABB by ``pad`` on every side (same as occupancy lattice)."""
45
+ bounds = np.asarray(bounds, dtype=np.float64)
46
+ out = bounds.copy()
47
+ out[0] -= pad
48
+ out[1] += pad
49
+ return out
50
+
51
+
52
+ def _uniform_grid_points(
53
+ bounds: np.ndarray,
54
+ spacing: float,
55
+ *,
56
+ max_points: int = MAX_FILL_POINTS,
57
+ ) -> tuple[np.ndarray, float, tuple[int, int, int]]:
58
+ """Regular XYZ lattice; coarsen spacing by 1.25 until under ``max_points``."""
59
+ if spacing <= 0:
60
+ raise ValueError("point_spacing must be > 0")
61
+ bmin = bounds[0].astype(np.float64)
62
+ bmax = bounds[1].astype(np.float64)
63
+ extents = np.maximum(bmax - bmin, 1e-12)
64
+ used = float(spacing)
65
+
66
+ def counts(step: float) -> tuple[int, int, int]:
67
+ return tuple(max(2, int(np.floor(extents[i] / step)) + 1) for i in range(3))
68
+
69
+ nx, ny, nz = counts(used)
70
+ while nx * ny * nz > max_points:
71
+ used *= 1.25
72
+ nx, ny, nz = counts(used)
73
+
74
+ xs = np.linspace(bmin[0], bmax[0], nx, dtype=np.float64)
75
+ ys = np.linspace(bmin[1], bmax[1], ny, dtype=np.float64)
76
+ zs = np.linspace(bmin[2], bmax[2], nz, dtype=np.float64)
77
+ xx, yy, zz = np.meshgrid(xs, ys, zs, indexing="ij")
78
+ points = np.column_stack([xx.ravel(), yy.ravel(), zz.ravel()])
79
+ return points, used, (nx, ny, nz)
80
+
81
+
82
+ def triangles_from_obj_text(text: str) -> tuple[np.ndarray, np.ndarray]:
83
+ """Parse Wavefront text via a temp file (same loader as training)."""
84
+ raw = str(text or "")
85
+ if not raw.strip():
86
+ raise ValueError("OBJ is empty")
87
+ fd, path = tempfile.mkstemp(suffix=".obj")
88
+ try:
89
+ os.write(fd, raw.encode("utf-8"))
90
+ os.close(fd)
91
+ fd = -1
92
+ return load_obj_triangles(path, cache=False)
93
+ finally:
94
+ if fd >= 0:
95
+ try:
96
+ os.close(fd)
97
+ except OSError:
98
+ pass
99
+ try:
100
+ os.unlink(path)
101
+ except OSError:
102
+ pass
103
+
104
+
105
+ def fill_aabb_lattice(
106
+ vertices: np.ndarray,
107
+ spacing: float,
108
+ *,
109
+ max_points: int = MAX_FILL_POINTS,
110
+ ) -> tuple[np.ndarray, float, tuple[int, int, int]]:
111
+ """
112
+ Regular grid in a padded mesh AABB. Pad = spacing (outside shell), no jitter.
113
+ """
114
+ step = clamp_spacing(spacing)
115
+ verts = np.asarray(vertices, dtype=np.float64)
116
+ if verts.ndim != 2 or verts.shape[1] != 3 or verts.shape[0] < 1:
117
+ raise ValueError(f"vertices must be (V, 3), got {tuple(verts.shape)}")
118
+ bounds = np.stack([verts.min(axis=0), verts.max(axis=0)])
119
+ sample_bounds = _padded_bounds(bounds, step)
120
+ points, used, grid = _uniform_grid_points(
121
+ sample_bounds, step, max_points=int(max_points)
122
+ )
123
+ return np.ascontiguousarray(points, dtype=np.float32), float(used), grid
124
+
125
+
126
+ def fill_from_obj_text(
127
+ obj_text: str,
128
+ spacing: float,
129
+ ) -> dict:
130
+ """Return lattice points (float32) and grid metadata."""
131
+ vertices, _faces = triangles_from_obj_text(obj_text)
132
+ points, used, grid = fill_aabb_lattice(vertices, spacing)
133
+ return {
134
+ "n": int(points.shape[0]),
135
+ "used_spacing": used,
136
+ "grid": [int(grid[0]), int(grid[1]), int(grid[2])],
137
+ "points_b64": base64.b64encode(np.ascontiguousarray(points)).decode("ascii"),
138
+ }
src/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ """Occupancy fill: catalog training, encoder, and inspect helper."""
src/config.py ADDED
@@ -0,0 +1,611 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Hybrid config loader for catalog occupancy training.
2
+
3
+ Static experiment knobs live in ``config.yaml``. ``device`` is resolved
4
+ here from CUDA availability. Training knobs (epochs, lr, batch_size,
5
+ optimizer, catalog, val split) are YAML-owned so ``train_multi_npz``
6
+ does not hardcode them.
7
+
8
+ This module does not read ``.env`` and does not open NPZ files.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import os
14
+ from dataclasses import dataclass
15
+ from pathlib import Path
16
+ from typing import Any, Mapping, TypedDict
17
+
18
+ import torch
19
+ import yaml
20
+
21
+ # Repo root: src/config.py → parents[1].
22
+ _REPO_ROOT = Path(__file__).resolve().parents[1]
23
+ _DEFAULT_YAML = _REPO_ROOT / "config.yaml"
24
+
25
+ _REQUIRED_YAML_KEYS = (
26
+ "hidden",
27
+ "depth",
28
+ "seed",
29
+ "data_dir",
30
+ "epochs",
31
+ "lr",
32
+ )
33
+
34
+
35
+ class YamlKnobs(TypedDict):
36
+ """Subset of OccupancyConfig that is stored in YAML (paths as strings)."""
37
+
38
+ hidden: int
39
+ depth: int
40
+ seed: int
41
+ data_dir: str
42
+ epochs: int
43
+ lr: float
44
+ val_fraction: float
45
+ latent_dim: int | None
46
+ npz_glob: str
47
+ npz_paths: tuple[str, ...]
48
+ npz_catalog: tuple[tuple[str, int | None], ...]
49
+ max_files_per_shape: int | None
50
+ run_name: str
51
+ checkpoint_metric: str
52
+ batch_size: int
53
+ optimizer: str
54
+ n_surface: int
55
+ knn_k: int
56
+ knn_local_dim: int | None
57
+ shape_encoder: str
58
+ # Explicit BCE pos_weight; None when omitted or when auto is set.
59
+ pos_weight: float | None
60
+ pos_weight_auto: bool
61
+
62
+
63
+ def get_device() -> torch.device:
64
+ """
65
+ CUDA when a GPU is visible; otherwise CPU.
66
+
67
+ Hugging Face CPU Spaces have no CUDA: infer still runs (slower).
68
+ ``SCATTERINGNET_DEVICE=cpu`` forces CPU even if a GPU exists.
69
+ ``SCATTERINGNET_DEVICE=cuda`` uses CUDA only when ``is_available()``;
70
+ otherwise it falls back to CPU (no crash).
71
+
72
+ Training scripts should still warn on a long catalog run on CPU.
73
+ """
74
+ forced = os.environ.get("SCATTERINGNET_DEVICE", "").strip().lower()
75
+ if forced == "cpu":
76
+ return torch.device("cpu")
77
+ if forced == "cuda":
78
+ if torch.cuda.is_available():
79
+ return torch.device("cuda")
80
+ return torch.device("cpu")
81
+ if torch.cuda.is_available():
82
+ return torch.device("cuda")
83
+ return torch.device("cpu")
84
+
85
+
86
+ def repo_root() -> Path:
87
+ """Git / project root (folder that contains ``src/`` and ``config.yaml``)."""
88
+ return _REPO_ROOT
89
+
90
+
91
+ def gpu_name(device: torch.device | None = None) -> str | None:
92
+ """
93
+ Human GPU name for the run snapshot (``None`` on CPU).
94
+
95
+ Uses ``cfg.device`` when given so a forced-CPU train does not stamp a
96
+ card that was not used.
97
+ """
98
+ dev = device if device is not None else get_device()
99
+ if dev.type != "cuda" or not torch.cuda.is_available():
100
+ return None
101
+ index = 0 if dev.index is None else int(dev.index)
102
+ if index < 0 or index >= torch.cuda.device_count():
103
+ return None
104
+ name = str(torch.cuda.get_device_name(index)).strip()
105
+ return name or None
106
+
107
+
108
+ def _as_positive_int(name: str, value: Any) -> int:
109
+ """YAML may yield int or (rarely) str; occupancy dims must be int >= 1."""
110
+ try:
111
+ parsed = int(value)
112
+ except (TypeError, ValueError) as exc:
113
+ raise ValueError(f"{name} must be an integer, got {value!r}") from exc
114
+ if parsed < 1:
115
+ raise ValueError(f"{name} must be >= 1, got {parsed}")
116
+ return parsed
117
+
118
+
119
+ def _as_int_in_range(name: str, value: Any, lo: int, hi: int) -> int:
120
+ """Inclusive integer range (YAML ``knn_k`` is 0–4096)."""
121
+ try:
122
+ parsed = int(value)
123
+ except (TypeError, ValueError) as exc:
124
+ raise ValueError(f"{name} must be an integer, got {value!r}") from exc
125
+ if parsed < lo or parsed > hi:
126
+ raise ValueError(f"{name} must be in [{lo}, {hi}], got {parsed}")
127
+ return parsed
128
+
129
+
130
+ def _as_positive_float(name: str, value: Any) -> float:
131
+ """Learning-rate style knobs must be a finite float > 0."""
132
+ try:
133
+ parsed = float(value)
134
+ except (TypeError, ValueError) as exc:
135
+ raise ValueError(f"{name} must be a float, got {value!r}") from exc
136
+ if parsed <= 0.0 or parsed != parsed:
137
+ raise ValueError(f"{name} must be > 0, got {parsed}")
138
+ return parsed
139
+
140
+
141
+ def _as_pos_weight_pair(raw: Mapping[str, Any]) -> tuple[float | None, bool]:
142
+ """
143
+ YAML ``pos_weight``: omit / null → unweighted BCE.
144
+
145
+ ``auto`` → compute n_outside / n_inside on the train split at train time.
146
+ A finite float > 0 is used as-is (1.0 is unweighted).
147
+ """
148
+ if "pos_weight" not in raw:
149
+ return None, False
150
+ value = raw["pos_weight"]
151
+ if value is None or value is False:
152
+ return None, False
153
+ if isinstance(value, str):
154
+ text = value.strip().lower()
155
+ if text in ("", "none", "off", "false"):
156
+ return None, False
157
+ if text == "auto":
158
+ return None, True
159
+ parsed = _as_positive_float("pos_weight", value)
160
+ return parsed, False
161
+
162
+
163
+ def _as_open_unit_interval(name: str, value: Any) -> float:
164
+ """Hold-out fractions must be in (0, 1) so both splits are non-empty."""
165
+ try:
166
+ parsed = float(value)
167
+ except (TypeError, ValueError) as exc:
168
+ raise ValueError(f"{name} must be a float, got {value!r}") from exc
169
+ if parsed != parsed or parsed <= 0.0 or parsed >= 1.0:
170
+ raise ValueError(f"{name} must be in (0, 1), got {parsed}")
171
+ return parsed
172
+
173
+
174
+ def _as_nonempty_path_string(name: str, value: Any) -> str:
175
+ if value is None or (isinstance(value, str) and not value.strip()):
176
+ raise ValueError(f"{name} must be a non-empty path string in config.yaml")
177
+ return str(value).strip()
178
+
179
+
180
+ def _as_data_dir_string(value: Any) -> str:
181
+ return _as_nonempty_path_string("data_dir", value)
182
+
183
+
184
+ def _as_optional_positive_int(name: str, value: Any) -> int | None:
185
+ """YAML null → unlimited catalog cap; otherwise int >= 1."""
186
+ if value is None:
187
+ return None
188
+ return _as_positive_int(name, value)
189
+
190
+
191
+ def _as_run_name(value: Any) -> str:
192
+ """Optional YAML suffix for ``runs/<timestamp>_<name>/``; empty → ``run``."""
193
+ if value is None:
194
+ return "run"
195
+ text = str(value).strip()
196
+ return text if text else "run"
197
+
198
+
199
+ _CHECKPOINT_METRIC_ALIASES = {
200
+ "test_acc": "val_acc",
201
+ "test_iou": "val_iou",
202
+ }
203
+
204
+
205
+ def _as_checkpoint_metric(value: Any) -> str:
206
+ """Name of the scalar used to decide ``best.pt`` (strict improve)."""
207
+ if value is None or (isinstance(value, str) and not value.strip()):
208
+ raise ValueError("checkpoint_metric must be a non-empty string")
209
+ name = str(value).strip()
210
+ return _CHECKPOINT_METRIC_ALIASES.get(name, name)
211
+
212
+
213
+ def _as_val_fraction(raw: Mapping[str, Any]) -> float:
214
+ """Prefer ``val_fraction``; accept legacy ``test_fraction``."""
215
+ if "val_fraction" in raw:
216
+ return _as_open_unit_interval("val_fraction", raw["val_fraction"])
217
+ if "test_fraction" in raw:
218
+ return _as_open_unit_interval("test_fraction", raw["test_fraction"])
219
+ raise ValueError("Config YAML missing keys: val_fraction")
220
+
221
+
222
+ def _as_optional_latent_dim(raw: Mapping[str, Any]) -> int | None:
223
+ """YAML omit / null → use ``hidden`` at train time."""
224
+ if "latent_dim" not in raw or raw["latent_dim"] is None:
225
+ return None
226
+ return _as_positive_int("latent_dim", raw["latent_dim"])
227
+
228
+
229
+ # Names accepted in config.yaml ``optimizer``. Used by train_multi_npz.
230
+ _ALLOWED_OPTIMIZERS = ("adam", "adamw", "sgd")
231
+ # ``none`` = OccupancyMLP (xyz only). ``surface`` = envelope encoder.
232
+ _ALLOWED_SHAPE_ENCODERS = ("none", "surface")
233
+
234
+
235
+ def _as_optimizer(value: Any) -> str:
236
+ """Optimizer family for multi-NPZ train; default Adam."""
237
+ if value is None or (isinstance(value, str) and not str(value).strip()):
238
+ return "adam"
239
+ name = str(value).strip().lower()
240
+ if name not in _ALLOWED_OPTIMIZERS:
241
+ allowed = ", ".join(_ALLOWED_OPTIMIZERS)
242
+ raise ValueError(f"optimizer must be one of {allowed}, got {value!r}")
243
+ return name
244
+
245
+
246
+ def _as_shape_encoder(value: Any) -> str:
247
+ """Occupancy head family; default xyz-only so older YAML still loads."""
248
+ if value is None or (isinstance(value, str) and not str(value).strip()):
249
+ return "none"
250
+ name = str(value).strip().lower()
251
+ if name not in _ALLOWED_SHAPE_ENCODERS:
252
+ allowed = ", ".join(_ALLOWED_SHAPE_ENCODERS)
253
+ raise ValueError(f"shape_encoder must be one of {allowed}, got {value!r}")
254
+ return name
255
+
256
+
257
+ def _as_npz_paths(value: Any) -> tuple[str, ...]:
258
+ """Explicit NPZ list relative to data_dir (empty → use glob)."""
259
+ if value is None:
260
+ return ()
261
+ if isinstance(value, str):
262
+ item = value.strip()
263
+ return (item,) if item else ()
264
+ if not isinstance(value, list):
265
+ raise ValueError(f"npz_paths must be a list of strings or null, got {type(value).__name__}")
266
+ out: list[str] = []
267
+ for i, raw in enumerate(value):
268
+ text = _as_nonempty_path_string(f"npz_paths[{i}]", raw)
269
+ out.append(text)
270
+ return tuple(out)
271
+
272
+
273
+ def _as_npz_catalog(value: Any) -> tuple[tuple[str, int | None], ...]:
274
+ """Union of globs; optional ``max_shapes`` is unique meshes per glob."""
275
+ if value is None:
276
+ return ()
277
+ if not isinstance(value, list):
278
+ raise ValueError(
279
+ f"npz_catalog must be a list or null, got {type(value).__name__}"
280
+ )
281
+ out: list[tuple[str, int | None]] = []
282
+ for i, raw in enumerate(value):
283
+ if isinstance(raw, str):
284
+ glob_s = _as_nonempty_path_string(f"npz_catalog[{i}]", raw)
285
+ out.append((glob_s, None))
286
+ continue
287
+ if not isinstance(raw, Mapping):
288
+ raise ValueError(
289
+ f"npz_catalog[{i}] must be a glob string or mapping, "
290
+ f"got {type(raw).__name__}"
291
+ )
292
+ if "glob" not in raw:
293
+ raise ValueError(f"npz_catalog[{i}] missing glob")
294
+ glob_s = _as_nonempty_path_string(f"npz_catalog[{i}].glob", raw["glob"])
295
+ max_shapes: int | None = None
296
+ if "max_shapes" in raw and raw["max_shapes"] is not None:
297
+ max_shapes = _as_positive_int(
298
+ f"npz_catalog[{i}].max_shapes", raw["max_shapes"]
299
+ )
300
+ out.append((glob_s, max_shapes))
301
+ return tuple(out)
302
+
303
+
304
+ def as_repo_relative(path: Path | str, *, root: Path | None = None) -> str:
305
+ """
306
+ POSIX string relative to the git repo when ``path`` is inside it.
307
+
308
+ Already-relative inputs are returned as POSIX. Absolute paths on another
309
+ drive (the dataset disk) cannot be repo-relative and stay absolute POSIX.
310
+ """
311
+ text = str(path).strip()
312
+ if not text:
313
+ return text
314
+ parsed = Path(text)
315
+ if not parsed.is_absolute():
316
+ return parsed.as_posix()
317
+ base = (root or _REPO_ROOT).resolve()
318
+ try:
319
+ return parsed.resolve().relative_to(base).as_posix()
320
+ except ValueError:
321
+ return parsed.resolve().as_posix()
322
+
323
+
324
+ def as_data_relative(path: Path | str, data_dir: Path | str) -> str:
325
+ """
326
+ POSIX string relative to ``data_dir`` (``exports/...``, not ``E:/...``).
327
+
328
+ Already-relative inputs are returned as POSIX. Paths outside ``data_dir``
329
+ (unit-test temp trees) fall back to absolute POSIX.
330
+ """
331
+ text = str(path).strip()
332
+ if not text:
333
+ return text
334
+ parsed = Path(text)
335
+ if not parsed.is_absolute():
336
+ return parsed.as_posix()
337
+ root = Path(data_dir).expanduser().resolve()
338
+ resolved = parsed.expanduser().resolve()
339
+ try:
340
+ return resolved.relative_to(root).as_posix()
341
+ except ValueError:
342
+ return resolved.as_posix()
343
+
344
+
345
+ def load_yaml_knobs(path: Path) -> YamlKnobs:
346
+ """
347
+ Read YAML settings. Does not check that data_dir exists on disk.
348
+
349
+ Parameters
350
+ ----------
351
+ path:
352
+ Path to ``config.yaml``.
353
+
354
+ Returns
355
+ -------
356
+ YamlKnobs
357
+ Typed dict of experiment knobs (paths still strings).
358
+ """
359
+ if not path.is_file():
360
+ raise FileNotFoundError(f"Config YAML not found: {path}")
361
+ raw = yaml.safe_load(path.read_text(encoding="utf-8"))
362
+ if not isinstance(raw, Mapping):
363
+ raise ValueError(f"Config YAML must be a mapping, got {type(raw).__name__}")
364
+ missing = [k for k in _REQUIRED_YAML_KEYS if k not in raw]
365
+ if missing:
366
+ raise ValueError(f"Config YAML missing keys: {', '.join(missing)}")
367
+ pos_weight, pos_weight_auto = _as_pos_weight_pair(raw)
368
+ return {
369
+ "hidden": _as_positive_int("hidden", raw["hidden"]),
370
+ "depth": _as_positive_int("depth", raw["depth"]),
371
+ "seed": _as_positive_int("seed", raw["seed"]),
372
+ "data_dir": _as_data_dir_string(raw["data_dir"]),
373
+ "epochs": _as_positive_int("epochs", raw["epochs"]),
374
+ "lr": _as_positive_float("lr", raw["lr"]),
375
+ "val_fraction": _as_val_fraction(raw),
376
+ "latent_dim": _as_optional_latent_dim(raw),
377
+ "npz_glob": (
378
+ _as_nonempty_path_string("npz_glob", raw["npz_glob"])
379
+ if "npz_glob" in raw
380
+ else "exports/dataset/*.npz"
381
+ ),
382
+ "npz_paths": _as_npz_paths(raw.get("npz_paths")),
383
+ "npz_catalog": _as_npz_catalog(raw.get("npz_catalog")),
384
+ "max_files_per_shape": (
385
+ _as_optional_positive_int("max_files_per_shape", raw["max_files_per_shape"])
386
+ if "max_files_per_shape" in raw
387
+ else 2
388
+ ),
389
+ "run_name": (
390
+ _as_run_name(raw["run_name"]) if "run_name" in raw else "run"
391
+ ),
392
+ "checkpoint_metric": (
393
+ _as_checkpoint_metric(raw["checkpoint_metric"])
394
+ if "checkpoint_metric" in raw
395
+ else "val_acc"
396
+ ),
397
+ "batch_size": (
398
+ _as_positive_int("batch_size", raw["batch_size"])
399
+ if "batch_size" in raw
400
+ else 1024
401
+ ),
402
+ "optimizer": (
403
+ _as_optimizer(raw["optimizer"]) if "optimizer" in raw else "adam"
404
+ ),
405
+ "n_surface": (
406
+ _as_positive_int("n_surface", raw["n_surface"])
407
+ if "n_surface" in raw
408
+ else 1024
409
+ ),
410
+ "knn_k": (
411
+ _as_int_in_range("knn_k", raw["knn_k"], 0, 4096)
412
+ if "knn_k" in raw
413
+ else 0
414
+ ),
415
+ "knn_local_dim": (
416
+ _as_positive_int("knn_local_dim", raw["knn_local_dim"])
417
+ if "knn_local_dim" in raw
418
+ else None
419
+ ),
420
+ "shape_encoder": (
421
+ _as_shape_encoder(raw["shape_encoder"])
422
+ if "shape_encoder" in raw
423
+ else "none"
424
+ ),
425
+ "pos_weight": pos_weight,
426
+ "pos_weight_auto": pos_weight_auto,
427
+ }
428
+
429
+
430
+ def _warn_missing_data_dir(data_dir: Path, yaml_path: Path) -> None:
431
+ """Print a terminal hint so the user can fix config.yaml (no .env involved)."""
432
+ print(
433
+ "\n"
434
+ "Dataset folder not found.\n"
435
+ f" Looked for: {data_dir}\n"
436
+ "\n"
437
+ "Update `data_dir` in config.yaml to the folder that contains "
438
+ "`exports/` and `meshes/`.\n"
439
+ f" Config file: {yaml_path}\n"
440
+ )
441
+
442
+
443
+ def require_data_dir(data_dir: Path, *, yaml_path: Path) -> None:
444
+ """Validate the dataset root before training or NPZ loading."""
445
+ if data_dir.is_dir():
446
+ return
447
+ _warn_missing_data_dir(data_dir, yaml_path)
448
+ raise FileNotFoundError(
449
+ f"Dataset directory does not exist: {data_dir}. "
450
+ f"Set data_dir in {yaml_path}."
451
+ )
452
+
453
+
454
+ @dataclass(frozen=True)
455
+ class OccupancyConfig:
456
+ """Resolved experiment settings from YAML plus detected device."""
457
+
458
+ data_dir: Path
459
+ device: torch.device
460
+ hidden: int
461
+ depth: int
462
+ seed: int
463
+ epochs: int
464
+ lr: float
465
+ # Fraction of catalog **meshes** held out as val (selection split, not a locked test).
466
+ val_fraction: float
467
+ # Catalog knobs (optional in YAML; omitted keys keep these defaults).
468
+ npz_glob: str = "exports/dataset/*.npz"
469
+ npz_paths: tuple[str, ...] = ()
470
+ # ``(glob, max_shapes)`` rows. Empty → use ``npz_glob``. ``max_shapes``
471
+ # None keeps every mesh that glob hits (after ``max_files_per_shape``).
472
+ npz_catalog: tuple[tuple[str, int | None], ...] = ()
473
+ max_files_per_shape: int | None = 2
474
+ # Suffix for runs/<timestamp>_<name>/ (device stays runtime-only).
475
+ run_name: str = "run"
476
+ # Which logged scalar selects best.pt (strict improve).
477
+ checkpoint_metric: str = "val_acc"
478
+ # Encoder latent width; None → use ``hidden`` at train / infer time.
479
+ latent_dim: int | None = None
480
+ # Mini-batch size and optimizer family (train_multi_npz).
481
+ batch_size: int = 1024
482
+ optimizer: str = "adam"
483
+ # Envelope sample count (YAML). Used when ``shape_encoder`` is ``surface``.
484
+ n_surface: int = 1024
485
+ # 0 = global envelope z only. >0 = that many nearest envelope dots per query.
486
+ knn_k: int = 0
487
+ # Width of z_local; None → same as occupancy latent_dim / hidden.
488
+ knn_local_dim: int | None = None
489
+ # ``none`` keeps OccupancyMLP; ``surface`` uses the envelope PointNet.
490
+ shape_encoder: str = "none"
491
+ # BCE inside-class weight. None + auto=False = unweighted (legacy).
492
+ pos_weight: float | None = None
493
+ pos_weight_auto: bool = False
494
+
495
+
496
+ def load_config(
497
+ yaml_path: Path | None = None,
498
+ *,
499
+ require_existing_data_dir: bool = True,
500
+ ) -> OccupancyConfig:
501
+ """
502
+ Compose OccupancyConfig from ``config.yaml`` (not from ``.env``).
503
+
504
+ When ``require_existing_data_dir`` is True (default), a missing folder
505
+ prints a short instruction and then raises FileNotFoundError — used for
506
+ training and data loading.
507
+
508
+ Parameters
509
+ ----------
510
+ yaml_path:
511
+ Config file; default is repo-root ``config.yaml``.
512
+ require_existing_data_dir:
513
+ If True, refuse to return a config whose ``data_dir`` is missing.
514
+
515
+ Returns
516
+ -------
517
+ OccupancyConfig
518
+ YAML knobs plus detected ``device``.
519
+ """
520
+ cfg_path = yaml_path or _DEFAULT_YAML
521
+ knobs = load_yaml_knobs(cfg_path)
522
+ data_dir = Path(knobs["data_dir"])
523
+ if require_existing_data_dir:
524
+ require_data_dir(data_dir, yaml_path=cfg_path)
525
+ return OccupancyConfig(
526
+ data_dir=data_dir,
527
+ device=get_device(),
528
+ hidden=knobs["hidden"],
529
+ depth=knobs["depth"],
530
+ seed=knobs["seed"],
531
+ epochs=knobs["epochs"],
532
+ lr=knobs["lr"],
533
+ val_fraction=knobs["val_fraction"],
534
+ latent_dim=knobs["latent_dim"],
535
+ npz_glob=knobs["npz_glob"],
536
+ npz_paths=knobs["npz_paths"],
537
+ npz_catalog=knobs["npz_catalog"],
538
+ max_files_per_shape=knobs["max_files_per_shape"],
539
+ run_name=knobs["run_name"],
540
+ checkpoint_metric=knobs["checkpoint_metric"],
541
+ batch_size=knobs["batch_size"],
542
+ optimizer=knobs["optimizer"],
543
+ n_surface=knobs["n_surface"],
544
+ knn_k=knobs["knn_k"],
545
+ knn_local_dim=knobs["knn_local_dim"],
546
+ shape_encoder=knobs["shape_encoder"],
547
+ pos_weight=knobs["pos_weight"],
548
+ pos_weight_auto=knobs["pos_weight_auto"],
549
+ )
550
+
551
+
552
+ def encoder_latent_dim(cfg: OccupancyConfig) -> int:
553
+ """OccupancyEncoder ``z`` width: YAML ``latent_dim`` or ``hidden``."""
554
+ if cfg.latent_dim is None:
555
+ return int(cfg.hidden)
556
+ return int(cfg.latent_dim)
557
+
558
+
559
+ def encoder_knn_local_dim(cfg: OccupancyConfig) -> int:
560
+ """Local envelope code width: YAML ``knn_local_dim`` or global latent."""
561
+ if cfg.knn_local_dim is None:
562
+ return encoder_latent_dim(cfg)
563
+ return int(cfg.knn_local_dim)
564
+
565
+
566
+ def format_config(cfg: OccupancyConfig) -> str:
567
+ """
568
+ Pretty-print for CLI smoke checks.
569
+
570
+ Parameters
571
+ ----------
572
+ cfg:
573
+ Resolved config.
574
+
575
+ Returns
576
+ -------
577
+ str
578
+ Multi-line ``OccupancyConfig(...)`` dump.
579
+ """
580
+ return (
581
+ f"OccupancyConfig(\n"
582
+ f" data_dir={cfg.data_dir}\n"
583
+ f" device={cfg.device}\n"
584
+ f" gpu={gpu_name(cfg.device)}\n"
585
+ f" hidden={cfg.hidden}\n"
586
+ f" depth={cfg.depth}\n"
587
+ f" seed={cfg.seed}\n"
588
+ f" epochs={cfg.epochs}\n"
589
+ f" lr={cfg.lr}\n"
590
+ f" val_fraction={cfg.val_fraction}\n"
591
+ f" latent_dim={cfg.latent_dim}\n"
592
+ f" npz_glob={cfg.npz_glob}\n"
593
+ f" npz_paths={list(cfg.npz_paths)}\n"
594
+ f" npz_catalog={list(cfg.npz_catalog)}\n"
595
+ f" max_files_per_shape={cfg.max_files_per_shape}\n"
596
+ f" run_name={cfg.run_name}\n"
597
+ f" checkpoint_metric={cfg.checkpoint_metric}\n"
598
+ f" batch_size={cfg.batch_size}\n"
599
+ f" optimizer={cfg.optimizer}\n"
600
+ f" n_surface={cfg.n_surface}\n"
601
+ f" knn_k={cfg.knn_k}\n"
602
+ f" knn_local_dim={cfg.knn_local_dim}\n"
603
+ f" shape_encoder={cfg.shape_encoder}\n"
604
+ f" pos_weight={cfg.pos_weight}\n"
605
+ f" pos_weight_auto={cfg.pos_weight_auto}\n"
606
+ f")"
607
+ )
608
+
609
+
610
+ if __name__ == "__main__":
611
+ print(format_config(load_config()))
src/data_npz.py ADDED
@@ -0,0 +1,460 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Load occupancy query points and labels from scatter NPZ files.
2
+
3
+ :func:`load_points_labels` reads **one** file. A catalog resolver lists
4
+ many NPZs (glob or explicit paths) without training.
5
+
6
+ ``load_points_labels`` still returns only ``points`` and ``labels``.
7
+ ``load_points_labels_mesh`` also resolves ``mesh_path`` against ``data_dir``.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ import glob as globlib
13
+ import random
14
+ from pathlib import Path
15
+ from typing import Sequence
16
+
17
+ import numpy as np
18
+ from numpy.typing import NDArray
19
+
20
+ # Labels are converted here (not deferred to the Dataset) so every caller gets
21
+ # the same dtypes: float32 XYZ and float32 {0, 1} occupancy.
22
+ PointsArray = NDArray[np.float32]
23
+ LabelsArray = NDArray[np.float32]
24
+
25
+
26
+ def load_points_labels(path: Path) -> tuple[PointsArray, LabelsArray]:
27
+ """
28
+ Read query coordinates and inside/outside labels from one NPZ file.
29
+
30
+ Parameters
31
+ ----------
32
+ path:
33
+ Path to a ``.npz`` with arrays ``points`` ``(N, 3)`` and
34
+ ``labels`` ``(N,)`` (typically uint8 0/1).
35
+
36
+ Returns
37
+ -------
38
+ points:
39
+ ``float32`` array of shape ``(N, 3)``.
40
+ labels:
41
+ ``float32`` array of shape ``(N,)`` with values in ``{0.0, 1.0}``
42
+ (0 = outside, 1 = inside).
43
+ """
44
+ npz_path = Path(path)
45
+ if not npz_path.is_file():
46
+ raise FileNotFoundError(f"NPZ not found: {npz_path}")
47
+
48
+ # allow_pickle=False: we only need numeric arrays, not object payloads.
49
+ with np.load(npz_path, allow_pickle=False) as raw:
50
+ files = set(raw.files)
51
+ if "points" not in files or "labels" not in files:
52
+ raise KeyError(
53
+ f"NPZ must contain 'points' and 'labels', got {sorted(files)} "
54
+ f"in {npz_path}"
55
+ )
56
+ points = np.asarray(raw["points"])
57
+ labels = np.asarray(raw["labels"])
58
+
59
+ if points.ndim != 2 or points.shape[1] != 3:
60
+ raise ValueError(
61
+ f"points must have shape (N, 3), got {tuple(points.shape)} in {npz_path}"
62
+ )
63
+ n = int(points.shape[0])
64
+ if labels.shape != (n,):
65
+ raise ValueError(
66
+ f"labels must have shape (N,) with N={n}, got {tuple(labels.shape)} "
67
+ f"in {npz_path}"
68
+ )
69
+
70
+ points_f32 = np.asarray(points, dtype=np.float32)
71
+ labels_f32 = np.asarray(labels, dtype=np.float32)
72
+ unique = np.unique(labels_f32)
73
+ if not np.all((unique == 0.0) | (unique == 1.0)):
74
+ raise ValueError(
75
+ f"labels must be in {{0, 1}}, got unique={unique.tolist()} in {npz_path}"
76
+ )
77
+ return points_f32, labels_f32
78
+
79
+
80
+ def count_npz_points(path: Path | str) -> int:
81
+ """
82
+ Query count in one NPZ without keeping the arrays.
83
+
84
+ Catalog construct uses this so ``len(part)`` / ``n_points`` do not
85
+ require loading every lattice into RAM.
86
+ """
87
+ npz_path = Path(path)
88
+ if not npz_path.is_file():
89
+ raise FileNotFoundError(f"NPZ not found: {npz_path}")
90
+ with np.load(npz_path, allow_pickle=False) as raw:
91
+ files = set(raw.files)
92
+ if "points" not in files or "labels" not in files:
93
+ raise KeyError(
94
+ f"NPZ must contain 'points' and 'labels', got {sorted(files)} "
95
+ f"in {npz_path}"
96
+ )
97
+ n = int(np.asarray(raw["points"]).shape[0])
98
+ n_y = int(np.asarray(raw["labels"]).shape[0])
99
+ if n_y != n:
100
+ raise ValueError(
101
+ f"labels must have shape (N,) with N={n}, got N={n_y} in {npz_path}"
102
+ )
103
+ return n
104
+
105
+
106
+ def count_npz_labels(path: Path | str) -> tuple[int, int]:
107
+ """
108
+ Outside / inside counts in one NPZ without keeping the point cloud.
109
+
110
+ Used for ``pos_weight: auto`` (n_outside / n_inside on the train split).
111
+ """
112
+ npz_path = Path(path)
113
+ if not npz_path.is_file():
114
+ raise FileNotFoundError(f"NPZ not found: {npz_path}")
115
+ with np.load(npz_path, allow_pickle=False) as raw:
116
+ if "labels" not in set(raw.files):
117
+ raise KeyError(f"NPZ must contain 'labels', got {sorted(raw.files)} in {npz_path}")
118
+ labels = np.asarray(raw["labels"]).reshape(-1)
119
+ # Labels are 0/1 occupancy; nonzero is inside.
120
+ n_in = int(np.count_nonzero(labels))
121
+ n_out = int(labels.size) - n_in
122
+ return n_out, n_in
123
+
124
+
125
+ def read_npz_mesh_path(path: Path | str) -> str:
126
+ """
127
+ Read the stored ``mesh_path`` string from one occupancy NPZ.
128
+
129
+ Step 2 writes a ``data_dir``-relative POSIX path (for example
130
+ ``meshes/Primitives/Sphere/sphere_r0p5_sa16_sh16.obj``).
131
+
132
+ Parameters
133
+ ----------
134
+ path:
135
+ Occupancy ``.npz`` that contains ``mesh_path``.
136
+
137
+ Returns
138
+ -------
139
+ str
140
+ Stored path string (relative or absolute). Not resolved here.
141
+ """
142
+ npz_path = Path(path)
143
+ if not npz_path.is_file():
144
+ raise FileNotFoundError(f"NPZ not found: {npz_path}")
145
+ # allow_pickle=True: some exports store a 0-d string / object array.
146
+ with np.load(npz_path, allow_pickle=True) as raw:
147
+ if "mesh_path" not in raw.files:
148
+ raise KeyError(f"NPZ has no 'mesh_path' in {npz_path}")
149
+ stored = str(np.asarray(raw["mesh_path"]).item()).strip()
150
+ if not stored:
151
+ raise ValueError(f"mesh_path is empty in {npz_path}")
152
+ return stored
153
+
154
+
155
+ def resolve_mesh_path(stored: str, data_dir: Path | str) -> Path:
156
+ """
157
+ Resolve a stored ``mesh_path`` against ``data_dir``.
158
+
159
+ Relative entries are joined to ``data_dir``. Absolute entries are
160
+ used as-is. Missing files raise ``FileNotFoundError``.
161
+
162
+ Parameters
163
+ ----------
164
+ stored:
165
+ Value from :func:`read_npz_mesh_path`.
166
+ data_dir:
167
+ Dataset root (``config.yaml`` ``data_dir``).
168
+
169
+ Returns
170
+ -------
171
+ Path
172
+ Existing resolved mesh file.
173
+ """
174
+ text = str(stored).strip()
175
+ if not text:
176
+ raise ValueError("mesh_path is empty")
177
+ item = Path(text)
178
+ root = Path(data_dir)
179
+ resolved = item if item.is_absolute() else (root / item)
180
+ resolved = resolved.resolve()
181
+ if not resolved.is_file():
182
+ raise FileNotFoundError(f"mesh not found: {resolved} (stored={text!r})")
183
+ return resolved
184
+
185
+
186
+ def load_points_labels_mesh(
187
+ path: Path | str,
188
+ data_dir: Path | str,
189
+ ) -> tuple[PointsArray, LabelsArray, Path]:
190
+ """
191
+ Read occupancy arrays and resolve the source OBJ.
192
+
193
+ Keeps :func:`load_points_labels` unchanged (xyz + labels only).
194
+
195
+ Parameters
196
+ ----------
197
+ path:
198
+ Occupancy ``.npz`` with ``points``, ``labels``, and ``mesh_path``.
199
+ data_dir:
200
+ Root used to resolve a relative ``mesh_path``.
201
+
202
+ Returns
203
+ -------
204
+ points, labels, mesh_path:
205
+ Same arrays as :func:`load_points_labels`, plus the existing OBJ.
206
+ """
207
+ npz_path = Path(path)
208
+ points, labels = load_points_labels(npz_path)
209
+ stored = read_npz_mesh_path(npz_path)
210
+ mesh_path = resolve_mesh_path(stored, data_dir)
211
+ return points, labels, mesh_path
212
+
213
+
214
+ def _is_combo_npz(path: Path) -> bool:
215
+ """True when the filename looks like a combo dump (excluded from the catalog)."""
216
+ return "combo" in path.name.lower()
217
+
218
+
219
+ def shape_key(path: Path) -> str:
220
+ """
221
+ Group NPZs that belong to the same mesh.
222
+
223
+ Uses the stem before ``__`` (dataset_builder tag), else the full stem.
224
+
225
+ Parameters
226
+ ----------
227
+ path:
228
+ NPZ path.
229
+
230
+ Returns
231
+ -------
232
+ str
233
+ Stable key for ``max_files_per_shape``.
234
+ """
235
+ stem = Path(path).stem
236
+ if "__" in stem:
237
+ return stem.split("__", 1)[0]
238
+ return stem
239
+
240
+
241
+ def _cap_per_shape(
242
+ paths: Sequence[Path],
243
+ max_files_per_shape: int | None,
244
+ ) -> list[Path]:
245
+ """Keep at most ``max_files_per_shape`` files per :func:`shape_key` (sorted order)."""
246
+ if max_files_per_shape is None:
247
+ return list(paths)
248
+ if max_files_per_shape < 1:
249
+ raise ValueError(f"max_files_per_shape must be >= 1 or None, got {max_files_per_shape}")
250
+ counts: dict[str, int] = {}
251
+ out: list[Path] = []
252
+ for path in paths:
253
+ key = shape_key(path)
254
+ taken = counts.get(key, 0)
255
+ if taken >= max_files_per_shape:
256
+ continue
257
+ counts[key] = taken + 1
258
+ out.append(path)
259
+ return out
260
+
261
+
262
+ def _glob_npz(root: Path, pattern: str) -> list[Path]:
263
+ """Match ``pattern`` under ``root`` (``*`` / ``**``)."""
264
+ full = str(root / pattern)
265
+ recursive = "**" in pattern.replace("\\", "/")
266
+ found = globlib.glob(full, recursive=recursive)
267
+ return [Path(p).resolve() for p in found if Path(p).is_file()]
268
+
269
+
270
+ def _is_parameterized_stem(path: Path) -> bool:
271
+ """
272
+ Maya catalog names are ``family_param_...``. Varied one-off stems
273
+ (``Cone.obj`` → ``Cone__occupancy.npz``) have no ``_`` in the shape key
274
+ and must not ride along when Windows glob is case-insensitive.
275
+ """
276
+ return "_" in shape_key(path)
277
+
278
+
279
+ def _subsample_shapes(
280
+ paths: Sequence[Path],
281
+ max_shapes: int,
282
+ seed: int,
283
+ ) -> list[Path]:
284
+ """Keep NPZs for at most ``max_shapes`` unique :func:`shape_key` values."""
285
+ if max_shapes < 1:
286
+ raise ValueError(f"max_shapes must be >= 1, got {max_shapes}")
287
+ keys: list[str] = []
288
+ seen: set[str] = set()
289
+ for path in paths:
290
+ key = shape_key(path)
291
+ if key in seen:
292
+ continue
293
+ seen.add(key)
294
+ keys.append(key)
295
+ if max_shapes >= len(keys):
296
+ return list(paths)
297
+ # Sort then sample so the same seed always picks the same meshes.
298
+ chosen_keys = set(random.Random(int(seed)).sample(sorted(keys), max_shapes))
299
+ return [path for path in paths if shape_key(path) in chosen_keys]
300
+
301
+
302
+ def resolve_npz_catalog(
303
+ data_dir: Path | str,
304
+ *,
305
+ npz_glob: str = "exports/dataset/*.npz",
306
+ npz_paths: Sequence[str | Path] | None = None,
307
+ npz_catalog: Sequence[tuple[str, int | None]] | None = None,
308
+ max_files_per_shape: int | None = 2,
309
+ exclude_combo: bool = True,
310
+ seed: int = 1,
311
+ ) -> list[Path]:
312
+ """
313
+ Resolve occupancy NPZ paths under ``data_dir`` (no point loading).
314
+
315
+ Priority: explicit ``npz_paths``, else ``npz_catalog`` (union of globs),
316
+ else ``npz_glob``. Relative entries are joined to ``data_dir``.
317
+ Missing files in ``npz_paths`` raise ``FileNotFoundError``.
318
+
319
+ Parameters
320
+ ----------
321
+ data_dir:
322
+ Dataset root (``config.yaml`` ``data_dir``).
323
+ npz_glob:
324
+ Single glob relative to ``data_dir`` (``*`` and ``**`` allowed).
325
+ npz_paths:
326
+ Explicit relative or absolute NPZ paths. Empty / None → use glob(s).
327
+ npz_catalog:
328
+ ``(glob, max_shapes)`` rows. ``max_shapes`` is unique meshes after
329
+ the per-shape file cap; ``None`` keeps every mesh the glob hits.
330
+ max_files_per_shape:
331
+ Cap per :func:`shape_key` after sort. ``None`` = no cap.
332
+ exclude_combo:
333
+ Drop filenames containing ``combo``.
334
+ seed:
335
+ RNG for ``max_shapes`` subsampling (YAML ``seed``).
336
+
337
+ Returns
338
+ -------
339
+ list[Path]
340
+ Sorted existing ``.npz`` files.
341
+ """
342
+ root = Path(data_dir)
343
+ chosen: list[Path]
344
+ if npz_paths:
345
+ chosen = []
346
+ for raw in npz_paths:
347
+ item = Path(raw)
348
+ resolved = item if item.is_absolute() else (root / item)
349
+ if not resolved.is_file():
350
+ raise FileNotFoundError(f"NPZ not found: {resolved}")
351
+ chosen.append(resolved.resolve())
352
+ elif npz_catalog:
353
+ # Union in YAML order. Same file from two globs is kept once.
354
+ seen: set[Path] = set()
355
+ chosen = []
356
+ for pattern, max_shapes in npz_catalog:
357
+ hit = _glob_npz(root, str(pattern))
358
+ if exclude_combo:
359
+ hit = [p for p in hit if not _is_combo_npz(p)]
360
+ hit = [
361
+ p
362
+ for p in hit
363
+ if p.suffix.lower() == ".npz" and _is_parameterized_stem(p)
364
+ ]
365
+ hit = sorted(hit)
366
+ hit = _cap_per_shape(hit, max_files_per_shape)
367
+ if max_shapes is not None:
368
+ hit = _subsample_shapes(hit, int(max_shapes), int(seed))
369
+ for path in hit:
370
+ if path in seen:
371
+ continue
372
+ seen.add(path)
373
+ chosen.append(path)
374
+ else:
375
+ chosen = _glob_npz(root, npz_glob)
376
+
377
+ npz_only = [p for p in chosen if p.suffix.lower() == ".npz"]
378
+ if exclude_combo:
379
+ npz_only = [p for p in npz_only if not _is_combo_npz(p)]
380
+ if npz_catalog and not npz_paths:
381
+ # Already capped per glob; keep YAML union order (not a global sort).
382
+ capped = npz_only
383
+ else:
384
+ npz_only = sorted(npz_only)
385
+ capped = _cap_per_shape(npz_only, max_files_per_shape)
386
+ if not capped:
387
+ raise FileNotFoundError(
388
+ f"No occupancy NPZ files matched under {root} "
389
+ f"(glob={npz_glob!r}, catalog={bool(npz_catalog)}, "
390
+ f"explicit={bool(npz_paths)})"
391
+ )
392
+ return capped
393
+
394
+
395
+ def _summarize(points: PointsArray, labels: LabelsArray) -> str:
396
+ n = int(points.shape[0])
397
+ inside = float(labels.mean()) if n else float("nan")
398
+ xyz_min = points.min(axis=0) if n else np.full(3, np.nan, dtype=np.float32)
399
+ xyz_max = points.max(axis=0) if n else np.full(3, np.nan, dtype=np.float32)
400
+ return (
401
+ f"N={n}\n"
402
+ f"inside_fraction={inside:.6f}\n"
403
+ f"xyz_min={xyz_min.tolist()}\n"
404
+ f"xyz_max={xyz_max.tolist()}"
405
+ )
406
+
407
+
408
+ # Default smoke-check file from the v2 plan (dataset_test sphere).
409
+ _SAMPLE_RELATIVE = Path("exports") / "dataset_test" / "sphere__raycast_z_raut_s0.15_inout.npz"
410
+
411
+
412
+ if __name__ == "__main__":
413
+ import sys
414
+
415
+ from scatteringnet.config import load_config
416
+
417
+ cfg = load_config()
418
+ if "--catalog" in sys.argv:
419
+ paths = resolve_npz_catalog(
420
+ cfg.data_dir,
421
+ npz_glob=cfg.npz_glob,
422
+ npz_paths=cfg.npz_paths or None,
423
+ npz_catalog=cfg.npz_catalog or None,
424
+ max_files_per_shape=cfg.max_files_per_shape,
425
+ seed=cfg.seed,
426
+ )
427
+ print(f"data_dir={cfg.data_dir}")
428
+ print(f"npz_glob={cfg.npz_glob}")
429
+ print(f"npz_catalog={list(cfg.npz_catalog)}")
430
+ print(f"max_files_per_shape={cfg.max_files_per_shape}")
431
+ print(f"files={len(paths)}")
432
+ # Load a few files only — the full catalog can be thousands of NPZs.
433
+ preview = paths[:3]
434
+ n_all = 0
435
+ n_in = 0
436
+ for path in preview:
437
+ pts, labs = load_points_labels(path)
438
+ n = int(pts.shape[0])
439
+ inside = int((labs == 1.0).sum())
440
+ n_all += n
441
+ n_in += inside
442
+ print(
443
+ f" {path.name} N={n} inside={inside} outside={n - inside}"
444
+ )
445
+ print(
446
+ f"preview_files={len(preview)} preview_N={n_all} "
447
+ f"preview_inside={n_in} preview_outside={n_all - n_in}"
448
+ )
449
+ from scatteringnet.dataset import OccupancyMultiNpzDataset, make_dataloader
450
+
451
+ ds = OccupancyMultiNpzDataset(paths)
452
+ loader = make_dataloader(ds.parts[0], batch_size=8, shuffle=False)
453
+ xyz, y = next(iter(loader))
454
+ print(f"files={len(ds)} n_points={ds.n_points} parts={len(ds.parts)}")
455
+ print(f"batch xyz={tuple(xyz.shape)} y={tuple(y.shape)}")
456
+ else:
457
+ sample = cfg.data_dir / _SAMPLE_RELATIVE
458
+ pts, labs = load_points_labels(sample)
459
+ print(f"file={sample}")
460
+ print(_summarize(pts, labs))
src/geometry/__init__.py ADDED
@@ -0,0 +1,10 @@
 
 
 
 
 
 
 
 
 
 
 
1
+ """Training-side geometry: OBJ triangles and envelope samples."""
2
+
3
+ from scatteringnet.geometry.mesh_io import clear_triangle_cache
4
+ from scatteringnet.geometry.surface import clear_envelope_cache
5
+
6
+
7
+ def clear_geometry_caches() -> None:
8
+ """Drop OBJ triangle and envelope process caches."""
9
+ clear_triangle_cache()
10
+ clear_envelope_cache()
src/geometry/mesh_io.py ADDED
@@ -0,0 +1,80 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Load OBJ triangle meshes for the NPZ ↔ mesh join.
2
+
3
+ This module only returns ``vertices (V, 3)`` and ``faces (T, 3)``.
4
+ It does **not** sample the envelope.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ from pathlib import Path
10
+ import numpy as np
11
+ import trimesh
12
+ from numpy.typing import NDArray
13
+
14
+ from scatteringnet.geometry.trimesh_util import as_trimesh
15
+
16
+ VerticesArray = NDArray[np.float32]
17
+ FacesArray = NDArray[np.int32]
18
+
19
+ # Same resolved OBJ can back several NPZs (``max_files_per_shape``). Cache
20
+ # the triangle arrays so catalog load does not re-parse the file.
21
+ _TRIANGLE_CACHE: dict[str, tuple[VerticesArray, FacesArray]] = {}
22
+
23
+
24
+ def clear_triangle_cache() -> None:
25
+ """Drop cached OBJ arrays (tests / long-lived notebooks)."""
26
+ _TRIANGLE_CACHE.clear()
27
+
28
+
29
+ def load_obj_triangles(
30
+ path: Path | str,
31
+ *,
32
+ cache: bool = True,
33
+ ) -> tuple[VerticesArray, FacesArray]:
34
+ """
35
+ Read one OBJ as triangle vertices and face indices.
36
+
37
+ Parameters
38
+ ----------
39
+ path:
40
+ Existing ``.obj`` file.
41
+ cache:
42
+ Reuse arrays for the same resolved path (catalog load).
43
+
44
+ Returns
45
+ -------
46
+ vertices:
47
+ ``float32`` array of shape ``(V, 3)``.
48
+ faces:
49
+ ``int32`` array of shape ``(T, 3)`` (0-based vertex indices).
50
+ """
51
+ obj_path = Path(path)
52
+ if not obj_path.is_file():
53
+ raise FileNotFoundError(f"OBJ not found: {obj_path}")
54
+ if obj_path.suffix.lower() != ".obj":
55
+ raise ValueError(f"expected .obj, got {obj_path.suffix!r} ({obj_path})")
56
+
57
+ cache_key = str(obj_path.resolve())
58
+ if cache and cache_key in _TRIANGLE_CACHE:
59
+ return _TRIANGLE_CACHE[cache_key]
60
+
61
+ # process=False keeps the authored vertices; we only need the join.
62
+ loaded = trimesh.load(obj_path, force=None, process=False)
63
+ mesh = as_trimesh(loaded)
64
+ vertices = np.asarray(mesh.vertices, dtype=np.float32)
65
+ faces = np.asarray(mesh.faces, dtype=np.int32)
66
+ if vertices.ndim != 2 or vertices.shape[1] != 3:
67
+ raise ValueError(
68
+ f"vertices must have shape (V, 3), got {tuple(vertices.shape)} in {obj_path}"
69
+ )
70
+ if faces.ndim != 2 or faces.shape[1] != 3:
71
+ raise ValueError(
72
+ f"faces must have shape (T, 3), got {tuple(faces.shape)} in {obj_path}"
73
+ )
74
+ if int(faces.shape[0]) < 1:
75
+ raise ValueError(f"OBJ has no triangles: {obj_path}")
76
+
77
+ arrays = (vertices, faces)
78
+ if cache:
79
+ _TRIANGLE_CACHE[cache_key] = arrays
80
+ return arrays
src/geometry/surface.py ADDED
@@ -0,0 +1,154 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Surface envelope samples from a triangle mesh.
2
+
3
+ Area-weighted face darts (larger triangles get more samples). Each
4
+ sample is ``(x, y, z, nx, ny, nz)``: position plus the unit normal of
5
+ the triangle it sits on. World clouds are cached per
6
+ ``(mesh_key, n_surface, seed)``. AABB is applied to XYZ only.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import numpy as np
12
+ import trimesh
13
+ from numpy.typing import NDArray
14
+ from trimesh.sample import sample_surface
15
+
16
+ from scatteringnet.normalize import apply_normalization
17
+
18
+ PointsArray = NDArray[np.float32]
19
+ ENVELOPE_XYZ_DIM = 3
20
+ ENVELOPE_FEAT_DIM = 6
21
+
22
+ # Same OBJ + count + seed → same world samples (always 6-D).
23
+ _ENVELOPE_CACHE: dict[tuple[str, int, int], PointsArray] = {}
24
+
25
+
26
+ def clear_envelope_cache() -> None:
27
+ """Drop cached world-space envelopes (tests / long-lived notebooks)."""
28
+ _ENVELOPE_CACHE.clear()
29
+
30
+
31
+ def _triangle_normals(verts: np.ndarray, faces: np.ndarray) -> NDArray[np.float64]:
32
+ v0 = verts[faces[:, 0]]
33
+ v1 = verts[faces[:, 1]]
34
+ v2 = verts[faces[:, 2]]
35
+ cross = np.cross(v1 - v0, v2 - v0)
36
+ length = np.linalg.norm(cross, axis=1, keepdims=True)
37
+ ok = length[:, 0] > 1e-12
38
+ normals = np.zeros_like(cross)
39
+ normals[ok] = cross[ok] / length[ok]
40
+ return normals
41
+
42
+
43
+ def _pack_xyz_normal(xyz: np.ndarray, normals: np.ndarray) -> PointsArray:
44
+ """Concatenate XYZ with unit face normals → ``(N, 6)``."""
45
+ pos = np.asarray(xyz, dtype=np.float32)
46
+ nrm = np.asarray(normals, dtype=np.float32)
47
+ if pos.shape != nrm.shape or pos.ndim != 2 or pos.shape[1] != 3:
48
+ raise ValueError(
49
+ f"xyz/normals must be (N, 3), got {tuple(pos.shape)} / {tuple(nrm.shape)}"
50
+ )
51
+ length = np.linalg.norm(nrm, axis=1, keepdims=True)
52
+ ok = length[:, 0] > 1e-12
53
+ unit = np.zeros_like(nrm)
54
+ unit[ok] = nrm[ok] / length[ok]
55
+ return np.concatenate([pos, unit], axis=1)
56
+
57
+
58
+ def apply_envelope_aabb(
59
+ env: np.ndarray, center: np.ndarray, scale: float
60
+ ) -> PointsArray:
61
+ """AABB-normalize XYZ; leave normals as unit directions."""
62
+ arr = np.asarray(env, dtype=np.float32)
63
+ if arr.ndim != 2 or arr.shape[1] not in (ENVELOPE_XYZ_DIM, ENVELOPE_FEAT_DIM):
64
+ raise ValueError(
65
+ f"envelope must be (N, 3) or (N, 6), got {tuple(arr.shape)}"
66
+ )
67
+ out = arr.copy()
68
+ out[:, :3] = apply_normalization(out[:, :3], center, scale)
69
+ return out
70
+
71
+
72
+ def undo_envelope_aabb(
73
+ env: np.ndarray, center: np.ndarray, scale: float
74
+ ) -> PointsArray:
75
+ """Undo AABB on XYZ only."""
76
+ arr = np.asarray(env, dtype=np.float32)
77
+ if arr.ndim != 2 or arr.shape[1] not in (ENVELOPE_XYZ_DIM, ENVELOPE_FEAT_DIM):
78
+ raise ValueError(
79
+ f"envelope must be (N, 3) or (N, 6), got {tuple(arr.shape)}"
80
+ )
81
+ out = arr.copy()
82
+ c = np.asarray(center, dtype=np.float32).reshape(3)
83
+ out[:, :3] = out[:, :3] * np.float32(scale) + c
84
+ return out
85
+
86
+
87
+ def project_envelope_dim(env: np.ndarray, dim: int) -> PointsArray:
88
+ """Keep XYZ+normal or drop to XYZ for an older checkpoint."""
89
+ arr = np.asarray(env, dtype=np.float32)
90
+ want = int(dim)
91
+ if want == ENVELOPE_FEAT_DIM:
92
+ if arr.ndim != 2 or arr.shape[1] != ENVELOPE_FEAT_DIM:
93
+ raise ValueError(f"expected (N, 6) envelope, got {tuple(arr.shape)}")
94
+ return arr
95
+ if want == ENVELOPE_XYZ_DIM:
96
+ if arr.ndim != 2 or arr.shape[1] not in (ENVELOPE_XYZ_DIM, ENVELOPE_FEAT_DIM):
97
+ raise ValueError(f"expected (N, 3|6) envelope, got {tuple(arr.shape)}")
98
+ return arr[:, :3]
99
+ raise ValueError(f"envelope dim must be 3 or 6, got {want}")
100
+
101
+
102
+ def _sample_area(
103
+ verts: np.ndarray, tris: np.ndarray, count: int, seed: int
104
+ ) -> PointsArray:
105
+ mesh = trimesh.Trimesh(vertices=verts, faces=tris, process=False)
106
+ points, face_idx = sample_surface(mesh, count, seed=int(seed))
107
+ nrm = _triangle_normals(verts, tris)[np.asarray(face_idx, dtype=np.int64)]
108
+ out = _pack_xyz_normal(points, nrm)
109
+ if out.shape != (count, ENVELOPE_FEAT_DIM):
110
+ raise ValueError(
111
+ f"expected envelope shape {(count, ENVELOPE_FEAT_DIM)}, got {tuple(out.shape)}"
112
+ )
113
+ return out
114
+
115
+
116
+ def sample_surface_points(
117
+ vertices: np.ndarray,
118
+ faces: np.ndarray,
119
+ n_surface: int,
120
+ *,
121
+ seed: int = 1,
122
+ cache_key: str | None = None,
123
+ ) -> PointsArray:
124
+ """
125
+ Sample ``n_surface`` envelope points as ``(N, 6)`` XYZ + unit normal.
126
+
127
+ Face-area weighted: larger triangles receive more darts. No crease
128
+ or fold path — unused ``envelope_mix`` on old checkpoints is ignored.
129
+ """
130
+ count = int(n_surface)
131
+ if count < 1:
132
+ raise ValueError(f"n_surface must be >= 1, got {count}")
133
+ verts = np.asarray(vertices, dtype=np.float64)
134
+ tris = np.asarray(faces, dtype=np.int64)
135
+ if verts.ndim != 2 or verts.shape[1] != 3:
136
+ raise ValueError(f"vertices must have shape (V, 3), got {tuple(verts.shape)}")
137
+ if tris.ndim != 2 or tris.shape[1] != 3:
138
+ raise ValueError(f"faces must have shape (T, 3), got {tuple(tris.shape)}")
139
+
140
+ key: tuple[str, int, int] | None = None
141
+ if cache_key is not None:
142
+ key = (str(cache_key), count, int(seed))
143
+ cached = _ENVELOPE_CACHE.get(key)
144
+ if cached is not None:
145
+ return cached
146
+
147
+ out = _sample_area(verts, tris, count, int(seed))
148
+ if out.shape != (count, ENVELOPE_FEAT_DIM):
149
+ raise ValueError(
150
+ f"expected envelope shape {(count, ENVELOPE_FEAT_DIM)}, got {tuple(out.shape)}"
151
+ )
152
+ if key is not None:
153
+ _ENVELOPE_CACHE[key] = out
154
+ return out
src/geometry/trimesh_util.py ADDED
@@ -0,0 +1,19 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Shared trimesh flatten. No Open3D — safe for train-side mesh_io."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from typing import Any
6
+
7
+ import trimesh
8
+
9
+
10
+ def as_trimesh(mesh: Any) -> trimesh.Trimesh:
11
+ """Flatten a Trimesh or a Scene of triangle meshes to one Trimesh."""
12
+ if isinstance(mesh, trimesh.Scene):
13
+ geoms = [g for g in mesh.geometry.values() if isinstance(g, trimesh.Trimesh)]
14
+ if not geoms:
15
+ raise ValueError("Scene contains no triangle meshes")
16
+ mesh = trimesh.util.concatenate(geoms)
17
+ if not isinstance(mesh, trimesh.Trimesh):
18
+ raise TypeError(f"Unsupported mesh type: {type(mesh)!r}")
19
+ return mesh
src/gradio/app.py ADDED
@@ -0,0 +1,503 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Gradio occupancy demo: upload an OBJ, fill the AABB, classify with best.pt.
2
+
3
+ The 3D pane is a persistent Babylon canvas (``orbit.js`` via ``launch(js=)``).
4
+ Gradio ``Model3D`` remounts WebGL on every new GLB — that is the gray flash
5
+ and the camera snap. ``gr.HTML`` scripts are stripped, so the viewer JS is
6
+ not put in the HTML component. The inspect tool remains ``src/viewer``.
7
+
8
+ Launch (conda env scatteringNet):
9
+
10
+ python src/gradio/app.py
11
+ """
12
+
13
+ from __future__ import annotations
14
+
15
+ import html
16
+ import os
17
+ import sys
18
+ import tempfile
19
+ from pathlib import Path
20
+
21
+ # ZeroGPU scans for ``@spaces.GPU`` at import. Import this before torch.
22
+ try:
23
+ import spaces
24
+ except ImportError:
25
+ spaces = None
26
+
27
+ import numpy as np
28
+
29
+ # Pip package first — this directory must not shadow it.
30
+ import gradio as gr
31
+
32
+ _HERE = Path(__file__).resolve().parent
33
+ # Repo root (Space ships ``scatteringnet/`` here; local uses pip ``-e .``).
34
+ _ROOT = _HERE.parents[1]
35
+ if str(_ROOT) not in sys.path:
36
+ sys.path.insert(0, str(_ROOT))
37
+ if str(_HERE) not in sys.path:
38
+ sys.path.insert(0, str(_HERE))
39
+
40
+ from figure import ( # noqa: E402
41
+ DEFAULT_MESH_OPACITY,
42
+ DEFAULT_POINT_SIZE,
43
+ empty_figure,
44
+ occupancy_figure,
45
+ )
46
+ from pipeline import ( # noqa: E402
47
+ DEFAULT_CUT,
48
+ DEFAULT_DENSITY,
49
+ apply_cut,
50
+ default_run_id,
51
+ fill_and_infer,
52
+ )
53
+ from scatteringnet.viewer.obj_fill import triangles_from_obj_text # noqa: E402
54
+
55
+ # Shipped sample meshes in examples/ (not in the training catalog).
56
+ _EXAMPLE_OBJS = (
57
+ "Obese.obj",
58
+ "horse.obj",
59
+ "Player.obj",
60
+ "dog.obj",
61
+ "Helix_bend.obj",
62
+ "TorusX3_box.obj",
63
+ )
64
+
65
+ # Inspect viewer: scene.background = 0x2a2a32 (not near-black).
66
+ _BG_HEX = "#2a2a32"
67
+ _ORBIT_JS = _HERE / "orbit.js"
68
+ # Host never goes in event outputs — remounting it would flash again.
69
+ _ORBIT_HOST = (
70
+ f'<div id="sn-orbit-host" style="position:relative;width:100%;height:640px;'
71
+ f'background:{_BG_HEX};border-radius:8px;overflow:hidden;">'
72
+ '<canvas id="sn-orbit" style="width:100%;height:100%;display:block;"></canvas>'
73
+ '<div id="sn-drop-hint">Drop an OBJ file here</div>'
74
+ "</div>"
75
+ )
76
+ # Shared by Blocks (Hugging Face finds ``demo``) and local launch().
77
+ _ORBIT_CSS = (
78
+ "#sn-cmd-wrap { display: none !important; }"
79
+ "#sn-obj-file-slot {"
80
+ " position: fixed !important; left: -100vw !important; top: 0 !important;"
81
+ " width: 8px !important; height: 8px !important; overflow: hidden !important;"
82
+ " opacity: 0 !important; pointer-events: none !important;"
83
+ "}"
84
+ "#sn-obj-file-slot .file-preview-holder, #sn-obj-file-slot table,"
85
+ " #sn-obj-file-slot .filename { display: none !important; }"
86
+ "#sn-orbit-host.sn-drop-over { outline: 2px solid #f5a623; outline-offset: -2px; }"
87
+ "#sn-drop-hint {"
88
+ " position: absolute; left: 0; right: 0; top: 12px;"
89
+ " text-align: center; pointer-events: none; z-index: 2;"
90
+ " font-size: 13px; line-height: 1.3; color: #aabbcc;"
91
+ " text-shadow: 0 1px 2px #1a1a20;"
92
+ "}"
93
+ )
94
+
95
+
96
+ def _make_blocks() -> gr.Blocks:
97
+ """Gradio 6 moved ``js`` / ``css`` off Blocks onto ``launch()``."""
98
+ return gr.Blocks(title="scatteringNet occupancy")
99
+
100
+
101
+ def _patch_launch(blocks: gr.Blocks) -> None:
102
+ """HF calls ``demo.launch()`` with no kwargs; inject orbit JS / CSS."""
103
+ orbit_js = _ORBIT_JS.read_text(encoding="utf-8")
104
+ orig = blocks.launch
105
+
106
+ def launch(*args, **kwargs):
107
+ kwargs.setdefault("js", orbit_js)
108
+ kwargs.setdefault("css", _ORBIT_CSS)
109
+ kwargs.setdefault("allowed_paths", [str(_glb_dir())])
110
+ try:
111
+ return orig(*args, **kwargs)
112
+ except TypeError:
113
+ kwargs.pop("js", None)
114
+ return orig(*args, **kwargs)
115
+
116
+ blocks.launch = launch
117
+
118
+
119
+ def _glb_dir() -> Path:
120
+ """Same folder ``figure._export_glb`` writes; ``allowed_paths`` serves it."""
121
+ folder = Path(tempfile.gettempdir()) / "scatteringnet_gradio"
122
+ folder.mkdir(parents=True, exist_ok=True)
123
+ return folder
124
+
125
+
126
+ def _cmd_html(
127
+ glb: str,
128
+ reset_n: int,
129
+ point_size: float = DEFAULT_POINT_SIZE,
130
+ wire: bool = False,
131
+ force_default: bool = False,
132
+ ) -> str:
133
+ """Tiny span orbit.js polls. Remounting this does not remount the canvas.
134
+
135
+ ``force_default`` (last field) is **Reset view** only. Sample / drop /
136
+ Load OBJ / Run keep the current orbit — same as the original load fix.
137
+ """
138
+ path = html.escape(str(Path(glb).resolve()))
139
+ px = max(1, min(24, int(round(float(point_size)))))
140
+ w = 1 if wire else 0
141
+ d = 1 if force_default else 0
142
+ return f'<span id="sn-cmd">{path}|{int(reset_n)}|{px}|{w}|{d}</span>'
143
+
144
+
145
+ def _held_path(file_obj) -> str | None:
146
+ """File widget / UploadButton → filepath string for later Run / restyle."""
147
+ if file_obj is None:
148
+ return None
149
+ return str(getattr(file_obj, "name", None) or file_obj)
150
+
151
+
152
+ def _read_upload(file_obj) -> tuple[str, str]:
153
+ """Gradio File (filepath) → (basename, OBJ text)."""
154
+ if file_obj is None:
155
+ raise ValueError("upload an .obj")
156
+ path = Path(getattr(file_obj, "name", None) or str(file_obj))
157
+ if path.suffix.lower() != ".obj":
158
+ raise ValueError("file must be an .obj")
159
+ return path.name, path.read_text(encoding="utf-8", errors="replace")
160
+
161
+
162
+ def _status_md(result: dict, *, cut: float, shown_out: bool) -> str:
163
+ """Short inspect-style summary under the view."""
164
+ t = result.get("timings") or {}
165
+ extra = " (outside hidden)" if not shown_out else ""
166
+ return (
167
+ f"**{result.get('obj_name', 'OBJ')}** · `{result.get('run_id', '')}` \n"
168
+ f"Inside **{result['n_inside']}** / outside **{result['n_outside']}** "
169
+ f"(cut {cut:.2f}){extra} \n"
170
+ f"Lattice `{result['n']}` pts · spacing `{result['used_spacing']:.3f}` · "
171
+ f"grid `{result['grid']}` \n"
172
+ f"Times s: fill+forward total `{t.get('server_total', 0):.3f}` · "
173
+ f"envelope `{t.get('envelope', 0):.3f}` · run `{t.get('forward', 0):.3f}` · "
174
+ f"device `{t.get('device', '?')}`"
175
+ )
176
+
177
+
178
+ def _view(result: dict | None, cut: float, show_outside: bool, mesh_opacity: float):
179
+ """Rebuild the occupancy GLB. The canvas stays; only the cmd span changes."""
180
+ if not result:
181
+ return str(empty_figure()), "Upload an OBJ, then **Run model**."
182
+ pred, n_in, n_out = apply_cut(result["probs"], float(cut))
183
+ view = dict(result)
184
+ view["pred"] = pred
185
+ view["n_inside"] = n_in
186
+ view["n_outside"] = n_out
187
+ glb = occupancy_figure(
188
+ view["vertices"],
189
+ view["faces"],
190
+ view["points"],
191
+ pred,
192
+ show_outside=bool(show_outside),
193
+ title=str(view.get("obj_name") or "occupancy fill"),
194
+ mesh_opacity=float(mesh_opacity),
195
+ )
196
+ return str(glb), _status_md(view, cut=float(cut), shown_out=bool(show_outside))
197
+
198
+
199
+ def run_model(file_obj, obj_held, density, cut, show_outside, mesh_opacity):
200
+ """Fill + infer with the single default checkpoint (no model picker)."""
201
+ src = file_obj or obj_held
202
+ held = _held_path(src)
203
+ try:
204
+ obj_name, obj_text = _read_upload(src)
205
+ result = fill_and_infer(
206
+ obj_text,
207
+ obj_name=obj_name,
208
+ run_id=default_run_id(),
209
+ density=float(density),
210
+ )
211
+ glb, md = _view(result, float(cut), bool(show_outside), float(mesh_opacity))
212
+ return result, glb, md, None, held
213
+ except Exception as exc:
214
+ return None, str(empty_figure(str(exc))), f"**Error:** {exc}", None, held
215
+
216
+
217
+ def preview_upload(file_obj, mesh_opacity=DEFAULT_MESH_OPACITY):
218
+ """Show the uploaded mesh on the floor immediately (no occupancy yet)."""
219
+ if file_obj is None:
220
+ return None, str(empty_figure()), "Upload an OBJ, then click **Run model**."
221
+ try:
222
+ obj_name, obj_text = _read_upload(file_obj)
223
+ vertices, faces = triangles_from_obj_text(obj_text)
224
+ glb = occupancy_figure(
225
+ vertices,
226
+ faces,
227
+ np.zeros((0, 3), dtype=np.float32),
228
+ np.zeros((0,), dtype=np.uint8),
229
+ title=obj_name,
230
+ mesh_opacity=float(mesh_opacity),
231
+ )
232
+ return (
233
+ None,
234
+ str(glb),
235
+ f"**{obj_name}** is on the floor. Click **Run model** to classify the fill.",
236
+ )
237
+ except Exception as exc:
238
+ return None, str(empty_figure(str(exc))), f"**Error:** {exc}"
239
+
240
+
241
+ def accept_obj(
242
+ file_obj,
243
+ mesh_opacity=DEFAULT_MESH_OPACITY,
244
+ reset_n=0,
245
+ current_glb=None,
246
+ ):
247
+ """
248
+ Load an OBJ, then clear the File box so the next drop can fire change.
249
+
250
+ Keep ``reset_n`` so orbit.js does not reframe. Clearing the picker
251
+ (``file_obj is None``) must not write a new empty GLB — that changed
252
+ the cmd path and made the camera reload. **Reset view** is the only bump.
253
+ """
254
+ keep = int(reset_n or 0)
255
+ if file_obj is None:
256
+ glb = current_glb or str(empty_figure())
257
+ yield None, str(glb), keep, "Upload an OBJ, then click **Run model**.", None, None
258
+ return
259
+ held = _held_path(file_obj)
260
+ state, glb, md = preview_upload(file_obj, mesh_opacity)
261
+ yield state, str(glb), keep, md, None, held
262
+
263
+
264
+ def restyle(result, file_obj, obj_held, cut, show_outside, mesh_opacity):
265
+ """Rebuild GLB after cut / opacity / outside. Reset counter is not touched."""
266
+ try:
267
+ if result:
268
+ return _view(result, float(cut), bool(show_outside), float(mesh_opacity))
269
+ src = file_obj or obj_held
270
+ if src is None:
271
+ return str(empty_figure()), "Upload an OBJ, then **Run model**."
272
+ _, glb, md = preview_upload(src, mesh_opacity)
273
+ return str(glb), md
274
+ except Exception as exc:
275
+ return str(empty_figure(str(exc))), f"**Error:** {exc}"
276
+
277
+
278
+ def bump_reset(glb_held, reset_n, point_size, wire):
279
+ """Reset view: same GLB, increment token so JS frames 45° / 70°."""
280
+ nxt = int(reset_n or 0) + 1
281
+ path = glb_held or str(empty_figure())
282
+ return nxt, _cmd_html(path, nxt, point_size, wire, force_default=True)
283
+
284
+
285
+ def accept_obj_ui(
286
+ file_obj,
287
+ mesh_opacity=DEFAULT_MESH_OPACITY,
288
+ reset_n=0,
289
+ point_size=DEFAULT_POINT_SIZE,
290
+ wire=False,
291
+ glb_held=None,
292
+ ):
293
+ """accept_obj plus the #sn-cmd span orbit.js reads."""
294
+ for state, glb, nxt, md, cleared, held in accept_obj(
295
+ file_obj, mesh_opacity, reset_n, current_glb=glb_held
296
+ ):
297
+ yield state, glb, nxt, _cmd_html(glb, nxt, point_size, wire), md, cleared, held
298
+
299
+
300
+ def _maybe_gpu(fn):
301
+ """ZeroGPU requires at least one ``@spaces.GPU`` at startup. No-op locally."""
302
+ if spaces is None:
303
+ return fn
304
+ return spaces.GPU(duration=30)(fn)
305
+
306
+
307
+ @_maybe_gpu
308
+ def run_model_ui(
309
+ file_obj, obj_held, density, cut, show_outside, mesh_opacity, reset_n, point_size, wire
310
+ ):
311
+ """Run: new GLB, same reset token, so the canvas keeps the orbit."""
312
+ result, glb, md, cleared, held = run_model(
313
+ file_obj, obj_held, density, cut, show_outside, mesh_opacity
314
+ )
315
+ return result, glb, _cmd_html(glb, int(reset_n or 0), point_size, wire), md, cleared, held
316
+
317
+
318
+ def restyle_ui(
319
+ result, file_obj, obj_held, cut, show_outside, mesh_opacity, reset_n, point_size, wire
320
+ ):
321
+ """Restyle: new GLB, same reset token."""
322
+ glb, md = restyle(result, file_obj, obj_held, cut, show_outside, mesh_opacity)
323
+ return glb, _cmd_html(glb, int(reset_n or 0), point_size, wire), md
324
+
325
+
326
+ def set_view_flags(glb_held, reset_n, point_size, wire):
327
+ """Dot size / wireframe — no new GLB, orbit.js applies both."""
328
+ path = glb_held or str(empty_figure())
329
+ return _cmd_html(path, int(reset_n or 0), point_size, wire)
330
+
331
+
332
+ def build_demo() -> gr.Blocks:
333
+ """One-page Gradio occupancy fill."""
334
+ # Only list files that are actually on disk (partial checkout still launches).
335
+ example_list = [
336
+ str(path)
337
+ for name in _EXAMPLE_OBJS
338
+ if (path := _HERE / "examples" / name).is_file()
339
+ ] or None
340
+ empty_glb = empty_figure()
341
+
342
+ with _make_blocks() as demo:
343
+ gr.Markdown(
344
+ """
345
+ # scatteringNet — occupancy fill
346
+
347
+ A trained occupancy network that fills a 3D mesh with points it labels
348
+ **inside** the solid (not outside).
349
+
350
+ [Code](https://github.com/PerryGu/scatteringNet)
351
+
352
+ [Video](https://youtu.be/vU45O0Mu0o4)
353
+ """.strip()
354
+ )
355
+ state = gr.State(None)
356
+ obj_held = gr.State(None)
357
+ glb_held = gr.State(str(Path(empty_glb).resolve()))
358
+ reset_n = gr.State(0)
359
+ with gr.Row():
360
+ with gr.Column(scale=1):
361
+ load_btn = gr.Button("Load OBJ", elem_id="sn-load-obj")
362
+ density_in = gr.Slider(
363
+ 0,
364
+ 100,
365
+ value=DEFAULT_DENSITY,
366
+ step=1,
367
+ label="Density (higher = denser lattice)",
368
+ )
369
+ cut_in = gr.Slider(
370
+ 0.05,
371
+ 0.95,
372
+ value=DEFAULT_CUT,
373
+ step=0.01,
374
+ label="Inside cut",
375
+ )
376
+ opacity_in = gr.Slider(
377
+ 0,
378
+ 100,
379
+ value=DEFAULT_MESH_OPACITY,
380
+ step=1,
381
+ label="Mesh opacity",
382
+ info="Default 50% so occupancy points show through the shell.",
383
+ )
384
+ psize_in = gr.Slider(
385
+ 1,
386
+ 24,
387
+ value=DEFAULT_POINT_SIZE,
388
+ step=1,
389
+ label="Dot size",
390
+ info="Occupancy points in the 3D view (pixels).",
391
+ )
392
+ with gr.Row():
393
+ show_out = gr.Checkbox(label="Show outside points", value=False)
394
+ wire_in = gr.Checkbox(label="Wireframe", value=False)
395
+ with gr.Row():
396
+ run_btn = gr.Button("Run model", variant="primary")
397
+ reset_btn = gr.Button("Reset view")
398
+ status = gr.Markdown(
399
+ "Drop an OBJ on the **3D view**, or click **Load OBJ**. "
400
+ "**Run model** fills the volume."
401
+ )
402
+ with gr.Column(scale=2):
403
+ # Never list this HTML in outputs — a new value remounts the canvas.
404
+ gr.HTML(value=_ORBIT_HOST, elem_id="sn-orbit-wrap")
405
+ # Plumbing only. The cmd span is a temp GLB path; if it sits in the
406
+ # layout it paints through Load OBJ (even with height:0).
407
+ with gr.Column(elem_id="sn-obj-file-slot"):
408
+ cmd = gr.HTML(
409
+ value=_cmd_html(str(empty_glb), 0, DEFAULT_POINT_SIZE),
410
+ elem_id="sn-cmd-wrap",
411
+ visible="hidden",
412
+ container=False,
413
+ )
414
+ obj_in = gr.File(
415
+ label="OBJ",
416
+ file_types=[".obj"],
417
+ type="filepath",
418
+ show_label=False,
419
+ container=False,
420
+ elem_id="sn-obj-file",
421
+ )
422
+ _accept_out = [state, glb_held, reset_n, cmd, status, obj_in, obj_held]
423
+ _run_out = [state, glb_held, cmd, status, obj_in, obj_held]
424
+ _style_out = [glb_held, cmd, status]
425
+ if example_list:
426
+ # One file column; the other inputs keep the live sliders / Wireframe.
427
+ # inputs=obj_in alone called accept_obj_ui with wire=False.
428
+ gr.Markdown(
429
+ "These geometries were not in the model's training catalog."
430
+ )
431
+ gr.Examples(
432
+ examples=[[p] for p in example_list],
433
+ inputs=[obj_in, opacity_in, reset_n, psize_in, wire_in, glb_held],
434
+ outputs=_accept_out,
435
+ fn=accept_obj_ui,
436
+ run_on_click=True,
437
+ cache_examples=False,
438
+ label="Sample OBJ",
439
+ )
440
+ # Picker is a plain Button: UploadButton paints the path over the label.
441
+ # orbit.js clicks #sn-obj-file's input; File.upload runs accept.
442
+ # Do not bind File.clear — we empty the box on purpose so DND stays.
443
+ obj_in.upload(
444
+ accept_obj_ui,
445
+ inputs=[obj_in, opacity_in, reset_n, psize_in, wire_in, glb_held],
446
+ outputs=_accept_out,
447
+ )
448
+ run_btn.click(
449
+ run_model_ui,
450
+ inputs=[
451
+ obj_in,
452
+ obj_held,
453
+ density_in,
454
+ cut_in,
455
+ show_out,
456
+ opacity_in,
457
+ reset_n,
458
+ psize_in,
459
+ wire_in,
460
+ ],
461
+ outputs=_run_out,
462
+ )
463
+ reset_btn.click(
464
+ bump_reset,
465
+ inputs=[glb_held, reset_n, psize_in, wire_in],
466
+ outputs=[reset_n, cmd],
467
+ )
468
+ _style_in = [
469
+ state,
470
+ obj_in,
471
+ obj_held,
472
+ cut_in,
473
+ show_out,
474
+ opacity_in,
475
+ reset_n,
476
+ psize_in,
477
+ wire_in,
478
+ ]
479
+ cut_in.release(restyle_ui, inputs=_style_in, outputs=_style_out)
480
+ opacity_in.release(restyle_ui, inputs=_style_in, outputs=_style_out)
481
+ show_out.change(restyle_ui, inputs=_style_in, outputs=_style_out)
482
+ _flag_in = [glb_held, reset_n, psize_in, wire_in]
483
+ psize_in.release(set_view_flags, inputs=_flag_in, outputs=cmd)
484
+ wire_in.change(set_view_flags, inputs=_flag_in, outputs=cmd)
485
+ return demo
486
+
487
+
488
+ # HF / ZeroGPU look for a module-level ``demo`` and for ``@spaces.GPU``.
489
+ demo = build_demo()
490
+ demo.queue()
491
+ _patch_launch(demo)
492
+
493
+
494
+ def main() -> None:
495
+ """Local binds localhost; Spaces set ``PORT`` and need ``0.0.0.0``."""
496
+ on_space = bool(os.environ.get("SPACE_ID") or os.environ.get("PORT"))
497
+ port = int(os.environ["PORT"]) if os.environ.get("PORT") else 7860
498
+ host = "0.0.0.0" if on_space else "127.0.0.1"
499
+ demo.launch(server_name=host, server_port=port)
500
+
501
+
502
+ if __name__ == "__main__":
503
+ main()
src/gradio/examples/Helix_bend.obj ADDED
The diff for this file is too large to render. See raw diff
 
src/gradio/examples/Obese.obj ADDED
The diff for this file is too large to render. See raw diff
 
src/gradio/examples/Player.obj ADDED
The diff for this file is too large to render. See raw diff
 
src/gradio/examples/TorusX3_box.obj ADDED
The diff for this file is too large to render. See raw diff
 
src/gradio/examples/dog.obj ADDED
The diff for this file is too large to render. See raw diff
 
src/gradio/examples/horse.obj ADDED
The diff for this file is too large to render. See raw diff
 
src/gradio/figure.py ADDED
@@ -0,0 +1,255 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """GLB scene for Gradio ``Model3D`` (Babylon.js orbit, Y-up).
2
+
3
+ Plotly 3D was the wrong tool: its camera is a unit-cube, not world Y-up, so
4
+ the floor never sat on the ground. This writes a real mesh (floor grid +
5
+ OBJ + occupancy points) and lets Gradio's 3D viewer frame it.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import tempfile
11
+ import uuid
12
+ from pathlib import Path
13
+
14
+ import numpy as np
15
+ import trimesh
16
+ from trimesh.path.entities import Line
17
+ from trimesh.path.path import Path3D
18
+ from trimesh.visual.material import PBRMaterial
19
+ from trimesh.visual.texture import TextureVisuals
20
+
21
+ # Inspect COLOR_INSIDE is 0xffaa00. Babylon treats GLB vertex colors as
22
+ # linear, so that same RGB reads as lemon-yellow. Drop the green channel
23
+ # so the on-screen fill matches the inspect orange.
24
+ COLOR_INSIDE = np.array([255, 136, 0, 255], dtype=np.uint8)
25
+ COLOR_OUTSIDE = np.array([74, 109, 140, 220], dtype=np.uint8)
26
+ COLOR_MESH_RGB = (0.60, 0.64, 0.70)
27
+ # THREE.GridHelper(span, 20, 0xaabbcc, 0x556677) — center vs cells.
28
+ COLOR_GRID = np.array([85, 102, 119, 255], dtype=np.uint8)
29
+ COLOR_GRID_CENTER = np.array([170, 187, 204, 255], dtype=np.uint8)
30
+ DEFAULT_MESH_OPACITY = 50
31
+ # Babylon POINTS are 1 px in the GLB; orbit.js applies this pixel size.
32
+ DEFAULT_POINT_SIZE = 8
33
+ # Match pipeline.MAX_FILL_POINTS so the GLB is not a random subset of the lattice.
34
+ MAX_PLOT_POINTS = 80_000
35
+ FLOOR_DIVS = 20
36
+
37
+
38
+ def _floor_bounds(vertices: np.ndarray | None) -> tuple[float, float, float, float, float]:
39
+ """XZ rectangle and Y of the floor (content ymin, else y=0)."""
40
+ if vertices is None or int(np.asarray(vertices).size) == 0:
41
+ return -2.0, 2.0, -2.0, 2.0, 0.0
42
+ verts = np.asarray(vertices, dtype=np.float64)
43
+ vmin = verts.min(axis=0)
44
+ vmax = verts.max(axis=0)
45
+ radius = 0.5 * float(np.max(vmax - vmin))
46
+ span = max(radius * 4.0, 4.0)
47
+ cx = 0.5 * float(vmin[0] + vmax[0])
48
+ cz = 0.5 * float(vmin[2] + vmax[2])
49
+ y = float(vmin[1])
50
+ half = 0.5 * span
51
+ return cx - half, cx + half, cz - half, cz + half, y
52
+
53
+
54
+ def _clamp_opacity(value: float) -> float:
55
+ """Slider 0–100 → 0–1. Missing / NaN → default 50%."""
56
+ try:
57
+ t = float(value)
58
+ except (TypeError, ValueError):
59
+ t = float(DEFAULT_MESH_OPACITY)
60
+ if not np.isfinite(t):
61
+ t = float(DEFAULT_MESH_OPACITY)
62
+ return min(1.0, max(0.0, t / 100.0))
63
+
64
+
65
+ def _shell_mesh(vertices: np.ndarray, faces: np.ndarray, opacity: float) -> trimesh.Trimesh:
66
+ """
67
+ OBJ shell with a GLTF BLEND material.
68
+
69
+ Vertex-color alpha is ignored by Babylon Model3D (that is why the
70
+ torus looked solid). ``alphaMode=BLEND`` + baseColorFactor.a is the
71
+ channel the viewer actually uses.
72
+ """
73
+ mesh = trimesh.Trimesh(vertices=vertices, faces=faces, process=False)
74
+ alpha = _clamp_opacity(opacity)
75
+ mesh.visual = TextureVisuals(
76
+ material=PBRMaterial(
77
+ name="shell",
78
+ baseColorFactor=[COLOR_MESH_RGB[0], COLOR_MESH_RGB[1], COLOR_MESH_RGB[2], alpha],
79
+ metallicFactor=0.0,
80
+ roughnessFactor=0.85,
81
+ alphaMode="BLEND",
82
+ doubleSided=True,
83
+ )
84
+ )
85
+ mesh.metadata["name"] = "occ_shell"
86
+ return mesh
87
+
88
+
89
+ def _grid_mesh(vertices: np.ndarray | None, *, n: int = FLOOR_DIVS) -> Path3D:
90
+ """
91
+ XZ floor as GL_LINES (1 px in Babylon), not extruded boxes.
92
+
93
+ Boxes read as a thick waffle. The inspect viewer uses
94
+ ``THREE.GridHelper(span, 20, 0xaabbcc, 0x556677)`` — same count
95
+ and the lighter centre cross.
96
+ """
97
+ x0, x1, z0, z1, y = _floor_bounds(vertices)
98
+ n = max(2, int(n))
99
+ mid = n // 2
100
+ verts: list[list[float]] = []
101
+ entities: list[Line] = []
102
+ colors: list[np.ndarray] = []
103
+ for i in range(n + 1):
104
+ t = i / n
105
+ x = x0 + (x1 - x0) * t
106
+ z = z0 + (z1 - z0) * t
107
+ color = COLOR_GRID_CENTER if i == mid else COLOR_GRID
108
+ i0 = len(verts)
109
+ verts.extend(([x, y, z0], [x, y, z1]))
110
+ entities.append(Line(points=[i0, i0 + 1]))
111
+ colors.append(color)
112
+ i1 = len(verts)
113
+ verts.extend(([x0, y, z], [x1, y, z]))
114
+ entities.append(Line(points=[i1, i1 + 1]))
115
+ colors.append(color)
116
+ return Path3D(
117
+ entities=entities,
118
+ vertices=np.asarray(verts, dtype=np.float64),
119
+ colors=np.asarray(colors, dtype=np.uint8),
120
+ process=False,
121
+ )
122
+
123
+
124
+ def _axis_mesh(vertices: np.ndarray | None) -> Path3D:
125
+ """
126
+ RGB triad as GL_LINES — same idea as ``THREE.AxesHelper``.
127
+
128
+ Length is ``max(radius * 0.45, 1)`` with ``span = 4 * radius``,
129
+ matching the inspect helper. Boxes looked like fat sticks.
130
+ """
131
+ x0, x1, z0, z1, y = _floor_bounds(vertices)
132
+ span = max(x1 - x0, z1 - z0, 1.0)
133
+ length = max(0.45 * (span / 4.0), 1.0)
134
+ origin = np.array([0.5 * (x0 + x1), y, 0.5 * (z0 + z1)], dtype=np.float64)
135
+ # AxesHelper: +X red, +Y green, +Z blue.
136
+ specs = (
137
+ (np.array([length, 0.0, 0.0]), (255, 0, 0, 255)),
138
+ (np.array([0.0, length, 0.0]), (0, 255, 0, 255)),
139
+ (np.array([0.0, 0.0, length]), (0, 0, 255, 255)),
140
+ )
141
+ verts: list[np.ndarray] = []
142
+ entities: list[Line] = []
143
+ colors: list[tuple[int, int, int, int]] = []
144
+ for delta, rgba in specs:
145
+ i0 = len(verts)
146
+ verts.extend((origin, origin + delta))
147
+ entities.append(Line(points=[i0, i0 + 1]))
148
+ colors.append(rgba)
149
+ return Path3D(
150
+ entities=entities,
151
+ vertices=np.asarray(verts, dtype=np.float64),
152
+ colors=np.asarray(colors, dtype=np.uint8),
153
+ process=False,
154
+ )
155
+
156
+
157
+ def _points_mesh(xyz: np.ndarray, rgba: np.ndarray) -> trimesh.Trimesh | None:
158
+ """Occupancy dots as a vertex cloud (GLB point primitive)."""
159
+ if xyz.ndim != 2 or xyz.shape[0] < 1:
160
+ return None
161
+ cloud = trimesh.points.PointCloud(xyz.astype(np.float64), colors=rgba)
162
+ # Node name survives GLB import so orbit.js can set pointSize.
163
+ cloud.metadata["name"] = "occ_points"
164
+ return cloud
165
+
166
+
167
+ def _subsample(points: np.ndarray, pred: np.ndarray, cap: int) -> tuple[np.ndarray, np.ndarray]:
168
+ n = int(points.shape[0])
169
+ if n <= cap:
170
+ return points, pred
171
+ rng = np.random.default_rng(0)
172
+ inside = np.flatnonzero(pred > 0)
173
+ outside = np.flatnonzero(pred == 0)
174
+ n_in = min(int(inside.shape[0]), cap)
175
+ take_in = rng.choice(inside, size=n_in, replace=False) if n_in else np.empty(0, dtype=int)
176
+ remain = cap - n_in
177
+ n_out = min(int(outside.shape[0]), remain)
178
+ take_out = (
179
+ rng.choice(outside, size=n_out, replace=False) if n_out else np.empty(0, dtype=int)
180
+ )
181
+ idx = np.concatenate([take_in, take_out])
182
+ return points[idx], pred[idx]
183
+
184
+
185
+ def _placeholder_geom() -> trimesh.Trimesh:
186
+ """Tiny triangle so trimesh can export when the draw list is empty."""
187
+ dummy = trimesh.Trimesh(
188
+ vertices=np.array([[0.0, 0.0, 0.0], [1e-4, 0.0, 0.0], [0.0, 0.0, 1e-4]]),
189
+ faces=np.array([[0, 1, 2]], dtype=np.int64),
190
+ process=False,
191
+ )
192
+ dummy.metadata["name"] = "occ_empty"
193
+ return dummy
194
+
195
+
196
+ def _export_glb(geoms: list) -> Path:
197
+ """Write a unique GLB so the orbit fetch is not served from a stale cache."""
198
+ folder = Path(tempfile.gettempdir()) / "scatteringnet_gradio"
199
+ folder.mkdir(parents=True, exist_ok=True)
200
+ path = folder / f"view_{uuid.uuid4().hex[:10]}.glb"
201
+ scene = trimesh.Scene()
202
+ added = 0
203
+ for i, geom in enumerate(geoms):
204
+ if geom is None:
205
+ continue
206
+ meta = getattr(geom, "metadata", None) or {}
207
+ name = str(meta["name"]) if isinstance(meta, dict) and meta.get("name") else f"g{i}"
208
+ scene.add_geometry(geom, node_name=name)
209
+ added += 1
210
+ if added < 1:
211
+ scene.add_geometry(_placeholder_geom(), node_name="occ_empty")
212
+ scene.export(path)
213
+ return path
214
+
215
+
216
+ def empty_figure(message: str = "", **_kwargs) -> Path:
217
+ """
218
+ Placeholder GLB so the cmd path can change. The visible floor lives
219
+ in orbit.js and is never replaced (a per-mesh grid was the camera jump).
220
+ """
221
+ _ = message
222
+ return _export_glb([])
223
+
224
+
225
+ def occupancy_figure(
226
+ vertices: np.ndarray,
227
+ faces: np.ndarray,
228
+ points: np.ndarray,
229
+ pred: np.ndarray,
230
+ *,
231
+ show_outside: bool = False,
232
+ title: str = "",
233
+ mesh_opacity: float = DEFAULT_MESH_OPACITY,
234
+ **_kwargs,
235
+ ) -> Path:
236
+ """Mesh + occupancy points as one GLB. Floor stays in orbit.js."""
237
+ _ = title
238
+ verts = np.asarray(vertices, dtype=np.float32) if vertices is not None else np.zeros((0, 3), np.float32)
239
+ tris = np.asarray(faces, dtype=np.int32) if faces is not None else np.zeros((0, 3), np.int32)
240
+ geoms: list = []
241
+ if verts.ndim == 2 and verts.shape[0] > 0:
242
+ if tris.ndim == 2 and tris.shape[0] > 0 and tris.shape[1] == 3:
243
+ geoms.append(_shell_mesh(verts, tris, mesh_opacity))
244
+ xyz = np.asarray(points, dtype=np.float32)
245
+ if xyz.ndim == 2 and xyz.shape[0] > 0:
246
+ labels = np.asarray(pred, dtype=np.uint8).reshape(-1)
247
+ if not show_outside:
248
+ keep = labels > 0
249
+ xyz = xyz[keep]
250
+ labels = labels[keep]
251
+ xyz, labels = _subsample(xyz, labels, MAX_PLOT_POINTS)
252
+ if xyz.shape[0] > 0:
253
+ colors = np.where(labels[:, None] > 0, COLOR_INSIDE, COLOR_OUTSIDE)
254
+ geoms.append(_points_mesh(xyz, colors))
255
+ return _export_glb(geoms)
src/gradio/orbit.js ADDED
@@ -0,0 +1,564 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /**
2
+ * Persistent occupancy orbit (Y-up). Gradio Model3D remounts the whole
3
+ * WebGL context on every new GLB — that is the gray flash + camera snap.
4
+ * This script owns one Engine for the page lifetime and only replaces meshes.
5
+ *
6
+ * Gradio 6 strips <script> inside gr.HTML. This file is passed to launch(js=)
7
+ * so it actually runs. Python writes a tiny #sn-cmd span (path + reset).
8
+ * The canvas host HTML is never updated, so the Engine is not remounted.
9
+ */
10
+ (async function () {
11
+ /* Gradio can re-run launch(js=) when the cmd span updates. A second
12
+ Engine would start at 45°/70° and snapshotCam() would wipe the orbit. */
13
+ if (window.__snOrbit && window.__snOrbit.alive && window.__snOrbit.alive()) {
14
+ return;
15
+ }
16
+ const BG = [42 / 255, 42 / 255, 50 / 255, 1];
17
+ const START_ALPHA = 45;
18
+ const START_BETA = 70;
19
+ /* Pull back so samples (dog / horse / torus) fit in frame. Load still
20
+ does not move the camera; Reset view returns here. */
21
+ const START_RADIUS = 60;
22
+ const GRID_HALF = 16;
23
+ const GRID_DIVS = 20;
24
+ const AXIS_LEN = 4;
25
+ const BABYLON_SRCS = [
26
+ [
27
+ "https://cdn.babylonjs.com/babylon.js",
28
+ "https://cdn.babylonjs.com/loaders/babylonjs.loaders.min.js",
29
+ ],
30
+ [
31
+ "https://cdn.jsdelivr.net/npm/babylonjs@7/babylon.js",
32
+ "https://cdn.jsdelivr.net/npm/babylonjs@7/babylonjs.loaders.min.js",
33
+ ],
34
+ ];
35
+
36
+ function loadScript(src) {
37
+ return new Promise(function (resolve, reject) {
38
+ const el = document.createElement("script");
39
+ el.src = src;
40
+ el.async = true;
41
+ el.onload = function () {
42
+ resolve();
43
+ };
44
+ el.onerror = function () {
45
+ reject(new Error(src));
46
+ };
47
+ document.head.appendChild(el);
48
+ });
49
+ }
50
+
51
+ async function loadBabylon() {
52
+ if (window.BABYLON && window.BABYLON.SceneLoader) {
53
+ return;
54
+ }
55
+ let last = null;
56
+ for (let i = 0; i < BABYLON_SRCS.length; i += 1) {
57
+ try {
58
+ await loadScript(BABYLON_SRCS[i][0]);
59
+ await loadScript(BABYLON_SRCS[i][1]);
60
+ if (window.BABYLON && window.BABYLON.SceneLoader) {
61
+ return;
62
+ }
63
+ } catch (err) {
64
+ last = err;
65
+ }
66
+ }
67
+ throw last || new Error("Babylon.js failed to load");
68
+ }
69
+
70
+ function waitFor(id) {
71
+ return new Promise(function (resolve) {
72
+ function tick() {
73
+ const el = document.getElementById(id);
74
+ if (el) {
75
+ resolve(el);
76
+ return;
77
+ }
78
+ window.setTimeout(tick, 50);
79
+ }
80
+ tick();
81
+ });
82
+ }
83
+
84
+ function cmdSpan() {
85
+ const direct = document.getElementById("sn-cmd");
86
+ if (direct) {
87
+ return direct;
88
+ }
89
+ const wrap = document.getElementById("sn-cmd-wrap");
90
+ return wrap ? wrap.querySelector("span") : null;
91
+ }
92
+
93
+ const DEFAULT_PSIZE = 8;
94
+
95
+ function cmdParts() {
96
+ const el = cmdSpan();
97
+ if (!el) {
98
+ return ["", 0, DEFAULT_PSIZE, 0, 0];
99
+ }
100
+ const bits = String(el.textContent || "").trim().split("|");
101
+ const path = bits[0] || "";
102
+ const reset = bits.length > 1 ? Number(bits[1]) : 0;
103
+ const psize = bits.length > 2 ? Number(bits[2]) : DEFAULT_PSIZE;
104
+ const wire = bits.length > 3 ? Number(bits[3]) : 0;
105
+ const forceDefault = bits.length > 4 ? Number(bits[4]) : 0;
106
+ return [path, reset, psize, wire, forceDefault];
107
+ }
108
+
109
+ function fieldPath() {
110
+ return cmdParts()[0];
111
+ }
112
+
113
+ function fieldReset() {
114
+ return cmdParts()[1] || 0;
115
+ }
116
+
117
+ function fieldPsize() {
118
+ const n = cmdParts()[2];
119
+ if (!Number.isFinite(n) || n < 1) {
120
+ return DEFAULT_PSIZE;
121
+ }
122
+ return Math.max(1, Math.min(24, n));
123
+ }
124
+
125
+ function fieldWire() {
126
+ return cmdParts()[3] ? 1 : 0;
127
+ }
128
+
129
+ function fieldForceDefault() {
130
+ return cmdParts()[4] ? 1 : 0;
131
+ }
132
+
133
+ function fileUrl(absPath) {
134
+ const posix = String(absPath).replace(/\\/g, "/");
135
+ return "/gradio_api/file=" + posix;
136
+ }
137
+
138
+ function pickObj(fileList) {
139
+ if (!fileList) {
140
+ return null;
141
+ }
142
+ for (let i = 0; i < fileList.length; i += 1) {
143
+ if (fileList[i] && /\.obj$/i.test(fileList[i].name || "")) {
144
+ return fileList[i];
145
+ }
146
+ }
147
+ return null;
148
+ }
149
+
150
+ function objFileInput() {
151
+ const root = document.getElementById("sn-obj-file");
152
+ return root ? root.querySelector("input[type=file]") : null;
153
+ }
154
+
155
+ function sendObjToGradio(file) {
156
+ const input = objFileInput();
157
+ if (!input || !file) {
158
+ return;
159
+ }
160
+ const dt = new DataTransfer();
161
+ dt.items.add(file);
162
+ input.value = "";
163
+ input.files = dt.files;
164
+ input.dispatchEvent(new Event("input", { bubbles: true }));
165
+ input.dispatchEvent(new Event("change", { bubbles: true }));
166
+ }
167
+
168
+ function bindLoadBtn(wrap) {
169
+ if (!wrap || wrap.dataset.snLoadBound === "1") {
170
+ return;
171
+ }
172
+ wrap.dataset.snLoadBound = "1";
173
+ wrap.addEventListener(
174
+ "click",
175
+ function () {
176
+ const input = objFileInput();
177
+ if (input) {
178
+ input.click();
179
+ }
180
+ },
181
+ true
182
+ );
183
+ }
184
+
185
+ function bindViewDrop(host) {
186
+ if (!host || host.dataset.snDropBound === "1") {
187
+ return;
188
+ }
189
+ host.dataset.snDropBound = "1";
190
+ host.addEventListener(
191
+ "dragover",
192
+ function (e) {
193
+ if (!e.dataTransfer) {
194
+ return;
195
+ }
196
+ e.preventDefault();
197
+ e.dataTransfer.dropEffect = "copy";
198
+ host.classList.add("sn-drop-over");
199
+ },
200
+ true
201
+ );
202
+ host.addEventListener(
203
+ "dragleave",
204
+ function (e) {
205
+ if (!host.contains(e.relatedTarget)) {
206
+ host.classList.remove("sn-drop-over");
207
+ }
208
+ },
209
+ true
210
+ );
211
+ host.addEventListener(
212
+ "drop",
213
+ function (e) {
214
+ e.preventDefault();
215
+ host.classList.remove("sn-drop-over");
216
+ sendObjToGradio(pickObj(e.dataTransfer && e.dataTransfer.files));
217
+ },
218
+ true
219
+ );
220
+ }
221
+
222
+ waitFor("sn-orbit-host").then(bindViewDrop);
223
+ waitFor("sn-load-obj").then(bindLoadBtn);
224
+ await loadBabylon();
225
+ if (window.BABYLON && window.BABYLON.SceneLoader) {
226
+ window.BABYLON.SceneLoader.OnPluginActivatedObservable.add(function (plugin) {
227
+ if (plugin && String(plugin.name || "").toLowerCase() === "gltf") {
228
+ plugin.loadCameras = false;
229
+ }
230
+ });
231
+ }
232
+ let canvas = document.getElementById("sn-orbit");
233
+ if (!canvas) {
234
+ /* Gradio may strip <canvas> from HTML; the host div is enough. */
235
+ let host = document.getElementById("sn-orbit-host");
236
+ if (!host) {
237
+ const wrap = await waitFor("sn-orbit-wrap");
238
+ host = document.createElement("div");
239
+ host.id = "sn-orbit-host";
240
+ host.style.cssText =
241
+ "width:100%;height:640px;background:#2a2a32;border-radius:8px;overflow:hidden;";
242
+ wrap.appendChild(host);
243
+ }
244
+ bindViewDrop(host);
245
+ canvas = document.createElement("canvas");
246
+ canvas.id = "sn-orbit";
247
+ canvas.style.cssText = "width:100%;height:100%;display:block;";
248
+ host.appendChild(canvas);
249
+ }
250
+ if (canvas && canvas.parentElement && !document.getElementById("sn-drop-hint")) {
251
+ const hint = document.createElement("div");
252
+ hint.id = "sn-drop-hint";
253
+ hint.textContent = "Drop an OBJ file here";
254
+ canvas.parentElement.appendChild(hint);
255
+ }
256
+ const BABYLON = window.BABYLON;
257
+
258
+ const engine = new BABYLON.Engine(canvas, true, { preserveDrawingBuffer: true }, true);
259
+ const scene = new BABYLON.Scene(engine);
260
+ scene.useRightHandedSystem = true;
261
+ scene.clearColor = new BABYLON.Color4(BG[0], BG[1], BG[2], BG[3]);
262
+ scene.ambientColor = new BABYLON.Color3(0.35, 0.35, 0.38);
263
+
264
+ const camera = new BABYLON.ArcRotateCamera(
265
+ "cam",
266
+ BABYLON.Tools.ToRadians(START_ALPHA),
267
+ BABYLON.Tools.ToRadians(START_BETA),
268
+ START_RADIUS,
269
+ BABYLON.Vector3.Zero(),
270
+ scene
271
+ );
272
+ camera.lowerRadiusLimit = 0.05;
273
+ camera.minZ = 0.01;
274
+ camera.wheelPrecision = 40;
275
+ camera.attachControl(canvas, true);
276
+
277
+ const light = new BABYLON.HemisphericLight("hemi", new BABYLON.Vector3(0.2, 1, 0.15), scene);
278
+ light.intensity = 0.95;
279
+
280
+ /* Fixed world floor. Not part of the GLB — a per-mesh grid is what
281
+ made the camera look like it jumped when a new OBJ loaded. */
282
+ (function addWorldHelpers() {
283
+ const half = GRID_HALF;
284
+ const n = GRID_DIVS;
285
+ const mid = n / 2;
286
+ const cell = new BABYLON.Color3(85 / 255, 102 / 255, 119 / 255);
287
+ const center = new BABYLON.Color3(170 / 255, 187 / 255, 204 / 255);
288
+ for (let i = 0; i <= n; i += 1) {
289
+ const t = -half + (2 * half * i) / n;
290
+ const col = i === mid ? center : cell;
291
+ const gx = BABYLON.MeshBuilder.CreateLines(
292
+ "sn_gx_" + i,
293
+ { points: [new BABYLON.Vector3(t, 0, -half), new BABYLON.Vector3(t, 0, half)] },
294
+ scene
295
+ );
296
+ gx.color = col;
297
+ const gz = BABYLON.MeshBuilder.CreateLines(
298
+ "sn_gz_" + i,
299
+ { points: [new BABYLON.Vector3(-half, 0, t), new BABYLON.Vector3(half, 0, t)] },
300
+ scene
301
+ );
302
+ gz.color = col;
303
+ }
304
+ const axis = [
305
+ [new BABYLON.Vector3(AXIS_LEN, 0, 0), new BABYLON.Color3(1, 0, 0)],
306
+ [new BABYLON.Vector3(0, AXIS_LEN, 0), new BABYLON.Color3(0, 1, 0)],
307
+ [new BABYLON.Vector3(0, 0, AXIS_LEN), new BABYLON.Color3(0, 0, 1)],
308
+ ];
309
+ for (let i = 0; i < axis.length; i += 1) {
310
+ const line = BABYLON.MeshBuilder.CreateLines(
311
+ "sn_ax_" + i,
312
+ { points: [BABYLON.Vector3.Zero(), axis[i][0]] },
313
+ scene
314
+ );
315
+ line.color = axis[i][1];
316
+ }
317
+ })();
318
+
319
+ let imported = [];
320
+ let lastPath = "";
321
+ let lastReset = null;
322
+ let lastPsize = null;
323
+ let lastWire = null;
324
+ let loadGen = 0;
325
+ let loadBusy = false;
326
+
327
+ function applyCam(shot) {
328
+ if (!shot) {
329
+ return;
330
+ }
331
+ if (camera.inertialAlphaOffset !== undefined) {
332
+ camera.inertialAlphaOffset = 0;
333
+ camera.inertialBetaOffset = 0;
334
+ camera.inertialRadiusOffset = 0;
335
+ }
336
+ if (camera.inertialPanningX !== undefined) {
337
+ camera.inertialPanningX = 0;
338
+ camera.inertialPanningY = 0;
339
+ }
340
+ camera.alpha = shot.a;
341
+ camera.beta = shot.b;
342
+ camera.radius = shot.r;
343
+ camera.setTarget(new BABYLON.Vector3(shot.tx, shot.ty, shot.tz));
344
+ scene.activeCamera = camera;
345
+ }
346
+
347
+ function applyDefaultCam() {
348
+ applyCam({
349
+ a: BABYLON.Tools.ToRadians(START_ALPHA),
350
+ b: BABYLON.Tools.ToRadians(START_BETA),
351
+ r: START_RADIUS,
352
+ tx: 0,
353
+ ty: 0,
354
+ tz: 0,
355
+ });
356
+ }
357
+
358
+ function frameDefault() {
359
+ applyDefaultCam();
360
+ }
361
+
362
+ function isOccPoints(mesh) {
363
+ /* Grid/axes are GL_LINES with idx.length === vert count. Treating those
364
+ as a point cloud turns the floor into dots. Only occupancy points. */
365
+ if (!mesh) {
366
+ return false;
367
+ }
368
+ return String(mesh.name || "").indexOf("occ_points") >= 0;
369
+ }
370
+
371
+ function isOccShell(mesh) {
372
+ if (!mesh) {
373
+ return false;
374
+ }
375
+ return String(mesh.name || "").indexOf("occ_shell") >= 0;
376
+ }
377
+
378
+ function applyWireframe(on) {
379
+ /* Inspect viewer: THREE.EdgesGeometry(geom, 25) on the solid, not
380
+ material.wireframe (that draws every triangle). Babylon's epsilon is
381
+ cos(threshold): draw an edge when adjacent face normals diverge. */
382
+ const want = !!on;
383
+ const epsilon = Math.cos((25 * Math.PI) / 180);
384
+ const edgeColor = new BABYLON.Color4(228 / 255, 232 / 255, 238 / 255, 0.9);
385
+ for (let i = 0; i < imported.length; i += 1) {
386
+ const mesh = imported[i];
387
+ if (!isOccShell(mesh)) {
388
+ continue;
389
+ }
390
+ const mats = mesh.material
391
+ ? Array.isArray(mesh.material)
392
+ ? mesh.material
393
+ : [mesh.material]
394
+ : [];
395
+ for (let m = 0; m < mats.length; m += 1) {
396
+ if (mats[m]) {
397
+ mats[m].wireframe = false;
398
+ }
399
+ }
400
+ if (want) {
401
+ mesh.enableEdgesRendering(epsilon);
402
+ mesh.edgesWidth = 1.5;
403
+ mesh.edgesColor = edgeColor;
404
+ } else if (mesh.disableEdgesRendering) {
405
+ mesh.disableEdgesRendering();
406
+ }
407
+ }
408
+ }
409
+ function applyPointSize(px) {
410
+ const size = Math.max(1, Math.min(24, Number(px) || DEFAULT_PSIZE));
411
+ for (let i = 0; i < imported.length; i += 1) {
412
+ const mesh = imported[i];
413
+ if (!isOccPoints(mesh)) {
414
+ continue;
415
+ }
416
+ const mats = mesh.material
417
+ ? Array.isArray(mesh.material)
418
+ ? mesh.material
419
+ : [mesh.material]
420
+ : [];
421
+ for (let m = 0; m < mats.length; m += 1) {
422
+ const mat = mats[m];
423
+ if (!mat) {
424
+ continue;
425
+ }
426
+ mat.pointsCloud = true;
427
+ mat.pointSize = size;
428
+ mat.disableLighting = true;
429
+ if (mat.unlit !== undefined) {
430
+ mat.unlit = true;
431
+ }
432
+ }
433
+ }
434
+ }
435
+
436
+ function clearImported() {
437
+ for (let i = 0; i < imported.length; i += 1) {
438
+ if (imported[i] && imported[i].dispose) {
439
+ imported[i].dispose(false, true);
440
+ }
441
+ }
442
+ imported = [];
443
+ }
444
+
445
+ async function loadGlb(url, px, wire) {
446
+ /* Freeze the orbit. Import must not frame, target, or replace the camera. */
447
+ const pose = {
448
+ a: camera.alpha,
449
+ b: camera.beta,
450
+ r: camera.radius,
451
+ tx: camera.target.x,
452
+ ty: camera.target.y,
453
+ tz: camera.target.z,
454
+ };
455
+ const gen = (loadGen += 1);
456
+ let result;
457
+ try {
458
+ result = await BABYLON.SceneLoader.ImportMeshAsync("", url, "", scene);
459
+ } catch (err) {
460
+ applyCam(pose);
461
+ throw err;
462
+ }
463
+ if (gen !== loadGen) {
464
+ (result.meshes || []).forEach(function (m) {
465
+ m.dispose(false, true);
466
+ });
467
+ applyCam(pose);
468
+ return;
469
+ }
470
+ clearImported();
471
+ imported = [];
472
+ (result.meshes || []).forEach(function (m) {
473
+ const n = String(m.name || "");
474
+ if (n.indexOf("occ_empty") >= 0) {
475
+ m.dispose(false, true);
476
+ return;
477
+ }
478
+ imported.push(m);
479
+ });
480
+ (result.cameras || []).forEach(function (cam) {
481
+ if (cam && cam.dispose) {
482
+ cam.dispose();
483
+ }
484
+ });
485
+ scene.activeCamera = camera;
486
+ applyPointSize(px);
487
+ applyWireframe(wire);
488
+ applyCam(pose);
489
+ }
490
+
491
+ async function tick() {
492
+ if (loadBusy) {
493
+ return;
494
+ }
495
+ const path = fieldPath();
496
+ const resetN = fieldReset();
497
+ const px = fieldPsize();
498
+ const wire = fieldWire();
499
+ if (!path) {
500
+ return;
501
+ }
502
+ const pathChanged = path !== lastPath;
503
+ const resetChanged = lastReset !== null && resetN !== lastReset;
504
+ const sizeChanged = lastPsize !== null && px !== lastPsize;
505
+ const wireChanged = lastWire !== null && wire !== lastWire;
506
+ if (!pathChanged && !resetChanged && !sizeChanged && !wireChanged) {
507
+ lastReset = resetN;
508
+ lastPsize = px;
509
+ lastWire = wire;
510
+ return;
511
+ }
512
+ if (pathChanged) {
513
+ lastPath = path;
514
+ lastReset = resetN;
515
+ lastPsize = px;
516
+ lastWire = wire;
517
+ loadBusy = true;
518
+ try {
519
+ await loadGlb(fileUrl(path), px, wire);
520
+ } catch (err) {
521
+ console.error("sn-orbit load", err);
522
+ } finally {
523
+ loadBusy = false;
524
+ }
525
+ return;
526
+ }
527
+ lastReset = resetN;
528
+ lastPsize = px;
529
+ lastWire = wire;
530
+ if (resetChanged) {
531
+ applyDefaultCam();
532
+ }
533
+ if (sizeChanged || resetChanged) {
534
+ applyPointSize(px);
535
+ }
536
+ if (wireChanged || resetChanged) {
537
+ applyWireframe(wire);
538
+ }
539
+ }
540
+
541
+ window.addEventListener("resize", function () {
542
+ engine.resize();
543
+ });
544
+ engine.runRenderLoop(function () {
545
+ scene.render();
546
+ });
547
+
548
+ applyDefaultCam();
549
+
550
+ window.__snOrbit = {
551
+ alive: function () {
552
+ return !!(
553
+ engine &&
554
+ !engine.isDisposed &&
555
+ canvas &&
556
+ canvas.isConnected &&
557
+ engine.getRenderingCanvas() === canvas
558
+ );
559
+ },
560
+ };
561
+
562
+ await tick();
563
+ window.setInterval(tick, 200);
564
+ })();
src/gradio/pipeline.py ADDED
@@ -0,0 +1,133 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Gradio occupancy loop: fill an OBJ AABB, then classify with ``best.pt``.
2
+
3
+ Reuses the viewer helper (``obj_fill``, ``infer_job``, ``model_access``).
4
+ Does not import Gradio and does not change the Three.js viewer.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ import base64
10
+ from pathlib import Path
11
+ from typing import Any
12
+
13
+ import numpy as np
14
+
15
+ # Occupancy is the installed ``scatteringnet`` package (``pip install -e .``).
16
+ # This folder is not a package so it cannot shadow pip ``gradio``.
17
+ _GRADIO_DIR = Path(__file__).resolve().parent
18
+ _REPO = _GRADIO_DIR.parents[1]
19
+
20
+ from scatteringnet.config import load_config
21
+ from scatteringnet.viewer.infer_job import infer_uploaded_obj, pred_from_probs # noqa: E402
22
+ from scatteringnet.viewer.model_access import inspect_run_id, list_viewer_models # noqa: E402
23
+ from scatteringnet.viewer.obj_fill import ( # noqa: E402
24
+ fill_aabb_lattice,
25
+ spacing_from_slider,
26
+ triangles_from_obj_text,
27
+ )
28
+
29
+ # INSPECT alias from docs/inspect_checkpoint.yaml (not the newest best.pt).
30
+ INSPECT_RUN_ID = inspect_run_id()
31
+ # Viewer allows 200k; Plotly + Spaces need a tighter lattice.
32
+ MAX_FILL_POINTS = 80_000
33
+ MAX_OBJ_BYTES = 32 * 1024 * 1024
34
+ # Slider default matches the Three.js Density control (spacing ~0.15).
35
+ DEFAULT_DENSITY = 71
36
+ DEFAULT_CUT = 0.50
37
+
38
+
39
+ def models_root() -> Path:
40
+ """Repo ``models/`` (``models/<run_id>/best.pt``)."""
41
+ return _REPO / "models"
42
+
43
+
44
+ def list_run_ids(root: Path | None = None) -> list[str]:
45
+ """Newest-first run folder names that have a ``best.pt``."""
46
+ rows = list_viewer_models(root or models_root())
47
+ return [str(row["id"]) for row in rows]
48
+
49
+
50
+ def default_run_id(root: Path | None = None) -> str:
51
+ """INSPECT pointer if that ``best.pt`` exists, else newest, else empty."""
52
+ ids = list_run_ids(root)
53
+ wanted = inspect_run_id()
54
+ if wanted and wanted in ids:
55
+ return wanted
56
+ return ids[0] if ids else ""
57
+
58
+
59
+ def runtime_cfg(cfg: Any | None = None):
60
+ """
61
+ Device / batch from YAML. ``data_dir`` is not required (uploaded OBJ).
62
+ """
63
+ if cfg is not None:
64
+ return cfg
65
+ return load_config(require_existing_data_dir=False)
66
+
67
+
68
+ def apply_cut(probs: np.ndarray, threshold: float) -> tuple[np.ndarray, int, int]:
69
+ """Hard labels from stored sigmoid probs (no extra GPU pass)."""
70
+ pred = pred_from_probs(probs, threshold)
71
+ n_in = int(np.count_nonzero(pred > 0))
72
+ return pred, n_in, int(pred.shape[0]) - n_in
73
+
74
+
75
+ def fill_and_infer(
76
+ obj_text: str,
77
+ *,
78
+ obj_name: str,
79
+ run_id: str,
80
+ density: float,
81
+ models: Path | None = None,
82
+ cfg: Any | None = None,
83
+ ) -> dict[str, Any]:
84
+ """
85
+ Job B for Gradio: AABB lattice + occupancy forward.
86
+
87
+ Returns vertices/faces (for the Plotly mesh), query XYZ, sigmoid
88
+ probs, timings, and the run id. The UI recuts ``probs`` locally.
89
+ """
90
+ raw = str(obj_text or "")
91
+ if not raw.strip():
92
+ raise ValueError("OBJ is empty")
93
+ if len(raw.encode("utf-8")) > MAX_OBJ_BYTES:
94
+ raise ValueError(f"OBJ too large (max {MAX_OBJ_BYTES} bytes)")
95
+ name = str(run_id or "").strip()
96
+ if not name:
97
+ raise ValueError("select a checkpoint under models/<run_id>/best.pt")
98
+
99
+ vertices, faces = triangles_from_obj_text(raw)
100
+ spacing = spacing_from_slider(density)
101
+ points, used, grid = fill_aabb_lattice(
102
+ vertices, spacing, max_points=MAX_FILL_POINTS
103
+ )
104
+ ckpt_root = Path(models) if models is not None else models_root()
105
+ out = infer_uploaded_obj(
106
+ run_id=name,
107
+ models_root=ckpt_root,
108
+ data_dir=None,
109
+ obj_name=str(obj_name or "upload.obj"),
110
+ obj_text=raw,
111
+ points=points,
112
+ cfg=runtime_cfg(cfg),
113
+ )
114
+ probs = np.frombuffer(base64.b64decode(out["prob_b64"]), dtype=np.float32).reshape(
115
+ -1
116
+ )
117
+ pred, n_in, n_out = apply_cut(probs, DEFAULT_CUT)
118
+ return {
119
+ "vertices": np.ascontiguousarray(vertices, dtype=np.float32),
120
+ "faces": np.ascontiguousarray(faces, dtype=np.int32),
121
+ "points": np.ascontiguousarray(points, dtype=np.float32),
122
+ "probs": np.ascontiguousarray(probs, dtype=np.float32),
123
+ "pred": pred,
124
+ "n": int(out["n"]),
125
+ "n_inside": n_in,
126
+ "n_outside": n_out,
127
+ "used_spacing": float(used),
128
+ "grid": [int(grid[0]), int(grid[1]), int(grid[2])],
129
+ "run_id": str(out["run_id"]),
130
+ "shape_encoder": str(out["shape_encoder"]),
131
+ "timings": dict(out["timings"]),
132
+ "obj_name": str(obj_name or "upload.obj"),
133
+ }
src/infer_multi_npz.py ADDED
@@ -0,0 +1,202 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Classify occupancy NPZs from ``models/<run_id>/best.pt``.
2
+
3
+ Rebuilds ``OccupancyMLP`` or ``OccupancyEncoder`` from the checkpoint
4
+ ``kind``. Query XYZ (and envelope, when conditioned) use the stored
5
+ per-mesh AABB, not a fresh map. Face-token checkpoints are rejected.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ from pathlib import Path
11
+ from typing import Any
12
+
13
+ import numpy as np
14
+ import torch
15
+
16
+ from scatteringnet.config import OccupancyConfig, as_data_relative, load_config, repo_root
17
+ from scatteringnet.data_npz import load_points_labels, load_points_labels_mesh
18
+ from scatteringnet.geometry.mesh_io import load_obj_triangles
19
+ from scatteringnet.geometry.surface import apply_envelope_aabb, project_envelope_dim, sample_surface_points
20
+ from scatteringnet.metrics import occupancy_metrics
21
+ from scatteringnet.normalize import apply_normalization
22
+ from scatteringnet.occupancy_encoder import (
23
+ OccupancyEncoder,
24
+ envelope_dim_from_ckpt,
25
+ envelope_seed_from_ckpt,
26
+ )
27
+ from scatteringnet.occupancy_encoder import CHECKPOINT_KIND as ENCODER_KIND
28
+ from scatteringnet.occupancy_mlp import OccupancyMLP
29
+ from scatteringnet.occupancy_mlp import CHECKPOINT_KIND as MLP_KIND
30
+
31
+
32
+ def resolve_best_pt(
33
+ *,
34
+ checkpoint: Path | str | None = None,
35
+ run_id: str | None = None,
36
+ root: Path | None = None,
37
+ ) -> Path:
38
+ """Resolve ``best.pt`` from an explicit path or ``runs/<id>/checkpoint_dir.txt``."""
39
+ base = (root or repo_root()).resolve()
40
+ if checkpoint is not None:
41
+ path = Path(checkpoint)
42
+ if not path.is_absolute():
43
+ path = base / path
44
+ if not path.is_file():
45
+ raise FileNotFoundError(f"checkpoint not found: {path}")
46
+ return path
47
+ if not run_id:
48
+ raise ValueError("pass checkpoint= or run_id=")
49
+ pointer = base / "runs" / str(run_id) / "checkpoint_dir.txt"
50
+ if not pointer.is_file():
51
+ raise FileNotFoundError(f"run pointer not found: {pointer}")
52
+ rel = pointer.read_text(encoding="utf-8").strip()
53
+ path = Path(rel)
54
+ if not path.is_absolute():
55
+ path = base / path
56
+ best = path / "best.pt" if path.is_dir() else path
57
+ if not best.is_file():
58
+ raise FileNotFoundError(f"best.pt not found: {best}")
59
+ return best
60
+
61
+
62
+ def load_occupancy_model(
63
+ ckpt: dict[str, Any],
64
+ device: torch.device,
65
+ ) -> OccupancyMLP | OccupancyEncoder:
66
+ """Rebuild the head recorded in ``ckpt['kind']`` and load weights."""
67
+ hidden = int(ckpt["hidden"])
68
+ depth = int(ckpt["depth"])
69
+ kind = str(ckpt["kind"])
70
+ if kind == MLP_KIND:
71
+ model: OccupancyMLP | OccupancyEncoder = OccupancyMLP(hidden=hidden, depth=depth)
72
+ elif kind == ENCODER_KIND:
73
+ latent = int(ckpt["latent_dim"]) if ckpt.get("latent_dim") is not None else hidden
74
+ enc = str(ckpt.get("shape_encoder") or "surface").strip().lower()
75
+ if enc == "mesh":
76
+ raise ValueError(
77
+ "face-token occupancy checkpoints (shape_encoder='mesh') "
78
+ "are no longer supported"
79
+ )
80
+ knn_k = int(ckpt["knn_k"]) if ckpt.get("knn_k") is not None else 0
81
+ knn_local = None
82
+ if knn_k > 0 and ckpt.get("knn_local_dim") is not None:
83
+ knn_local = int(ckpt["knn_local_dim"])
84
+ model = OccupancyEncoder(
85
+ hidden=hidden,
86
+ depth=depth,
87
+ latent_dim=latent,
88
+ shape_encoder=enc,
89
+ knn_k=knn_k,
90
+ knn_local_dim=knn_local,
91
+ envelope_dim=envelope_dim_from_ckpt(ckpt),
92
+ )
93
+ else:
94
+ raise ValueError(f"unsupported checkpoint kind {kind!r}")
95
+ model.load_state_dict(ckpt["state_dict"])
96
+ return model.to(device).eval()
97
+
98
+
99
+ def _part_for_npz(ckpt: dict[str, Any], npz_path: Path, data_dir: Path) -> dict[str, Any]:
100
+ rel = as_data_relative(npz_path, data_dir)
101
+ name = npz_path.name
102
+ for part in ckpt.get("parts") or []:
103
+ stored = str(part.get("npz", ""))
104
+ if stored == rel or Path(stored).name == name:
105
+ return part
106
+ raise KeyError(f"no AABB part for {npz_path.name} in checkpoint")
107
+
108
+
109
+ def infer_npz(
110
+ cfg: OccupancyConfig,
111
+ *,
112
+ checkpoint: Path | str | None = None,
113
+ run_id: str | None = None,
114
+ npz_path: Path | str | None = None,
115
+ root: Path | None = None,
116
+ ) -> dict[str, float]:
117
+ """
118
+ Score one NPZ with ``best.pt``. Returns accuracy / IoU / F1.
119
+ """
120
+ ckpt_path = resolve_best_pt(checkpoint=checkpoint, run_id=run_id, root=root)
121
+ ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=False)
122
+ model = load_occupancy_model(ckpt, cfg.device)
123
+ if npz_path is None:
124
+ rels = ckpt.get("npz_paths") or []
125
+ if not rels:
126
+ raise ValueError("checkpoint has no npz_paths; pass npz_path=")
127
+ npz = cfg.data_dir / str(rels[0])
128
+ else:
129
+ npz = Path(npz_path)
130
+ if not npz.is_absolute():
131
+ npz = cfg.data_dir / npz
132
+ part = _part_for_npz(ckpt, npz, cfg.data_dir)
133
+ center = np.asarray(part["center"], dtype=np.float32)
134
+ scale = float(part["scale"])
135
+ geom = None
136
+ shape_id = None
137
+ enc = str(ckpt.get("shape_encoder", "none")).strip().lower()
138
+ if enc == "mesh":
139
+ raise ValueError(
140
+ "face-token occupancy checkpoints (shape_encoder='mesh') "
141
+ "are no longer supported"
142
+ )
143
+ if enc == "surface":
144
+ points, labels, mesh_path = load_points_labels_mesh(npz, cfg.data_dir)
145
+ vertices, faces = load_obj_triangles(mesh_path)
146
+ cache_key = str(mesh_path.resolve())
147
+ n_surface = int(ckpt.get("n_surface") or cfg.n_surface)
148
+ world = sample_surface_points(
149
+ vertices,
150
+ faces,
151
+ n_surface,
152
+ seed=envelope_seed_from_ckpt(ckpt),
153
+ cache_key=cache_key,
154
+ )
155
+ arr = apply_envelope_aabb(world, center, scale)
156
+ arr = project_envelope_dim(arr, envelope_dim_from_ckpt(ckpt))
157
+ geom = torch.from_numpy(arr).unsqueeze(0)
158
+ shape_id = torch.zeros((), dtype=torch.long)
159
+ else:
160
+ points, labels = load_points_labels(npz)
161
+ xyz = torch.from_numpy(apply_normalization(points, center, scale))
162
+ y = torch.from_numpy(np.asarray(labels, dtype=np.float32)).unsqueeze(1)
163
+ logits_rows: list[torch.Tensor] = []
164
+ with torch.no_grad():
165
+ for start in range(0, int(xyz.shape[0]), int(cfg.batch_size)):
166
+ sl = slice(start, start + int(cfg.batch_size))
167
+ batch_xyz = xyz[sl].to(cfg.device)
168
+ if geom is None:
169
+ logits_rows.append(model(batch_xyz).cpu())
170
+ else:
171
+ b = int(batch_xyz.shape[0])
172
+ geom_b = geom.expand(b, -1, -1).to(cfg.device)
173
+ sid = shape_id.expand(b).to(cfg.device)
174
+ logits_rows.append(model(batch_xyz, geom_b, sid).cpu())
175
+ logits = torch.cat(logits_rows, dim=0)
176
+ scores = occupancy_metrics(logits, y)
177
+ print(
178
+ f"checkpoint={ckpt_path} npz={npz.name} "
179
+ f"acc={scores.accuracy:.4f} iou={scores.inside_iou:.4f} "
180
+ f"f1={scores.inside_f1:.4f}"
181
+ )
182
+ return {
183
+ "accuracy": scores.accuracy,
184
+ "inside_iou": scores.inside_iou,
185
+ "inside_f1": scores.inside_f1,
186
+ }
187
+
188
+
189
+ if __name__ == "__main__":
190
+ import argparse
191
+
192
+ parser = argparse.ArgumentParser(description="Infer occupancy from models/<id>/best.pt")
193
+ parser.add_argument("--checkpoint", default=None, help="Path to best.pt")
194
+ parser.add_argument("--run-id", default=None, help="runs/<id> folder name")
195
+ parser.add_argument("--npz", default=None, help="NPZ path (data_dir-relative or absolute)")
196
+ args = parser.parse_args()
197
+ infer_npz(
198
+ load_config(),
199
+ checkpoint=args.checkpoint,
200
+ run_id=args.run_id,
201
+ npz_path=args.npz,
202
+ )
src/metrics.py ADDED
@@ -0,0 +1,156 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Occupancy classification metrics from logits vs labels.
2
+
3
+ The model emits raw logits (no sigmoid in ``forward``). Metrics apply
4
+ sigmoid only at decision time so they stay consistent with
5
+ ``BCEWithLogitsLoss`` and with inference (threshold ``0.5``).
6
+
7
+ Inside (label ``1``) is the product-relevant class: points that belong
8
+ in the interior. Precision / recall are therefore reported for that
9
+ class only. No extra packages (sklearn, etc.).
10
+ """
11
+
12
+ from __future__ import annotations
13
+
14
+ from dataclasses import dataclass
15
+
16
+ import torch
17
+ from torch import Tensor
18
+
19
+
20
+ @dataclass(frozen=True)
21
+ class OccupancyMetrics:
22
+ """Pointwise occupancy scores for one batch or a full set of queries."""
23
+
24
+ accuracy: float
25
+ inside_precision: float
26
+ inside_recall: float
27
+ inside_iou: float
28
+ inside_f1: float
29
+
30
+
31
+ def _safe_div(numerator: float, denominator: float) -> float:
32
+ """Return ``0.0`` when the count in the denominator is zero (no sklearn)."""
33
+ if denominator <= 0.0:
34
+ return 0.0
35
+ return numerator / denominator
36
+
37
+
38
+ def _flatten_pair(logits: Tensor, labels: Tensor) -> tuple[Tensor, Tensor]:
39
+ """
40
+ Collapse ``(B, 1)`` or ``(B,)`` logits / labels to a shared 1-D view.
41
+
42
+ Last dim of logits is 1 when it comes from ``OccupancyMLP``; labels from
43
+ the Dataset match that. A 1-D label vector is accepted so callers do not
44
+ have to unsqueeze.
45
+ """
46
+ if logits.numel() != labels.numel():
47
+ raise ValueError(
48
+ f"logits and labels must have the same number of elements, "
49
+ f"got logits={tuple(logits.shape)} labels={tuple(labels.shape)}"
50
+ )
51
+ return logits.reshape(-1), labels.reshape(-1)
52
+
53
+
54
+ def occupancy_metrics(
55
+ logits: Tensor,
56
+ labels: Tensor,
57
+ *,
58
+ threshold: float = 0.5,
59
+ ) -> OccupancyMetrics:
60
+ """
61
+ Accuracy plus inside precision / recall from occupancy logits.
62
+
63
+ Parameters
64
+ ----------
65
+ logits:
66
+ Unnormalized scores, shape ``(B, 1)`` or ``(B,)``. Positive → inside.
67
+ labels:
68
+ Float ``{0, 1}`` with the same number of elements as ``logits``.
69
+ threshold:
70
+ Decision cut on ``sigmoid(logit)``. Default ``0.5`` matches inference.
71
+
72
+ Returns
73
+ -------
74
+ OccupancyMetrics
75
+ Scalar floats on CPU (safe to print or average across batches).
76
+ """
77
+ # Sigmoid here only: training still uses BCE-with-logits on raw logits.
78
+ tp, fp, fn, correct, n = occupancy_counts(logits, labels, threshold=threshold)
79
+ return occupancy_metrics_from_counts(tp=tp, fp=fp, fn=fn, correct=correct, n=n)
80
+
81
+
82
+ def occupancy_counts(
83
+ logits: Tensor,
84
+ labels: Tensor,
85
+ *,
86
+ threshold: float = 0.5,
87
+ ) -> tuple[float, float, float, float, float]:
88
+ """Return ``(tp, fp, fn, correct, n)`` for a micro-average over points."""
89
+ logits_flat, labels_flat = _flatten_pair(logits, labels)
90
+ pred_inside = logits_flat.sigmoid() >= threshold
91
+ true_inside = labels_flat > 0.5
92
+ pred_f = pred_inside.to(dtype=torch.float32)
93
+ true_f = true_inside.to(dtype=torch.float32)
94
+ tp = float((pred_f * true_f).sum().item())
95
+ fp = float((pred_f * (1.0 - true_f)).sum().item())
96
+ fn = float(((1.0 - pred_f) * true_f).sum().item())
97
+ correct = float((pred_inside == true_inside).to(dtype=torch.float32).sum().item())
98
+ n = float(pred_inside.numel())
99
+ return tp, fp, fn, correct, n
100
+
101
+
102
+ def occupancy_metrics_from_counts(
103
+ *,
104
+ tp: float,
105
+ fp: float,
106
+ fn: float,
107
+ correct: float,
108
+ n: float,
109
+ ) -> OccupancyMetrics:
110
+ """Build metrics from accumulated confusion counts (point micro-average)."""
111
+ precision = _safe_div(tp, tp + fp)
112
+ recall = _safe_div(tp, tp + fn)
113
+ return OccupancyMetrics(
114
+ accuracy=_safe_div(correct, n),
115
+ inside_precision=precision,
116
+ inside_recall=recall,
117
+ inside_iou=_safe_div(tp, tp + fp + fn),
118
+ inside_f1=_safe_div(2.0 * precision * recall, precision + recall),
119
+ )
120
+
121
+
122
+ def accuracy_from_logits(
123
+ logits: Tensor,
124
+ labels: Tensor,
125
+ *,
126
+ threshold: float = 0.5,
127
+ ) -> float:
128
+ """Fraction of points whose thresholded sigmoid matches the 0/1 label.
129
+
130
+ Parameters
131
+ ----------
132
+ logits, labels, threshold:
133
+ Same meaning as :func:`occupancy_metrics`.
134
+
135
+ Returns
136
+ -------
137
+ float
138
+ Accuracy in ``[0, 1]``.
139
+ """
140
+ return occupancy_metrics(logits, labels, threshold=threshold).accuracy
141
+
142
+
143
+ if __name__ == "__main__":
144
+ # Deterministic 8-point batch: 4 true insides, 4 true outsides, all correct.
145
+ demo_logits = torch.tensor(
146
+ [[4.0], [3.0], [2.0], [1.0], [-1.0], [-2.0], [-3.0], [-4.0]]
147
+ )
148
+ demo_labels = torch.tensor(
149
+ [[1.0], [1.0], [1.0], [1.0], [0.0], [0.0], [0.0], [0.0]]
150
+ )
151
+ scores = occupancy_metrics(demo_logits, demo_labels)
152
+ print(f"accuracy={scores.accuracy:.4f}")
153
+ print(f"inside_precision={scores.inside_precision:.4f}")
154
+ print(f"inside_recall={scores.inside_recall:.4f}")
155
+ print(f"inside_iou={scores.inside_iou:.4f}")
156
+ print(f"inside_f1={scores.inside_f1:.4f}")
src/normalize.py ADDED
@@ -0,0 +1,104 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """AABB normalization for occupancy XYZ.
2
+
3
+ Maps a point cloud into a roughly ``[-1, 1]^3`` cube so the MLP sees
4
+ comparable coordinates across differently sized meshes.
5
+
6
+ center = midpoint of the axis-aligned bounding box
7
+ scale = maximum half-extent (longest AABB side / 2)
8
+
9
+ Normalized point: ``(xyz - center) / scale``.
10
+ This module does not touch the occupancy model.
11
+ """
12
+
13
+ from __future__ import annotations
14
+
15
+ import numpy as np
16
+ from numpy.typing import NDArray
17
+
18
+ PointsArray = NDArray[np.float32]
19
+
20
+
21
+ def compute_center_scale(points: np.ndarray) -> tuple[PointsArray, float]:
22
+ """
23
+ AABB center and max half-extent for an ``(N, 3)`` point array.
24
+
25
+ Parameters
26
+ ----------
27
+ points:
28
+ Query XYZ, shape ``(N, 3)``, at least one row.
29
+
30
+ Returns
31
+ -------
32
+ center:
33
+ ``float32`` vector of shape ``(3,)``.
34
+ scale:
35
+ Positive float (max half-extent). Raises if the cloud has no extent.
36
+ """
37
+ pts = np.asarray(points, dtype=np.float32)
38
+ if pts.ndim != 2 or pts.shape[1] != 3:
39
+ raise ValueError(f"points must have shape (N, 3), got {tuple(pts.shape)}")
40
+ if pts.shape[0] == 0:
41
+ raise ValueError("points must contain at least one row")
42
+
43
+ xyz_min = pts.min(axis=0)
44
+ xyz_max = pts.max(axis=0)
45
+ center = 0.5 * (xyz_min + xyz_max)
46
+ half_extents = 0.5 * (xyz_max - xyz_min)
47
+ scale = float(np.max(half_extents))
48
+ if scale <= 0.0:
49
+ raise ValueError(
50
+ "scale must be > 0; all points appear to share the same location"
51
+ )
52
+ return center.astype(np.float32, copy=False), scale
53
+
54
+
55
+ def apply_normalization(
56
+ points: np.ndarray,
57
+ center: np.ndarray,
58
+ scale: float,
59
+ ) -> PointsArray:
60
+ """
61
+ Return ``(points - center) / scale`` as ``float32 (N, 3)``.
62
+
63
+ Parameters
64
+ ----------
65
+ points:
66
+ Query XYZ, shape ``(N, 3)``.
67
+ center:
68
+ AABB midpoint, shape ``(3,)``.
69
+ scale:
70
+ Positive max half-extent.
71
+
72
+ Returns
73
+ -------
74
+ ndarray
75
+ Normalized points, ``float32 (N, 3)``.
76
+ """
77
+ if scale <= 0.0:
78
+ raise ValueError(f"scale must be > 0, got {scale}")
79
+ pts = np.asarray(points, dtype=np.float32)
80
+ if pts.ndim != 2 or pts.shape[1] != 3:
81
+ raise ValueError(f"points must have shape (N, 3), got {tuple(pts.shape)}")
82
+ c = np.asarray(center, dtype=np.float32).reshape(3)
83
+ return (pts - c) / np.float32(scale)
84
+
85
+
86
+ if __name__ == "__main__":
87
+ from scatteringnet.config import load_config
88
+ from scatteringnet.data_npz import load_points_labels
89
+
90
+ sample = (
91
+ load_config().data_dir
92
+ / "exports"
93
+ / "dataset_test"
94
+ / "sphere__raycast_z_raut_s0.15_inout.npz"
95
+ )
96
+ points, _labels = load_points_labels(sample)
97
+ center, scale = compute_center_scale(points)
98
+ normed = apply_normalization(points, center, scale)
99
+ recovered = normed[:3] * np.float32(scale) + center
100
+ print(f"file={sample}")
101
+ print(f"center={center.tolist()} scale={scale:.6f}")
102
+ print(f"normed_min={normed.min(axis=0).tolist()}")
103
+ print(f"normed_max={normed.max(axis=0).tolist()}")
104
+ print(f"inverse_ok={np.allclose(recovered, points[:3], rtol=1e-5, atol=1e-5)}")
src/occupancy_encoder.py ADDED
@@ -0,0 +1,248 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Geometry-conditioned occupancy: query XYZ plus a shape latent.
2
+
3
+ ``OccupancyMLP`` stays xyz-only. ``shape_encoder: surface`` is the
4
+ envelope PointNet over ``(N, 6)`` XYZ + face normal (or ``(N, 3)``
5
+ for older checkpoints). ``knn_k > 0`` adds per-query nearest-neighbor
6
+ features (XYZ offset, plus the neighbor normal when the cloud is 6-D).
7
+ The face-token head (``shape_encoder: mesh``) was removed.
8
+ """
9
+
10
+ from __future__ import annotations
11
+
12
+ import torch
13
+ import torch.nn as nn
14
+ from torch import Tensor
15
+
16
+ from scatteringnet.occupancy_mlp import build_mlp
17
+
18
+ # Distinct from OccupancyMLP so infer can tell the checkpoint apart.
19
+ CHECKPOINT_KIND = "occupancy_encoder"
20
+
21
+
22
+ class SurfaceEncoder(nn.Module):
23
+ """
24
+ Per-point MLP + max-pool (PointNet) over an envelope cloud.
25
+
26
+ Shapes
27
+ ------
28
+ envelope: ``(U, N, C)`` unique meshes; ``C`` is 3 (XYZ) or 6 (XYZ+n)
29
+ output: ``(U, D)`` one latent per mesh
30
+ """
31
+
32
+ def __init__(
33
+ self, latent_dim: int = 64, hidden: int = 64, *, in_dim: int = 3
34
+ ) -> None:
35
+ super().__init__()
36
+ if latent_dim < 1:
37
+ raise ValueError(f"latent_dim must be >= 1, got {latent_dim}")
38
+ if hidden < 1:
39
+ raise ValueError(f"hidden must be >= 1, got {hidden}")
40
+ dim = int(in_dim)
41
+ if dim not in (3, 6):
42
+ raise ValueError(f"in_dim must be 3 or 6, got {dim}")
43
+ self.latent_dim = latent_dim
44
+ self.hidden = hidden
45
+ self.in_dim = dim
46
+ self.point_mlp = nn.Sequential(
47
+ nn.Linear(dim, hidden),
48
+ nn.ReLU(inplace=True),
49
+ nn.Linear(hidden, latent_dim),
50
+ )
51
+
52
+ def forward(self, envelope: Tensor) -> Tensor:
53
+ """Max-pool per-point features → one vector per unique mesh."""
54
+ if envelope.ndim != 3 or envelope.shape[-1] != self.in_dim:
55
+ raise ValueError(
56
+ f"envelope must have shape (U, N, {self.in_dim}), "
57
+ f"got {tuple(envelope.shape)}"
58
+ )
59
+ features = self.point_mlp(envelope)
60
+ return features.max(dim=1).values
61
+
62
+
63
+ def knn_offsets(xyz: Tensor, envelope: Tensor, k: int) -> Tensor:
64
+ """
65
+ Offsets from each query to its ``k`` nearest envelope points.
66
+
67
+ Distances use XYZ only (AABB Euclidean). ``k`` is clamped to N.
68
+ If the cloud is ``(B, N, 6)``, each neighbor is
69
+ ``(dx, dy, dz, nx, ny, nz)`` — relative position plus that
70
+ neighbor's stored normal. Query points have no normal.
71
+
72
+ Shapes: ``xyz (B, 3)``, ``envelope (B, N, 3|6)`` → ``(B, k, 3|6)``.
73
+ """
74
+ if xyz.ndim != 2 or xyz.shape[-1] != 3:
75
+ raise ValueError(f"xyz must have shape (B, 3), got {tuple(xyz.shape)}")
76
+ if envelope.ndim != 3 or envelope.shape[-1] not in (3, 6):
77
+ raise ValueError(
78
+ f"envelope must have shape (B, N, 3 or 6), got {tuple(envelope.shape)}"
79
+ )
80
+ if int(xyz.shape[0]) != int(envelope.shape[0]):
81
+ raise ValueError(
82
+ f"xyz/envelope batch mismatch: {tuple(xyz.shape)} vs {tuple(envelope.shape)}"
83
+ )
84
+ n_env = int(envelope.shape[1])
85
+ if n_env < 1:
86
+ raise ValueError("envelope length N must be >= 1")
87
+ take = min(int(k), n_env)
88
+ if take < 1:
89
+ raise ValueError(f"k must be >= 1, got {k}")
90
+ feat = int(envelope.shape[-1])
91
+ # k-NN is position-only; extras (normals) ride along after the gather.
92
+ pos = envelope[..., :3]
93
+ dist = torch.linalg.norm(pos - xyz.unsqueeze(1), dim=-1)
94
+ idx = dist.topk(take, dim=-1, largest=False).indices
95
+ nbrs = torch.gather(envelope, 1, idx.unsqueeze(-1).expand(-1, -1, feat))
96
+ rel_xyz = nbrs[..., :3] - xyz.unsqueeze(1)
97
+ if feat == 3:
98
+ return rel_xyz
99
+ return torch.cat([rel_xyz, nbrs[..., 3:]], dim=-1)
100
+
101
+
102
+ # Catalog trains used YAML seed 1. Old best.pt files omit ``seed``.
103
+ DEFAULT_ENVELOPE_SEED = 1
104
+
105
+
106
+ def envelope_seed_from_ckpt(ckpt: dict) -> int:
107
+ """
108
+ Envelope RNG seed this checkpoint was trained with.
109
+
110
+ New trains store ``seed`` on ``best.pt``. Older files omit it; do
111
+ not fall back to live YAML (that knob may have changed since train).
112
+ """
113
+ raw = ckpt.get("seed")
114
+ if raw is None:
115
+ return DEFAULT_ENVELOPE_SEED
116
+ return int(raw)
117
+
118
+
119
+ def envelope_dim_from_ckpt(ckpt: dict) -> int:
120
+ """
121
+ Envelope channel count this checkpoint was trained with.
122
+
123
+ New trains store ``envelope_dim``. Older XYZ-only ``best.pt`` files
124
+ omit it; the first SurfaceEncoder Linear in-features is then 3.
125
+ """
126
+ raw = ckpt.get("envelope_dim")
127
+ if raw is not None:
128
+ dim = int(raw)
129
+ if dim not in (3, 6):
130
+ raise ValueError(f"envelope_dim must be 3 or 6, got {dim}")
131
+ return dim
132
+ weight = (ckpt.get("state_dict") or {}).get("surface.point_mlp.0.weight")
133
+ if weight is not None:
134
+ dim = int(weight.shape[1])
135
+ if dim in (3, 6):
136
+ return dim
137
+ return 3
138
+
139
+
140
+ class OccupancyEncoder(nn.Module):
141
+ """
142
+ Occupancy logits from query XYZ and an envelope code.
143
+
144
+ Unique ``shape_id`` values are encoded **once** per batch, then
145
+ broadcast. ``knn_k > 0`` concatenates a local envelope code.
146
+
147
+ Shapes
148
+ ------
149
+ xyz: ``(B, 3)``
150
+ geom: ``(B, N, C)`` envelope; ``C`` is ``envelope_dim`` (3 or 6)
151
+ shape_id: ``(B,)`` long
152
+ output: ``(B, 1)`` logits
153
+ """
154
+
155
+ def __init__(
156
+ self,
157
+ hidden: int = 64,
158
+ depth: int = 4,
159
+ latent_dim: int = 64,
160
+ *,
161
+ shape_encoder: str = "surface",
162
+ knn_k: int = 0,
163
+ knn_local_dim: int | None = None,
164
+ envelope_dim: int = 6,
165
+ ) -> None:
166
+ super().__init__()
167
+ if hidden < 1:
168
+ raise ValueError(f"hidden must be >= 1, got {hidden}")
169
+ if depth < 1:
170
+ raise ValueError(f"depth must be >= 1, got {depth}")
171
+ if latent_dim < 1:
172
+ raise ValueError(f"latent_dim must be >= 1, got {latent_dim}")
173
+ kind = str(shape_encoder).strip().lower()
174
+ if kind != "surface":
175
+ raise ValueError(
176
+ "OccupancyEncoder only supports shape_encoder='surface' "
177
+ f"(face-token 'mesh' was removed), got {shape_encoder!r}"
178
+ )
179
+ k = int(knn_k)
180
+ if k < 0:
181
+ raise ValueError(f"knn_k must be >= 0, got {k}")
182
+ local_dim = int(knn_local_dim) if knn_local_dim is not None else int(latent_dim)
183
+ if k > 0 and local_dim < 1:
184
+ raise ValueError(f"knn_local_dim must be >= 1, got {local_dim}")
185
+ ed = int(envelope_dim)
186
+ if ed not in (3, 6):
187
+ raise ValueError(f"envelope_dim must be 3 or 6, got {ed}")
188
+ self.hidden = hidden
189
+ self.depth = depth
190
+ self.latent_dim = latent_dim
191
+ self.shape_encoder = kind
192
+ self.knn_k = k
193
+ self.knn_local_dim = local_dim if k > 0 else 0
194
+ self.envelope_dim = ed
195
+ # Name ``surface`` is load-stable for existing envelope checkpoints.
196
+ self.surface = SurfaceEncoder(latent_dim=latent_dim, hidden=hidden, in_dim=ed)
197
+ self.token_dim = ed
198
+ head_in = 3 + latent_dim
199
+ if k > 0:
200
+ # Same PointNet block as the global envelope, over k neighbor features.
201
+ self.local = SurfaceEncoder(latent_dim=local_dim, hidden=hidden, in_dim=ed)
202
+ head_in += local_dim
203
+ self.head = build_mlp(head_in, hidden, depth)
204
+
205
+ def encode_unique(self, geom: Tensor, shape_id: Tensor) -> Tensor:
206
+ """
207
+ Encode each distinct ``shape_id`` once and scatter back to ``(B, D)``.
208
+
209
+ Parameters
210
+ ----------
211
+ geom, shape_id:
212
+ Batched envelope clouds and integer mesh ids (same length B).
213
+ """
214
+ if shape_id.ndim != 1 or int(shape_id.shape[0]) != int(geom.shape[0]):
215
+ raise ValueError(
216
+ f"shape_id must be (B,), got {tuple(shape_id.shape)} "
217
+ f"for geom {tuple(geom.shape)}"
218
+ )
219
+ unique_ids, inverse = torch.unique(shape_id, sorted=True, return_inverse=True)
220
+ hits = shape_id.unsqueeze(0) == unique_ids.unsqueeze(1)
221
+ first = hits.to(dtype=torch.int64).argmax(dim=1)
222
+ z_unique = self.surface(geom[first])
223
+ return z_unique[inverse]
224
+
225
+ def forward(
226
+ self,
227
+ xyz: Tensor,
228
+ geom: Tensor,
229
+ shape_id: Tensor,
230
+ ) -> Tensor:
231
+ """``cat(xyz, z_global[, z_local])`` → occupancy logit."""
232
+ if xyz.ndim != 2 or xyz.shape[-1] != 3:
233
+ raise ValueError(f"xyz must have shape (B, 3), got {tuple(xyz.shape)}")
234
+ if geom.ndim != 3 or geom.shape[-1] != self.token_dim:
235
+ raise ValueError(
236
+ f"geom must have shape (B, K, {self.token_dim}), got {tuple(geom.shape)}"
237
+ )
238
+ if int(xyz.shape[0]) != int(geom.shape[0]):
239
+ raise ValueError(
240
+ f"xyz/geom batch mismatch: {tuple(xyz.shape)} vs {tuple(geom.shape)}"
241
+ )
242
+ ids = shape_id.reshape(-1)
243
+ z_shape = self.encode_unique(geom, ids)
244
+ pieces = [xyz, z_shape]
245
+ if self.knn_k > 0:
246
+ rel = knn_offsets(xyz, geom, self.knn_k)
247
+ pieces.append(self.local(rel))
248
+ return self.head(torch.cat(pieces, dim=-1))
src/occupancy_mlp.py ADDED
@@ -0,0 +1,87 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Occupancy MLP: raw XYZ → inside/outside logit."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import torch
6
+ import torch.nn as nn
7
+ from torch import Tensor
8
+
9
+ # Checkpoint schema tag (not a YAML knob).
10
+ CHECKPOINT_KIND = "occupancy_mlp"
11
+
12
+
13
+ def build_mlp(in_dim: int, hidden: int, depth: int) -> nn.Sequential:
14
+ """Linear→ReLU × ``depth`` then a 1-logit head. Shared by OccupancyMLP."""
15
+ if in_dim < 1:
16
+ raise ValueError(f"in_dim must be >= 1, got {in_dim}")
17
+ if hidden < 1:
18
+ raise ValueError(f"hidden must be >= 1, got {hidden}")
19
+ if depth < 1:
20
+ raise ValueError(f"depth must be >= 1, got {depth}")
21
+ layers: list[nn.Module] = []
22
+ dim = in_dim
23
+ for _ in range(depth):
24
+ layers.append(nn.Linear(dim, hidden))
25
+ layers.append(nn.ReLU(inplace=True))
26
+ dim = hidden
27
+ layers.append(nn.Linear(dim, 1))
28
+ return nn.Sequential(*layers)
29
+
30
+
31
+ class OccupancyMLP(nn.Module):
32
+ """
33
+ Tiny fully-connected occupancy field.
34
+
35
+ Maps a batch of 3D query coordinates to a single unnormalized logit per
36
+ point. A later training step will apply ``binary_cross_entropy_with_logits``
37
+ (do not softmax / sigmoid inside ``forward``).
38
+
39
+ Shapes
40
+ ------
41
+ xyz: ``(B, 3)`` batch of query points (device follows the caller)
42
+ output: ``(B, 1)`` logits; positive → inside, negative → outside
43
+
44
+ Device
45
+ ------
46
+ Parameters live on whatever device the module was moved to
47
+ (``.to(device)`` / ``.cuda()``). ``xyz`` must already be on that same
48
+ device; this module does not copy tensors.
49
+ """
50
+
51
+ def __init__(self, hidden: int = 64, depth: int = 4) -> None:
52
+ """
53
+ Build Linear→ReLU blocks then a 1-logit head.
54
+
55
+ Parameters
56
+ ----------
57
+ hidden:
58
+ Channel width of each hidden Linear (must be ``>= 1``).
59
+ depth:
60
+ Number of hidden Linear+ReLU blocks (must be ``>= 1``).
61
+ """
62
+ super().__init__()
63
+ self.hidden = hidden
64
+ self.depth = depth
65
+ # First Linear is 3 → H; remaining blocks are H → H.
66
+ self.net = build_mlp(3, hidden, depth)
67
+
68
+ def forward(self, xyz: Tensor) -> Tensor:
69
+ """
70
+ Evaluate occupancy logits at query coordinates.
71
+
72
+ Parameters
73
+ ----------
74
+ xyz:
75
+ Float tensor of shape ``(B, 3)``. Last dim is Cartesian XYZ.
76
+
77
+ Returns
78
+ -------
79
+ Tensor
80
+ Float tensor of shape ``(B, 1)`` on the same device as ``xyz``.
81
+ """
82
+ if xyz.ndim != 2 or xyz.shape[-1] != 3:
83
+ raise ValueError(
84
+ f"xyz must have shape (B, 3), got {tuple(xyz.shape)}"
85
+ )
86
+ # Sequential Linear layers require matching dtype/device with parameters.
87
+ return self.net(xyz)
src/viewer/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ """Three.js inspect helper (HTTP + Job A / Job B)."""
src/viewer/infer_job.py ADDED
@@ -0,0 +1,536 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Job A / B: classify uploaded points with a ``models/*/best.pt``.
2
+
3
+ Reuses ``load_occupancy_model`` and AABB helpers. Envelope tokens follow
4
+ the **checkpoint**, not live ``config.yaml``. Face-token checkpoints
5
+ are rejected. Does not modify occupancy train/infer modules.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import base64
11
+ import threading
12
+ import time
13
+ from pathlib import Path
14
+ from typing import Any, TYPE_CHECKING
15
+
16
+ if TYPE_CHECKING:
17
+ from scatteringnet.config import OccupancyConfig
18
+
19
+ import numpy as np
20
+
21
+
22
+ from scatteringnet.infer_multi_npz import load_occupancy_model # noqa: E402
23
+ from scatteringnet.geometry.mesh_io import load_obj_triangles # noqa: E402
24
+ from scatteringnet.geometry.surface import ( # noqa: E402
25
+ apply_envelope_aabb,
26
+ project_envelope_dim,
27
+ sample_surface_points,
28
+ )
29
+ from scatteringnet.metrics import occupancy_metrics # noqa: E402
30
+ from scatteringnet.normalize import apply_normalization, compute_center_scale # noqa: E402
31
+ from scatteringnet.occupancy_encoder import CHECKPOINT_KIND as ENCODER_KIND # noqa: E402
32
+ from scatteringnet.occupancy_encoder import envelope_dim_from_ckpt # noqa: E402
33
+ from scatteringnet.occupancy_encoder import envelope_seed_from_ckpt # noqa: E402
34
+
35
+ from scatteringnet.viewer.mesh_access import resolve_viewer_mesh # noqa: E402
36
+ from scatteringnet.viewer.model_access import match_checkpoint_part, resolve_viewer_checkpoint # noqa: E402
37
+ from scatteringnet.viewer.obj_fill import triangles_from_obj_text # noqa: E402
38
+
39
+ MAX_POINTS = 2_000_000
40
+
41
+ # One occupancy head in this process. Same checkpoint + device → skip torch.load.
42
+ _CACHE_LOCK = threading.Lock()
43
+ _MODEL_CACHE: tuple[tuple[str, int, str], Any, dict[str, Any]] | None = None
44
+
45
+
46
+ def clear_model_cache() -> None:
47
+ """Drop the cached head (tests / swapped GPU)."""
48
+ global _MODEL_CACHE
49
+ with _CACHE_LOCK:
50
+ _MODEL_CACHE = None
51
+
52
+
53
+ def _model_cache_key(ckpt_path: Path, device: Any) -> tuple[str, int, str]:
54
+ stat = ckpt_path.stat()
55
+ return (str(ckpt_path.resolve()), int(stat.st_mtime_ns), str(device))
56
+
57
+
58
+ def load_cached_occupancy(
59
+ run_id: str,
60
+ models_root: Path,
61
+ device: Any,
62
+ ) -> tuple[Any, dict[str, Any], Path, bool]:
63
+ """
64
+ Load ``best.pt`` once per (path, mtime, device).
65
+
66
+ Returns ``(model, ckpt, path, cache_hit)``. Switching run_id replaces
67
+ the slot so GPU RAM does not keep every inspect checkpoint.
68
+ """
69
+ import torch
70
+
71
+ ckpt_path = resolve_viewer_checkpoint(run_id, models_root)
72
+ key = _model_cache_key(ckpt_path, device)
73
+ global _MODEL_CACHE
74
+ with _CACHE_LOCK:
75
+ slot = _MODEL_CACHE
76
+ if slot is not None and slot[0] == key:
77
+ return slot[1], slot[2], ckpt_path, True
78
+ ckpt = torch.load(ckpt_path, map_location="cpu", weights_only=False)
79
+ model = load_occupancy_model(ckpt, device)
80
+ with _CACHE_LOCK:
81
+ _MODEL_CACHE = (key, model, ckpt)
82
+ return model, ckpt, ckpt_path, False
83
+
84
+
85
+ def _sync_if_cuda(device) -> None:
86
+ """Wait for GPU work so lap times are not just the CPU launch."""
87
+ import torch
88
+
89
+ dev = device if hasattr(device, "type") else torch.device(str(device))
90
+ if str(dev.type) == "cuda":
91
+ torch.cuda.synchronize()
92
+
93
+
94
+ def _round_s(t0: float) -> float:
95
+ return round(time.perf_counter() - t0, 3)
96
+
97
+
98
+ def pred_from_probs(probs: np.ndarray, threshold: float = 0.5) -> np.ndarray:
99
+ """Hard inside labels: 1 iff sigmoid probability is at least ``threshold``."""
100
+ t = float(threshold)
101
+ if not np.isfinite(t):
102
+ t = 0.5
103
+ t = min(1.0, max(0.0, t))
104
+ return (np.asarray(probs, dtype=np.float32).reshape(-1) >= t).astype(np.uint8)
105
+
106
+
107
+ def _sigmoid_probs_and_pred(logits: Any) -> tuple[np.ndarray, np.ndarray]:
108
+ """
109
+ Float32 sigmoid of a 1-D logit tensor, plus a 0.5-cut pred for tests/compat.
110
+
111
+ The inspect page re-cuts from ``prob_b64``; occupancy train metrics stay at 0.5.
112
+ """
113
+ import torch
114
+
115
+ probs = (
116
+ torch.sigmoid(logits.reshape(-1))
117
+ .detach()
118
+ .cpu()
119
+ .numpy()
120
+ .astype(np.float32, copy=False)
121
+ )
122
+ return probs, pred_from_probs(probs, 0.5)
123
+
124
+
125
+ def aabb_for_viewer(
126
+ ckpt: dict[str, Any],
127
+ *,
128
+ npz_name: str,
129
+ mesh_path: str,
130
+ points: np.ndarray,
131
+ data_dir: Path | None,
132
+ vertices: np.ndarray | None = None,
133
+ uploaded_obj: bool = False,
134
+ ) -> tuple[np.ndarray, float, str]:
135
+ """
136
+ Checkpoint AABB when this catalog NPZ/mesh was trained; else mesh vertices.
137
+
138
+ Uploaded OBJ Fill (Job B) always uses this mesh's vertices. Catalog
139
+ ``parts`` matched by filename would apply another cube's box.
140
+ """
141
+ if not uploaded_obj:
142
+ part = match_checkpoint_part(ckpt, npz_name, mesh_path)
143
+ if part is not None:
144
+ center = np.asarray(part["center"], dtype=np.float32).reshape(3)
145
+ scale = float(part["scale"])
146
+ return center, scale, "checkpoint"
147
+ if vertices is not None and int(np.asarray(vertices).shape[0]) > 0:
148
+ center, scale = compute_center_scale(np.asarray(vertices, dtype=np.float32))
149
+ return center, scale, "mesh"
150
+ if mesh_path and data_dir is not None:
151
+ try:
152
+ resolved = resolve_viewer_mesh(mesh_path, data_dir)
153
+ vertices, _faces = load_obj_triangles(resolved)
154
+ center, scale = compute_center_scale(vertices)
155
+ return center, scale, "mesh"
156
+ except (PermissionError, FileNotFoundError, ValueError):
157
+ pass
158
+ center, scale = compute_center_scale(points)
159
+ return center, scale, "points"
160
+
161
+
162
+ def _ckpt_shape_encoder(ckpt: dict[str, Any]) -> str:
163
+ """
164
+ Encoder stored on ``best.pt``, not live YAML.
165
+
166
+ Missing ``shape_encoder`` on an occupancy-encoder checkpoint is the
167
+ original envelope head.
168
+ """
169
+ raw = str(ckpt.get("shape_encoder") or "").strip().lower()
170
+ if raw in ("surface", "none"):
171
+ return raw
172
+ if raw == "mesh":
173
+ raise ValueError(
174
+ "face-token occupancy checkpoints (shape_encoder='mesh') "
175
+ "are no longer supported"
176
+ )
177
+ kind = str(ckpt.get("kind") or "")
178
+ if kind == ENCODER_KIND:
179
+ return "surface"
180
+ return "none"
181
+
182
+
183
+ def _runtime_cfg(cfg: OccupancyConfig | None):
184
+ """Use the caller cfg (tests) or load repo YAML for device / batch / seed."""
185
+ if cfg is not None:
186
+ return cfg
187
+ from scatteringnet.config import load_config
188
+
189
+ return load_config()
190
+
191
+
192
+ def _forward_logits(model, cfg, xyz, geom, shape_id):
193
+ """Batched occupancy logits. ``geom`` is envelope ``(1, N, 3)`` or ``None``."""
194
+ import torch
195
+
196
+ n = int(xyz.shape[0])
197
+ logits_rows: list[Any] = []
198
+ with torch.no_grad():
199
+ for start in range(0, n, int(cfg.batch_size)):
200
+ sl = slice(start, start + int(cfg.batch_size))
201
+ batch_xyz = xyz[sl].to(cfg.device)
202
+ if geom is None:
203
+ logits_rows.append(model(batch_xyz).cpu())
204
+ else:
205
+ b = int(batch_xyz.shape[0])
206
+ geom_b = geom.expand(b, -1, -1).to(cfg.device)
207
+ sid = shape_id.reshape(()).expand(b).to(cfg.device)
208
+ logits_rows.append(model(batch_xyz, geom_b, sid).cpu())
209
+ return torch.cat(logits_rows, dim=0)
210
+
211
+
212
+ def _geom_from_mesh(ckpt, cfg, vertices, faces, center, scale, cache_key: str):
213
+ """
214
+ Build the envelope this checkpoint was trained with.
215
+
216
+ ``surface`` → ``(1, n_surface, C)`` envelope (C from checkpoint).
217
+ ``none`` → no geometry (xyz-only MLP).
218
+ Count comes from the checkpoint. Sampling is always area-weighted
219
+ (old ``envelope_mix`` on ``best.pt`` is ignored).
220
+ """
221
+ import torch
222
+
223
+ enc = _ckpt_shape_encoder(ckpt)
224
+ seed = envelope_seed_from_ckpt(ckpt)
225
+ if enc == "surface":
226
+ n_surface = (
227
+ int(ckpt["n_surface"]) if ckpt.get("n_surface") is not None else 1024
228
+ )
229
+ world = sample_surface_points(
230
+ vertices,
231
+ faces,
232
+ n_surface,
233
+ seed=seed,
234
+ cache_key=cache_key,
235
+ )
236
+ env = apply_envelope_aabb(world, center, scale)
237
+ env = project_envelope_dim(env, envelope_dim_from_ckpt(ckpt))
238
+ geom = torch.from_numpy(env).unsqueeze(0)
239
+ shape_id = torch.zeros(1, dtype=torch.long)
240
+ return geom, shape_id
241
+ return None, None
242
+
243
+
244
+ def decode_points_b64(text: str) -> np.ndarray:
245
+ """Little-endian float32 XYZ from standard base64."""
246
+ raw = base64.b64decode(text)
247
+ pts = np.frombuffer(raw, dtype=np.float32)
248
+ if pts.size % 3 != 0:
249
+ raise ValueError("points buffer length is not a multiple of 3")
250
+ return np.ascontiguousarray(pts.reshape(-1, 3))
251
+
252
+
253
+ def decode_labels_b64(text: str, n: int) -> np.ndarray:
254
+ """uint8 {0,1} labels, length ``n``."""
255
+ raw = base64.b64decode(text)
256
+ labels = np.frombuffer(raw, dtype=np.uint8)
257
+ if int(labels.shape[0]) != int(n):
258
+ raise ValueError(
259
+ f"labels length {labels.shape[0]} does not match points {n}"
260
+ )
261
+ return np.ascontiguousarray(labels)
262
+
263
+
264
+ def infer_uploaded_npz(
265
+ *,
266
+ run_id: str,
267
+ models_root: Path,
268
+ data_dir: Path | None,
269
+ npz_name: str,
270
+ mesh_path: str,
271
+ points: np.ndarray,
272
+ labels: np.ndarray,
273
+ cfg: OccupancyConfig | None = None,
274
+ ) -> dict[str, Any]:
275
+ """
276
+ Classify ``points`` with ``best.pt``. Returns pred bytes (base64) and scores.
277
+
278
+ Envelope rebuild follows ``ckpt['shape_encoder']``.
279
+ """
280
+ import torch
281
+
282
+ n = int(points.shape[0])
283
+ if n < 1:
284
+ raise ValueError("no query points")
285
+ if n > MAX_POINTS:
286
+ raise ValueError(f"too many points ({n}; max {MAX_POINTS})")
287
+ if points.ndim != 2 or points.shape[1] != 3:
288
+ raise ValueError(f"points must be (N, 3), got {tuple(points.shape)}")
289
+ if int(labels.shape[0]) != n:
290
+ raise ValueError("labels length does not match points")
291
+
292
+ t_all = time.perf_counter()
293
+ t0 = time.perf_counter()
294
+ cfg = _runtime_cfg(cfg)
295
+ model, ckpt, ckpt_path, cache_hit = load_cached_occupancy(
296
+ run_id, models_root, cfg.device
297
+ )
298
+ _sync_if_cuda(cfg.device)
299
+ load_s = _round_s(t0)
300
+ center, scale, aabb_src = aabb_for_viewer(
301
+ ckpt,
302
+ npz_name=npz_name,
303
+ mesh_path=mesh_path,
304
+ points=points,
305
+ data_dir=data_dir,
306
+ )
307
+ enc = _ckpt_shape_encoder(ckpt)
308
+ geom = None
309
+ shape_id = None
310
+ parse_s = 0.0
311
+ envelope_s = 0.0
312
+ if enc in ("surface", "mesh"):
313
+ if not mesh_path:
314
+ raise ValueError("this checkpoint needs a mesh_path for the geometry encoder")
315
+ if data_dir is None:
316
+ raise ValueError("helper has no data_dir; cannot load the OBJ for the encoder")
317
+ t0 = time.perf_counter()
318
+ mesh_file = resolve_viewer_mesh(mesh_path, data_dir)
319
+ vertices, faces = load_obj_triangles(mesh_file)
320
+ parse_s = _round_s(t0)
321
+ t0 = time.perf_counter()
322
+ geom, shape_id = _geom_from_mesh(
323
+ ckpt, cfg, vertices, faces, center, scale, str(mesh_file.resolve())
324
+ )
325
+ envelope_s = _round_s(t0)
326
+ if geom is None:
327
+ raise ValueError(f"checkpoint shape_encoder={enc!r} built no geometry tokens")
328
+
329
+ t0 = time.perf_counter()
330
+ xyz = torch.from_numpy(apply_normalization(points, center, scale))
331
+ y = torch.from_numpy(labels.astype(np.float32)).unsqueeze(1)
332
+ logits = _forward_logits(model, cfg, xyz, geom, shape_id)
333
+ _sync_if_cuda(cfg.device)
334
+ forward_s = _round_s(t0)
335
+ scores = occupancy_metrics(logits, y)
336
+ probs, pred = _sigmoid_probs_and_pred(logits)
337
+ timings = {
338
+ "parse_obj": parse_s,
339
+ "load_model": load_s,
340
+ "envelope": envelope_s,
341
+ "forward": forward_s,
342
+ "server_total": _round_s(t_all),
343
+ "load_cached": cache_hit,
344
+ "device": str(cfg.device),
345
+ }
346
+ print(
347
+ "viewer infer-npz timings (s) n=%s parse=%.3f load=%.3f envelope=%.3f "
348
+ "forward=%.3f total=%.3f cached=%s device=%s"
349
+ % (
350
+ n,
351
+ parse_s,
352
+ load_s,
353
+ envelope_s,
354
+ forward_s,
355
+ timings["server_total"],
356
+ cache_hit,
357
+ cfg.device,
358
+ ),
359
+ flush=True,
360
+ )
361
+ gt = labels > 0
362
+ pred_bool = pred > 0
363
+ n_fn = int(np.count_nonzero(gt & ~pred_bool))
364
+ n_fp = int(np.count_nonzero(~gt & pred_bool))
365
+ return {
366
+ "pred_b64": base64.b64encode(np.ascontiguousarray(pred)).decode("ascii"),
367
+ "prob_b64": base64.b64encode(np.ascontiguousarray(probs)).decode("ascii"),
368
+ "n": n,
369
+ "accuracy": float(scores.accuracy),
370
+ "inside_iou": float(scores.inside_iou),
371
+ "inside_f1": float(scores.inside_f1),
372
+ "kind": str(ckpt.get("kind") or ""),
373
+ "shape_encoder": enc,
374
+ "aabb": aabb_src,
375
+ "n_error": n_fn + n_fp,
376
+ "n_fn": n_fn,
377
+ "n_fp": n_fp,
378
+ "run_id": Path(ckpt_path).parent.name,
379
+ "timings": timings,
380
+ }
381
+
382
+
383
+ def infer_uploaded_obj(
384
+ *,
385
+ run_id: str,
386
+ models_root: Path,
387
+ data_dir: Path | None,
388
+ obj_name: str,
389
+ obj_text: str,
390
+ points: np.ndarray,
391
+ cfg: OccupancyConfig | None = None,
392
+ ) -> dict[str, Any]:
393
+ """
394
+ Classify fill-lattice XYZ with ``best.pt``.
395
+
396
+ Envelope tokens are built from the uploaded OBJ when the
397
+ checkpoint is ``surface``. No file labels (Job B): metrics are omitted.
398
+ """
399
+ import torch
400
+
401
+ n = int(points.shape[0])
402
+ if n < 1:
403
+ raise ValueError("no query points; Fill points first")
404
+ if n > MAX_POINTS:
405
+ raise ValueError(f"too many points ({n}; max {MAX_POINTS})")
406
+ if points.ndim != 2 or points.shape[1] != 3:
407
+ raise ValueError(f"points must be (N, 3), got {tuple(points.shape)}")
408
+
409
+ t_all = time.perf_counter()
410
+ t0 = time.perf_counter()
411
+ vertices, faces = triangles_from_obj_text(obj_text)
412
+ parse_s = _round_s(t0)
413
+ t0 = time.perf_counter()
414
+ cfg = _runtime_cfg(cfg)
415
+ model, ckpt, ckpt_path, cache_hit = load_cached_occupancy(
416
+ run_id, models_root, cfg.device
417
+ )
418
+ _sync_if_cuda(cfg.device)
419
+ load_s = _round_s(t0)
420
+ center, scale, aabb_src = aabb_for_viewer(
421
+ ckpt,
422
+ npz_name="",
423
+ mesh_path=str(obj_name or ""),
424
+ points=points,
425
+ data_dir=data_dir,
426
+ vertices=vertices,
427
+ uploaded_obj=True,
428
+ )
429
+ enc = _ckpt_shape_encoder(ckpt)
430
+ t0 = time.perf_counter()
431
+ geom, shape_id = _geom_from_mesh(
432
+ ckpt, cfg, vertices, faces, center, scale, "upload:" + str(obj_name or "obj")
433
+ )
434
+ envelope_s = _round_s(t0)
435
+ if enc in ("surface", "mesh") and geom is None:
436
+ raise ValueError(f"checkpoint shape_encoder={enc!r} built no geometry tokens")
437
+ t0 = time.perf_counter()
438
+ xyz = torch.from_numpy(apply_normalization(points, center, scale))
439
+ logits = _forward_logits(model, cfg, xyz, geom, shape_id)
440
+ _sync_if_cuda(cfg.device)
441
+ forward_s = _round_s(t0)
442
+ probs, pred = _sigmoid_probs_and_pred(logits)
443
+ n_in = int(np.count_nonzero(pred > 0))
444
+ timings = {
445
+ "parse_obj": parse_s,
446
+ "load_model": load_s,
447
+ "envelope": envelope_s,
448
+ "forward": forward_s,
449
+ "server_total": _round_s(t_all),
450
+ "load_cached": cache_hit,
451
+ "device": str(cfg.device),
452
+ }
453
+ print(
454
+ "viewer infer-obj timings (s) n=%s parse=%.3f load=%.3f envelope=%.3f "
455
+ "forward=%.3f total=%.3f cached=%s device=%s"
456
+ % (
457
+ n,
458
+ parse_s,
459
+ load_s,
460
+ envelope_s,
461
+ forward_s,
462
+ timings["server_total"],
463
+ cache_hit,
464
+ cfg.device,
465
+ ),
466
+ flush=True,
467
+ )
468
+ return {
469
+ "pred_b64": base64.b64encode(np.ascontiguousarray(pred)).decode("ascii"),
470
+ "prob_b64": base64.b64encode(np.ascontiguousarray(probs)).decode("ascii"),
471
+ "n": n,
472
+ "n_inside": n_in,
473
+ "n_outside": n - n_in,
474
+ "kind": str(ckpt.get("kind") or ""),
475
+ "shape_encoder": enc,
476
+ "aabb": aabb_src,
477
+ "run_id": Path(ckpt_path).parent.name,
478
+ "timings": timings,
479
+ }
480
+
481
+
482
+ _WARM_OBJ = (
483
+ "v 0 0 0\nv 1 0 0\nv 0 1 0\nv 0 0 1\n"
484
+ "f 1 2 3\nf 1 2 4\nf 1 3 4\nf 2 3 4\n"
485
+ )
486
+
487
+
488
+ def warmup_viewer_helper(models_root: Path) -> None:
489
+ """
490
+ Pay first-click costs before the browser opens: torch, CUDA, trimesh
491
+ fill, load the INSPECT ``best.pt`` (else newest), tiny envelope + forward.
492
+ """
493
+ t0 = time.perf_counter()
494
+ print("warmup: torch / CUDA…", flush=True)
495
+ import torch
496
+
497
+ from scatteringnet.config import load_config
498
+ from scatteringnet.viewer.obj_fill import fill_aabb_lattice
499
+
500
+ cfg = load_config()
501
+ print("warmup: occupancy device=" + str(cfg.device), flush=True)
502
+ if str(cfg.device.type) == "cuda":
503
+ torch.zeros(1, device=cfg.device)
504
+ torch.cuda.synchronize()
505
+ print("warmup: CUDA " + str(cfg.device), flush=True)
506
+ else:
507
+ print("warmup: CPU only", flush=True)
508
+
509
+ print("warmup: tiny fill…", flush=True)
510
+ verts, faces = triangles_from_obj_text(_WARM_OBJ)
511
+ points, _used, _grid = fill_aabb_lattice(verts, 0.40)
512
+ if int(points.shape[0]) < 1:
513
+ points = np.array([[0.2, 0.2, 0.2], [0.8, 0.2, 0.2]], dtype=np.float32)
514
+
515
+ from scatteringnet.viewer.model_access import inspect_run_id, list_viewer_models
516
+
517
+ rows = list_viewer_models(models_root)
518
+ if not rows:
519
+ print("warmup: no models/; skip infer", flush=True)
520
+ print("warmup: done in %.1fs" % (time.perf_counter() - t0), flush=True)
521
+ return
522
+ ids = [str(row["id"]) for row in rows]
523
+ wanted = inspect_run_id()
524
+ # Warm the inspect alias when that folder exists; otherwise newest.
525
+ run_id = wanted if wanted in ids else ids[0]
526
+ print("warmup: infer " + run_id + "…", flush=True)
527
+ infer_uploaded_obj(
528
+ run_id=run_id,
529
+ models_root=models_root,
530
+ data_dir=None,
531
+ obj_name="warmup.obj",
532
+ obj_text=_WARM_OBJ,
533
+ points=points,
534
+ cfg=cfg,
535
+ )
536
+ print("warmup: done in %.1fs" % (time.perf_counter() - t0), flush=True)
src/viewer/mesh_access.py ADDED
@@ -0,0 +1,52 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Resolve NPZ ``mesh_path`` for the occupancy viewer helper.
2
+
3
+ Uses :func:`data_npz.resolve_mesh_path` and then requires the file to sit
4
+ under ``data_dir`` (no arbitrary filesystem reads).
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ from pathlib import Path
10
+
11
+ from scatteringnet.data_npz import resolve_mesh_path
12
+
13
+
14
+ def resolve_viewer_mesh(stored: str, data_dir: Path | str) -> Path:
15
+ """
16
+ Return an existing ``.obj`` under ``data_dir`` for a stored mesh_path.
17
+
18
+ Parameters
19
+ ----------
20
+ stored:
21
+ Relative or absolute path from the NPZ ``mesh_path`` array.
22
+ data_dir:
23
+ Dataset root (``config.yaml`` ``data_dir``).
24
+
25
+ Returns
26
+ -------
27
+ Path
28
+ Resolved OBJ path.
29
+
30
+ Raises
31
+ ------
32
+ ValueError
33
+ Empty path or not an OBJ.
34
+ FileNotFoundError
35
+ Mesh does not exist (from :func:`resolve_mesh_path`).
36
+ PermissionError
37
+ Resolved path is outside ``data_dir``.
38
+ """
39
+ text = str(stored).strip()
40
+ if not text:
41
+ raise ValueError("mesh_path is empty")
42
+ root = Path(data_dir).resolve()
43
+ resolved = resolve_mesh_path(text, root)
44
+ try:
45
+ resolved.relative_to(root)
46
+ except ValueError as exc:
47
+ raise PermissionError(
48
+ f"mesh is outside data_dir: {resolved} (stored={text!r})"
49
+ ) from exc
50
+ if resolved.suffix.lower() != ".obj":
51
+ raise ValueError(f"mesh is not an OBJ: {resolved}")
52
+ return resolved
src/viewer/model_access.py ADDED
@@ -0,0 +1,183 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """List and resolve ``models/<run_id>/best.pt`` for the occupancy viewer.
2
+
3
+ Read-only: paths must stay under the repo ``models/`` folder.
4
+ Does not import torch.
5
+ """
6
+
7
+ from __future__ import annotations
8
+
9
+ from pathlib import Path
10
+
11
+ import yaml
12
+
13
+ # Repo root: this file lives at src/viewer/model_access.py.
14
+ _REPO_ROOT = Path(__file__).resolve().parents[2]
15
+ # Pointers (YAML). Occupancy train / infer math does not read these.
16
+ INSPECT_POINTER = _REPO_ROOT / "docs" / "inspect_checkpoint.yaml"
17
+ HOLDOUT_POINTER = _REPO_ROOT / "docs" / "locked_holdout_objs.yaml"
18
+ # Same id as the committed pointer; used if that file is missing.
19
+ _INSPECT_FALLBACK = "2026-09-14_07-43-34_prim_extruded_nr45_knn24_n2048_n6"
20
+
21
+
22
+ def inspect_run_id() -> str:
23
+ """INSPECT alias: ``run_id`` in ``docs/inspect_checkpoint.yaml``.
24
+
25
+ Viewers use this as the default ``models/<id>/best.pt`` when that file
26
+ exists. Occupancy train / infer math is unchanged.
27
+ """
28
+ try:
29
+ raw = yaml.safe_load(INSPECT_POINTER.read_text(encoding="utf-8"))
30
+ except OSError:
31
+ return _INSPECT_FALLBACK
32
+ if not isinstance(raw, dict):
33
+ return _INSPECT_FALLBACK
34
+ token = str(raw.get("run_id") or "").strip()
35
+ return token[:200] if token else _INSPECT_FALLBACK
36
+
37
+
38
+ def locked_holdout_objs() -> list[str]:
39
+ """OBJ basenames from ``docs/locked_holdout_objs.yaml``.
40
+
41
+ Train / catalog construction does not consult this list. It is the
42
+ locked inspect set for humans and for tests.
43
+ """
44
+ try:
45
+ raw = yaml.safe_load(HOLDOUT_POINTER.read_text(encoding="utf-8"))
46
+ except OSError:
47
+ return []
48
+ if not isinstance(raw, dict):
49
+ return []
50
+ objs = raw.get("objs") or []
51
+ if not isinstance(objs, list):
52
+ return []
53
+ names: list[str] = []
54
+ for item in objs:
55
+ name = str(item or "").strip()
56
+ if name:
57
+ names.append(name)
58
+ return names
59
+
60
+
61
+ def _shape_encoder_from_run(runs_root: Path | None, run_id: str) -> str:
62
+ """Read ``shape_encoder`` from ``runs/<id>/config.yaml`` (no torch)."""
63
+ if runs_root is None:
64
+ return ""
65
+ path = Path(runs_root) / run_id / "config.yaml"
66
+ try:
67
+ if not path.is_file():
68
+ return ""
69
+ for line in path.read_text(encoding="utf-8").splitlines():
70
+ stripped = line.strip()
71
+ if stripped.startswith("#") or not stripped.startswith("shape_encoder:"):
72
+ continue
73
+ raw = stripped.split(":", 1)[1].split("#", 1)[0].strip().strip("\"'")
74
+ kind = raw.lower()
75
+ if kind in ("surface", "mesh", "none"):
76
+ return kind
77
+ return ""
78
+ except OSError:
79
+ return ""
80
+ return ""
81
+
82
+
83
+ def list_viewer_models(
84
+ models_root: Path | str,
85
+ *,
86
+ runs_root: Path | str | None = None,
87
+ ) -> list[dict[str, str | int]]:
88
+ """
89
+ Return ``best.pt`` checkpoints one level under ``models_root``.
90
+
91
+ Each item: ``id`` (folder name), ``path`` (repo-relative), ``mtime``,
92
+ and ``shape_encoder`` when ``runs/<id>/config.yaml`` is present.
93
+ Newest first. Missing folder → empty list.
94
+ """
95
+ root = Path(models_root)
96
+ try:
97
+ root = root.resolve()
98
+ except OSError:
99
+ return []
100
+ if not root.is_dir():
101
+ return []
102
+ runs = Path(runs_root) if runs_root is not None else None
103
+ items: list[dict[str, str | int]] = []
104
+ for best in root.glob("*/best.pt"):
105
+ if not best.is_file():
106
+ continue
107
+ run_id = best.parent.name
108
+ try:
109
+ mtime = int(best.stat().st_mtime)
110
+ except OSError:
111
+ mtime = 0
112
+ enc = _shape_encoder_from_run(runs, run_id)
113
+ row: dict[str, str | int] = {
114
+ "id": run_id,
115
+ "path": "models/" + run_id + "/best.pt",
116
+ "mtime": mtime,
117
+ }
118
+ if enc:
119
+ row["shape_encoder"] = enc
120
+ items.append(row)
121
+ items.sort(key=lambda row: (-int(row["mtime"]), str(row["id"])))
122
+ return items
123
+
124
+
125
+ def resolve_viewer_checkpoint(run_id: str, models_root: Path | str) -> Path:
126
+ """
127
+ Return ``models_root / <run_id> / best.pt``.
128
+
129
+ ``run_id`` is a single folder name (no slashes). Also accepts
130
+ ``models/<run_id>/best.pt`` and strips it down to the folder name.
131
+
132
+ Raises
133
+ ------
134
+ ValueError
135
+ Empty or unsafe id.
136
+ FileNotFoundError
137
+ ``best.pt`` is missing.
138
+ PermissionError
139
+ Resolved path is outside ``models_root``.
140
+ """
141
+ raw = str(run_id or "").strip().replace("\\", "/")
142
+ if raw.endswith("/best.pt"):
143
+ raw = raw[: -len("/best.pt")]
144
+ if raw.startswith("models/"):
145
+ raw = raw[len("models/") :]
146
+ name = raw.strip("/")
147
+ if not name or "/" in name or name in (".", "..") or ".." in name:
148
+ raise ValueError("invalid model id")
149
+ root = Path(models_root).resolve()
150
+ best = (root / name / "best.pt").resolve()
151
+ try:
152
+ best.relative_to(root)
153
+ except ValueError as exc:
154
+ raise PermissionError(
155
+ f"checkpoint is outside models/: {best} (id={run_id!r})"
156
+ ) from exc
157
+ if not best.is_file():
158
+ raise FileNotFoundError(f"best.pt not found for {name!r}")
159
+ return best
160
+
161
+
162
+ def match_checkpoint_part(ckpt: dict, npz_name: str, mesh_path: str) -> dict | None:
163
+ """
164
+ Catalog AABB from ``ckpt['parts']``.
165
+
166
+ NPZ: full stored path or the occupancy filename (those stems are unique).
167
+ Mesh: exact stored path only. Basename matches (``cube.obj``) steal a
168
+ catalog box for an OOD upload of the same name — Job B must not do that.
169
+ """
170
+ parts = ckpt.get("parts") or []
171
+ name = Path(str(npz_name or "")).name
172
+ mesh = str(mesh_path or "").replace("\\", "/").strip()
173
+ if name:
174
+ for part in parts:
175
+ stored = str(part.get("npz") or "").replace("\\", "/")
176
+ if stored == name or Path(stored).name == name:
177
+ return part
178
+ if mesh:
179
+ for part in parts:
180
+ stored_m = str(part.get("mesh") or "").replace("\\", "/").strip()
181
+ if stored_m and stored_m == mesh:
182
+ return part
183
+ return None