Spaces:
Running on Zero
Running on Zero
scatteringnet-space commited on
Commit ·
bc4c433
0
Parent(s):
Slim Gradio Space: infer + demo only
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitignore +1 -0
- README.md +27 -0
- config.yaml +73 -0
- docs/inspect_checkpoint.yaml +5 -0
- models/.gitkeep +0 -0
- pyproject.toml +22 -0
- requirements.txt +4 -0
- scatteringnet/__init__.py +1 -0
- scatteringnet/config.py +611 -0
- scatteringnet/data_npz.py +460 -0
- scatteringnet/geometry/__init__.py +10 -0
- scatteringnet/geometry/mesh_io.py +80 -0
- scatteringnet/geometry/surface.py +154 -0
- scatteringnet/geometry/trimesh_util.py +19 -0
- scatteringnet/infer_multi_npz.py +202 -0
- scatteringnet/metrics.py +156 -0
- scatteringnet/normalize.py +104 -0
- scatteringnet/occupancy_encoder.py +248 -0
- scatteringnet/occupancy_mlp.py +87 -0
- scatteringnet/viewer/__init__.py +1 -0
- scatteringnet/viewer/infer_job.py +536 -0
- scatteringnet/viewer/mesh_access.py +52 -0
- scatteringnet/viewer/model_access.py +183 -0
- scatteringnet/viewer/obj_fill.py +138 -0
- src/__init__.py +1 -0
- src/config.py +611 -0
- src/data_npz.py +460 -0
- src/geometry/__init__.py +10 -0
- src/geometry/mesh_io.py +80 -0
- src/geometry/surface.py +154 -0
- src/geometry/trimesh_util.py +19 -0
- src/gradio/app.py +503 -0
- src/gradio/examples/Helix_bend.obj +0 -0
- src/gradio/examples/Obese.obj +0 -0
- src/gradio/examples/Player.obj +0 -0
- src/gradio/examples/TorusX3_box.obj +0 -0
- src/gradio/examples/dog.obj +0 -0
- src/gradio/examples/horse.obj +0 -0
- src/gradio/figure.py +255 -0
- src/gradio/orbit.js +564 -0
- src/gradio/pipeline.py +133 -0
- src/infer_multi_npz.py +202 -0
- src/metrics.py +156 -0
- src/normalize.py +104 -0
- src/occupancy_encoder.py +248 -0
- src/occupancy_mlp.py +87 -0
- src/viewer/__init__.py +1 -0
- src/viewer/infer_job.py +536 -0
- src/viewer/mesh_access.py +52 -0
- 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
|