Spaces:
Running on Zero
Running on Zero
Download src/data_npz.py from guyPerry/scatteringnet: direct link, hf CLI and curl.
- Browser
- Download file 15.3 kB
-
https://huggingface.co/spaces/guyPerry/scatteringnet/resolve/main/src/data_npz.py
- Command line
-
hf download hf://spaces/guyPerry/scatteringnet/src/data_npz.py
-
curl -L -o data_npz.py https://huggingface.co/spaces/guyPerry/scatteringnet/resolve/main/src/data_npz.py
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)) | |