scatteringnet / src /data_npz.py
scatteringnet-space
Slim Gradio Space: infer + demo only
bc4c433
Raw History Blame Contribute Delete
15.3 kB
"""Load occupancy query points and labels from scatter NPZ files.
:func:`load_points_labels` reads **one** file. A catalog resolver lists
many NPZs (glob or explicit paths) without training.
``load_points_labels`` still returns only ``points`` and ``labels``.
``load_points_labels_mesh`` also resolves ``mesh_path`` against ``data_dir``.
"""
from __future__ import annotations
import glob as globlib
import random
from pathlib import Path
from typing import Sequence
import numpy as np
from numpy.typing import NDArray
# Labels are converted here (not deferred to the Dataset) so every caller gets
# the same dtypes: float32 XYZ and float32 {0, 1} occupancy.
PointsArray = NDArray[np.float32]
LabelsArray = NDArray[np.float32]
def load_points_labels(path: Path) -> tuple[PointsArray, LabelsArray]:
"""
Read query coordinates and inside/outside labels from one NPZ file.
Parameters
----------
path:
Path to a ``.npz`` with arrays ``points`` ``(N, 3)`` and
``labels`` ``(N,)`` (typically uint8 0/1).
Returns
-------
points:
``float32`` array of shape ``(N, 3)``.
labels:
``float32`` array of shape ``(N,)`` with values in ``{0.0, 1.0}``
(0 = outside, 1 = inside).
"""
npz_path = Path(path)
if not npz_path.is_file():
raise FileNotFoundError(f"NPZ not found: {npz_path}")
# allow_pickle=False: we only need numeric arrays, not object payloads.
with np.load(npz_path, allow_pickle=False) as raw:
files = set(raw.files)
if "points" not in files or "labels" not in files:
raise KeyError(
f"NPZ must contain 'points' and 'labels', got {sorted(files)} "
f"in {npz_path}"
)
points = np.asarray(raw["points"])
labels = np.asarray(raw["labels"])
if points.ndim != 2 or points.shape[1] != 3:
raise ValueError(
f"points must have shape (N, 3), got {tuple(points.shape)} in {npz_path}"
)
n = int(points.shape[0])
if labels.shape != (n,):
raise ValueError(
f"labels must have shape (N,) with N={n}, got {tuple(labels.shape)} "
f"in {npz_path}"
)
points_f32 = np.asarray(points, dtype=np.float32)
labels_f32 = np.asarray(labels, dtype=np.float32)
unique = np.unique(labels_f32)
if not np.all((unique == 0.0) | (unique == 1.0)):
raise ValueError(
f"labels must be in {{0, 1}}, got unique={unique.tolist()} in {npz_path}"
)
return points_f32, labels_f32
def count_npz_points(path: Path | str) -> int:
"""
Query count in one NPZ without keeping the arrays.
Catalog construct uses this so ``len(part)`` / ``n_points`` do not
require loading every lattice into RAM.
"""
npz_path = Path(path)
if not npz_path.is_file():
raise FileNotFoundError(f"NPZ not found: {npz_path}")
with np.load(npz_path, allow_pickle=False) as raw:
files = set(raw.files)
if "points" not in files or "labels" not in files:
raise KeyError(
f"NPZ must contain 'points' and 'labels', got {sorted(files)} "
f"in {npz_path}"
)
n = int(np.asarray(raw["points"]).shape[0])
n_y = int(np.asarray(raw["labels"]).shape[0])
if n_y != n:
raise ValueError(
f"labels must have shape (N,) with N={n}, got N={n_y} in {npz_path}"
)
return n
def count_npz_labels(path: Path | str) -> tuple[int, int]:
"""
Outside / inside counts in one NPZ without keeping the point cloud.
Used for ``pos_weight: auto`` (n_outside / n_inside on the train split).
"""
npz_path = Path(path)
if not npz_path.is_file():
raise FileNotFoundError(f"NPZ not found: {npz_path}")
with np.load(npz_path, allow_pickle=False) as raw:
if "labels" not in set(raw.files):
raise KeyError(f"NPZ must contain 'labels', got {sorted(raw.files)} in {npz_path}")
labels = np.asarray(raw["labels"]).reshape(-1)
# Labels are 0/1 occupancy; nonzero is inside.
n_in = int(np.count_nonzero(labels))
n_out = int(labels.size) - n_in
return n_out, n_in
def read_npz_mesh_path(path: Path | str) -> str:
"""
Read the stored ``mesh_path`` string from one occupancy NPZ.
Step 2 writes a ``data_dir``-relative POSIX path (for example
``meshes/Primitives/Sphere/sphere_r0p5_sa16_sh16.obj``).
Parameters
----------
path:
Occupancy ``.npz`` that contains ``mesh_path``.
Returns
-------
str
Stored path string (relative or absolute). Not resolved here.
"""
npz_path = Path(path)
if not npz_path.is_file():
raise FileNotFoundError(f"NPZ not found: {npz_path}")
# allow_pickle=True: some exports store a 0-d string / object array.
with np.load(npz_path, allow_pickle=True) as raw:
if "mesh_path" not in raw.files:
raise KeyError(f"NPZ has no 'mesh_path' in {npz_path}")
stored = str(np.asarray(raw["mesh_path"]).item()).strip()
if not stored:
raise ValueError(f"mesh_path is empty in {npz_path}")
return stored
def resolve_mesh_path(stored: str, data_dir: Path | str) -> Path:
"""
Resolve a stored ``mesh_path`` against ``data_dir``.
Relative entries are joined to ``data_dir``. Absolute entries are
used as-is. Missing files raise ``FileNotFoundError``.
Parameters
----------
stored:
Value from :func:`read_npz_mesh_path`.
data_dir:
Dataset root (``config.yaml`` ``data_dir``).
Returns
-------
Path
Existing resolved mesh file.
"""
text = str(stored).strip()
if not text:
raise ValueError("mesh_path is empty")
item = Path(text)
root = Path(data_dir)
resolved = item if item.is_absolute() else (root / item)
resolved = resolved.resolve()
if not resolved.is_file():
raise FileNotFoundError(f"mesh not found: {resolved} (stored={text!r})")
return resolved
def load_points_labels_mesh(
path: Path | str,
data_dir: Path | str,
) -> tuple[PointsArray, LabelsArray, Path]:
"""
Read occupancy arrays and resolve the source OBJ.
Keeps :func:`load_points_labels` unchanged (xyz + labels only).
Parameters
----------
path:
Occupancy ``.npz`` with ``points``, ``labels``, and ``mesh_path``.
data_dir:
Root used to resolve a relative ``mesh_path``.
Returns
-------
points, labels, mesh_path:
Same arrays as :func:`load_points_labels`, plus the existing OBJ.
"""
npz_path = Path(path)
points, labels = load_points_labels(npz_path)
stored = read_npz_mesh_path(npz_path)
mesh_path = resolve_mesh_path(stored, data_dir)
return points, labels, mesh_path
def _is_combo_npz(path: Path) -> bool:
"""True when the filename looks like a combo dump (excluded from the catalog)."""
return "combo" in path.name.lower()
def shape_key(path: Path) -> str:
"""
Group NPZs that belong to the same mesh.
Uses the stem before ``__`` (dataset_builder tag), else the full stem.
Parameters
----------
path:
NPZ path.
Returns
-------
str
Stable key for ``max_files_per_shape``.
"""
stem = Path(path).stem
if "__" in stem:
return stem.split("__", 1)[0]
return stem
def _cap_per_shape(
paths: Sequence[Path],
max_files_per_shape: int | None,
) -> list[Path]:
"""Keep at most ``max_files_per_shape`` files per :func:`shape_key` (sorted order)."""
if max_files_per_shape is None:
return list(paths)
if max_files_per_shape < 1:
raise ValueError(f"max_files_per_shape must be >= 1 or None, got {max_files_per_shape}")
counts: dict[str, int] = {}
out: list[Path] = []
for path in paths:
key = shape_key(path)
taken = counts.get(key, 0)
if taken >= max_files_per_shape:
continue
counts[key] = taken + 1
out.append(path)
return out
def _glob_npz(root: Path, pattern: str) -> list[Path]:
"""Match ``pattern`` under ``root`` (``*`` / ``**``)."""
full = str(root / pattern)
recursive = "**" in pattern.replace("\\", "/")
found = globlib.glob(full, recursive=recursive)
return [Path(p).resolve() for p in found if Path(p).is_file()]
def _is_parameterized_stem(path: Path) -> bool:
"""
Maya catalog names are ``family_param_...``. Varied one-off stems
(``Cone.obj`` → ``Cone__occupancy.npz``) have no ``_`` in the shape key
and must not ride along when Windows glob is case-insensitive.
"""
return "_" in shape_key(path)
def _subsample_shapes(
paths: Sequence[Path],
max_shapes: int,
seed: int,
) -> list[Path]:
"""Keep NPZs for at most ``max_shapes`` unique :func:`shape_key` values."""
if max_shapes < 1:
raise ValueError(f"max_shapes must be >= 1, got {max_shapes}")
keys: list[str] = []
seen: set[str] = set()
for path in paths:
key = shape_key(path)
if key in seen:
continue
seen.add(key)
keys.append(key)
if max_shapes >= len(keys):
return list(paths)
# Sort then sample so the same seed always picks the same meshes.
chosen_keys = set(random.Random(int(seed)).sample(sorted(keys), max_shapes))
return [path for path in paths if shape_key(path) in chosen_keys]
def resolve_npz_catalog(
data_dir: Path | str,
*,
npz_glob: str = "exports/dataset/*.npz",
npz_paths: Sequence[str | Path] | None = None,
npz_catalog: Sequence[tuple[str, int | None]] | None = None,
max_files_per_shape: int | None = 2,
exclude_combo: bool = True,
seed: int = 1,
) -> list[Path]:
"""
Resolve occupancy NPZ paths under ``data_dir`` (no point loading).
Priority: explicit ``npz_paths``, else ``npz_catalog`` (union of globs),
else ``npz_glob``. Relative entries are joined to ``data_dir``.
Missing files in ``npz_paths`` raise ``FileNotFoundError``.
Parameters
----------
data_dir:
Dataset root (``config.yaml`` ``data_dir``).
npz_glob:
Single glob relative to ``data_dir`` (``*`` and ``**`` allowed).
npz_paths:
Explicit relative or absolute NPZ paths. Empty / None → use glob(s).
npz_catalog:
``(glob, max_shapes)`` rows. ``max_shapes`` is unique meshes after
the per-shape file cap; ``None`` keeps every mesh the glob hits.
max_files_per_shape:
Cap per :func:`shape_key` after sort. ``None`` = no cap.
exclude_combo:
Drop filenames containing ``combo``.
seed:
RNG for ``max_shapes`` subsampling (YAML ``seed``).
Returns
-------
list[Path]
Sorted existing ``.npz`` files.
"""
root = Path(data_dir)
chosen: list[Path]
if npz_paths:
chosen = []
for raw in npz_paths:
item = Path(raw)
resolved = item if item.is_absolute() else (root / item)
if not resolved.is_file():
raise FileNotFoundError(f"NPZ not found: {resolved}")
chosen.append(resolved.resolve())
elif npz_catalog:
# Union in YAML order. Same file from two globs is kept once.
seen: set[Path] = set()
chosen = []
for pattern, max_shapes in npz_catalog:
hit = _glob_npz(root, str(pattern))
if exclude_combo:
hit = [p for p in hit if not _is_combo_npz(p)]
hit = [
p
for p in hit
if p.suffix.lower() == ".npz" and _is_parameterized_stem(p)
]
hit = sorted(hit)
hit = _cap_per_shape(hit, max_files_per_shape)
if max_shapes is not None:
hit = _subsample_shapes(hit, int(max_shapes), int(seed))
for path in hit:
if path in seen:
continue
seen.add(path)
chosen.append(path)
else:
chosen = _glob_npz(root, npz_glob)
npz_only = [p for p in chosen if p.suffix.lower() == ".npz"]
if exclude_combo:
npz_only = [p for p in npz_only if not _is_combo_npz(p)]
if npz_catalog and not npz_paths:
# Already capped per glob; keep YAML union order (not a global sort).
capped = npz_only
else:
npz_only = sorted(npz_only)
capped = _cap_per_shape(npz_only, max_files_per_shape)
if not capped:
raise FileNotFoundError(
f"No occupancy NPZ files matched under {root} "
f"(glob={npz_glob!r}, catalog={bool(npz_catalog)}, "
f"explicit={bool(npz_paths)})"
)
return capped
def _summarize(points: PointsArray, labels: LabelsArray) -> str:
n = int(points.shape[0])
inside = float(labels.mean()) if n else float("nan")
xyz_min = points.min(axis=0) if n else np.full(3, np.nan, dtype=np.float32)
xyz_max = points.max(axis=0) if n else np.full(3, np.nan, dtype=np.float32)
return (
f"N={n}\n"
f"inside_fraction={inside:.6f}\n"
f"xyz_min={xyz_min.tolist()}\n"
f"xyz_max={xyz_max.tolist()}"
)
# Default smoke-check file from the v2 plan (dataset_test sphere).
_SAMPLE_RELATIVE = Path("exports") / "dataset_test" / "sphere__raycast_z_raut_s0.15_inout.npz"
if __name__ == "__main__":
import sys
from scatteringnet.config import load_config
cfg = load_config()
if "--catalog" in sys.argv:
paths = resolve_npz_catalog(
cfg.data_dir,
npz_glob=cfg.npz_glob,
npz_paths=cfg.npz_paths or None,
npz_catalog=cfg.npz_catalog or None,
max_files_per_shape=cfg.max_files_per_shape,
seed=cfg.seed,
)
print(f"data_dir={cfg.data_dir}")
print(f"npz_glob={cfg.npz_glob}")
print(f"npz_catalog={list(cfg.npz_catalog)}")
print(f"max_files_per_shape={cfg.max_files_per_shape}")
print(f"files={len(paths)}")
# Load a few files only — the full catalog can be thousands of NPZs.
preview = paths[:3]
n_all = 0
n_in = 0
for path in preview:
pts, labs = load_points_labels(path)
n = int(pts.shape[0])
inside = int((labs == 1.0).sum())
n_all += n
n_in += inside
print(
f" {path.name} N={n} inside={inside} outside={n - inside}"
)
print(
f"preview_files={len(preview)} preview_N={n_all} "
f"preview_inside={n_in} preview_outside={n_all - n_in}"
)
from scatteringnet.dataset import OccupancyMultiNpzDataset, make_dataloader
ds = OccupancyMultiNpzDataset(paths)
loader = make_dataloader(ds.parts[0], batch_size=8, shuffle=False)
xyz, y = next(iter(loader))
print(f"files={len(ds)} n_points={ds.n_points} parts={len(ds.parts)}")
print(f"batch xyz={tuple(xyz.shape)} y={tuple(y.shape)}")
else:
sample = cfg.data_dir / _SAMPLE_RELATIVE
pts, labs = load_points_labels(sample)
print(f"file={sample}")
print(_summarize(pts, labs))